三亩地 三亩地SAN MU DI · CODE DIARY
ARTICLE DETAIL

日记详情

真实记录编程学习的某一天,欢迎挑你感兴趣的翻一翻。

论文复现环境升级,先保存可对齐的旧基线

论文复现环境升级,先保存可对齐的旧基线

论文复现环境升级,先保存可对齐的旧基线

框架或 CUDA 升级会改变算子、随机性与性能。先保存旧环境基线,升级后的差异才有解释入口。

1. 把论文配置和本地实现分开登记

论文复现首先对齐数据处理、模型配置、随机状态和度量实现。公开结果可以作为参照,但不能替代对本地实验条件的逐项核验。

论文报告值、作者代码默认值和本地改动应分别记录。升级验证只比较能够对齐的部分,不用公开结果替代本地回归。

2. 按最小闭环验证

协作时应把输入格式、配置字段、产物和复核责任写成接口约定。无法复现的部分要明确标注缺失条件,不将推测写成结论。

升级前后先在固定样本上比较损失、主要指标和检查点加载结果。若算子行为改变,应保存最小复现和环境差异,不用调参掩盖偏差。

3. 参考实现与图示

下面的 PyTorch 示例用于建立升级前后的可对齐输出。运行记录应包含框架、CUDA、驱动和随机性设置,不能只保存最终分数。

import torch import torch.nn as nn import logging from typing import Tuple logging.basicConfig(level=logging.INFO) logger = logging.getLogger("PaperReplicationValidator") class BaselineAttention(nn.Module): """标准的 PyTorch 标准注意力实现 (Baseline)""" def __init__(self, embed_dim: int, num_heads: int): super().__init__() self.mha = nn.MultiheadAttention(embed_dim, num_heads, batch_first=True) def forward(self, x: torch.Tensor) -> torch.Tensor: out, _ = self.mha(x, x, x) return out class PaperNewAttention(nn.Module): """论文宣称的优化版注意力实现 (待验证的新算子)""" def __init__(self, embed_dim: int, num_heads: int): super().__init__() self.embed_dim = embed_dim self.num_heads = num_heads self.qkv_proj = nn.Linear(embed_dim, embed_dim * 3) self.out_proj = nn.Linear(embed_dim, embed_dim) def forward(self, x: torch.Tensor) -> torch.Tensor: # 伪造一个优化后的计算过程 B, N, C = x.shape qkv = self.qkv_proj(x) # 假设论文中使用了一些缩放 Hack out = self.out_proj(x) return out def run_differential_testing(batch_size: int = 4, seq_len: int = 128, embed_dim: int = 256) -> bool: """执行差分测试,比对论文新算子与 Baseline 在不同精度下的偏离程度""" device = "cuda" if torch.cuda.is_available() else "cpu" logger.info(f"正在 {device} 设备上执行论文新算子差分校验...") # 1. 固定随机数种子,减少随机性造成的差异 torch.manual_seed(42) baseline_mod = BaselineAttention(embed_dim, 8).to(device) paper_mod = PaperNewAttention(embed_dim, 8).to(device) # 构造相同的测试 Tensor input_tensor = torch.randn(batch_size, seq_len, embed_dim, device=device, dtype=torch.float32) # 前向计算 with torch.no_grad(): base_out = baseline_mod(input_tensor) paper_out = paper_mod(input_tensor) # 2. 计算相对误差与最大绝对误差 max_abs_diff = torch.max(torch.abs(base_out - paper_out)).item() mean_abs_diff = torch.mean(torch.abs(base_out - paper_out)).item() logger.info(f"最大绝对误差 (Max Abs Diff): {max_abs_diff:.6f}") logger.info(f"平均绝对误差 (Mean Abs Diff): {mean_abs_diff:.6f}") # 3. 严格断言:若最大偏离超过 1e-3,则认为论文算子引入了未预期的数值漂移 tolerance = 1e-3 if max_abs_diff > tolerance: logger.error(f"❌ 校验失败:新算子偏离超出容忍阀值 ({tolerance})!禁止直接升级上线。") return False logger.info("✅ 差分校验通过:新算子与 Baseline 在容忍度内保持数值一致。") return True if __name__ == "__main__": # 执行测试 (预期输出失败,因为论文伪算子与 Baseline 尚未对齐权重) run_differential_testing()

4. 复核清单

  • 旧环境的依赖锁与硬件信息是否保存。
  • 数据预处理和评测实现是否保持不变。
  • 固定小样本上的中间张量是否可比。
  • 差异是否区分框架变化与本地代码改动。

升级不是重新定义复现成功

升级后结果变化不一定是退化,也不能直接算进步。先对齐数据、配置和度量,再解释框架带来的差异。

← 返回列表