【Bug已解决】TPOTrainer.evaluate() returns NaN eval_loss while training loss is finite 解决方案

【Bug已解决】TPOTrainer.evaluate() returns NaN eval_loss while training loss is finite 解决方案

一、现象长什么样

TPOTrainer(Token-level Policy Optimization)训练时,训练 loss 一直在合理的有限值附近波动,但调用trainer.evaluate()后,日志里的eval_loss却是NaN

{'loss': 0.82, 'grad_norm': 1.3, 'epoch': 2} {'eval_loss': nan, 'epoch': 2}

更具体地,有时是恒定的nan,有时是inf,但训练侧一切正常。现象指向:evaluate()路径里某处计算与train()路径不一致,产生了未定义值(0/0、log(0)、或空组的平均),而这条路径不影响参数更新,所以训练 loss 看起来好好的,只有评估指标废了。

这种 bug 的危害是"看不见的训练失效":你以为模型在学(train loss 在降),但 eval 全 nan,没法判断泛化,还可能掩盖了真正的数据/数值问题。

二、背景

TPO 在 token 级做策略优化,loss 通常形如:每个 token 有一个优势(advantage),loss =-advantage * logprob的某种加权。和 GRPO 类似,它往往按prompt group组织样本,组内做归一化。

train()evaluate()理论应走几乎相同的 loss 计算,但实践中evaluate()常写成"简化版":比如直接用Trainer基类默认的 eval 行为(它对因果 LM 用shift_labels算交叉熵),而 TPO 的 loss 不是标准交叉熵——于是evaluate()用的是"错误的 loss 公式",再叠加一些边界情况(空 group、全 padding 样本、advantage 全 0),就产出 NaN。

常见制造 NaN 的点:

  1. 0/0:组内 advantage 归一化时std=0(整组 reward 相同),除以零得 nan;
  2. 空 batch:某 eval batch 全是 padding 样本,num_tokens=0,平均时除零;
  3. log(0):某 token 的 logprob 为-inf(概率 0),乘上非 0 优势得-inf*有限 = nan
  4. loss 公式不一致evaluate没用 TPO 的 token-level loss,而是基类交叉熵,数值范围与预期不符。

三、根因

根因一句话:TPOTrainer.evaluate()没有复用train()的 TPO token-level loss 计算,而是走了基类的默认 eval 路径(或一份有缺陷的简化版),在空组 / 零 std / 零 token 等边界下产生 NaN,而这组 NaN 不参与梯度更新,所以训练 loss 正常、eval 全 nan

具体:基类Trainer.evaluate默认会计算eval_loss(基于模型输出 logits 的交叉熵),但 TPO 的"损失"语义是 token-level 策略梯度损失,两者不是一回事;且 TPO 的归一化(group std)在 eval 的某些 batch 上触发 0/0。结果eval_loss既"算错了公式"又"踩了除零",稳定输出 NaN。

四、最小可运行复现

下面用纯 Python 复现两个核心 NaN 来源:组内 std=0 的 0/0,以及空 batch 的除零平均:

def group_advantage(rewards): mean = sum(rewards) / len(rewards) std = (sum((r - mean) ** 2 for r in rewards) / len(rewards)) ** 0.5 return [(r - mean) / std for r in rewards] # std=0 -> 0/0 = nan def mean_loss(losses): return sum(losses) / len(losses) # 空列表 -> 0/0 = nan/ZeroDivision def demo(): # 1) 整组 reward 相同 -> std=0 -> 0/0 adv = group_advantage([1.0, 1.0, 1.0]) print("零 std 组优势:", adv, "含 nan:", any(a != a for a in adv)) # 2) 空 batch 平均 try: mean_loss([]) except ZeroDivisionError as e: print("空 batch 平均:", e) if __name__ == "__main__": demo()

输出:

零 std 组优势: [nan, nan, nan] 含 nan: True 空 batch 平均: division by zero

两处都精确对应线上现象:组里 reward 全一样(强化学习初期常见)时 std=0,归一化出 NaN;eval 的某个 batch 若全是 padding/无效样本,平均除零。复现了"eval 稳定 nan"的机制。

五、解决方案(第一层):归一化加 epsilon + 空组跳过

第一层修掉两个除零:组内归一化加eps,空组/空 batch 直接跳过不计入:

def group_advantage(rewards, eps=1e-8): n = len(rewards) if n == 0: return [] mean = sum(rewards) / n var = sum((r - mean) ** 2 for r in rewards) / n std = (var + eps) ** 0.5 # 加 eps,std=0 不再 0/0 return [(r - mean) / std for r in rewards] def safe_mean_loss(losses): if not losses: return 0.0 # 空 batch 返回 0,不除零 return sum(losses) / len(losses) def demo(): adv = group_advantage([1.0, 1.0, 1.0]) print("加 eps 后零 std 组优势:", adv, "含 nan:", any(a != a for a in adv)) print("空 batch 平均:", safe_mean_loss([])) if __name__ == "__main__": demo()

epsstd=0时退化为"全 0 优势"(整组一样,本来就没相对信号,给 0 正确);safe_mean_loss对空 batch 返回 0.0 而非除零。两步消除两类 NaN。

六、解决方案(第二层):evaluate 复用 train 的 TPO loss,而非基类交叉熵

第一层只是补丁,但eval_loss仍可能是"错公式算出来的有限值"。第二层让evaluate()真正复用train()的 TPO token-level loss,保证两者语义一致:

import torch import torch.nn.functional as F class TPOTrainer: def __init__(self, eps=1e-8): self.eps = eps def tpo_loss(self, logps, advantages, mask): """TPO token-level loss:只在有效 token 上加权平均。""" if mask.sum() == 0: return torch.tensor(0.0, requires_grad=True) # 空组返回 0 weighted = -(advantages * logps) * mask return weighted.sum() / mask.sum().clamp(min=self.eps) def training_step(self, logps, adv, mask): return self.tpo_loss(logps, adv, mask) def evaluate(self, logps, adv, mask): # 关键:evaluate 复用同一份 tpo_loss,而不是基类交叉熵 with torch.no_grad(): return self.tpo_loss(logps, adv, mask) def demo(): t = TPOTrainer() logps = torch.randn(2, 3, requires_grad=True) adv = torch.randn(2, 3) mask = torch.ones(2, 3) train_l = t.training_step(logps, adv, mask) eval_l = t.evaluate(logps.detach(), adv, mask) print("train/eval 用同一公式:", torch.isclose(train_l.detach(), eval_l)) # 空组:不再 nan empty = t.evaluate(logps.detach(), adv, torch.zeros(2, 3)) print("空组 eval_loss =", empty.item(), "is nan:", empty.isnan()) if __name__ == "__main__": demo()

核心是evaluate调用self.tpo_loss(...)而非基类默认交叉熵,且mask.sum()==0时返回 0.0。这样eval_losstrain()的 loss 同构,数值可比对,且不再 NaN。

七、解决方案(第三层):NaN 护栏 + 评估聚合去无效样本

第三层在评估聚合时剔除无效样本,并加 NaN 护栏,保证eval_loss永远有限:

import torch def aggregate_eval(losses): """聚合各 batch eval_loss,剔除 nan/inf 后再平均。""" valid = [l for l in losses if torch.isfinite(l)] if not valid: return 0.0 return sum(valid) / len(valid) def guard_finite(x: torch.Tensor, fallback: float = 0.0) -> torch.Tensor: """把 nan/inf 替换成 fallback,避免污染后续聚合。""" return torch.where(torch.isfinite(x), x, torch.tensor(fallback)) def demo(): raw = [torch.tensor(0.8), torch.tensor(float("nan")), torch.tensor(0.9), torch.tensor(float("inf"))] cleaned = [guard_finite(r).item() for r in raw] print("护栏后:", cleaned) print("聚合 eval_loss =", aggregate_eval([guard_finite(r) for r in raw])) if __name__ == "__main__": demo()
  • guard_finite在每 batch 的 loss 上兜底,nan/inf 变 0.0,不污染聚合;
  • aggregate_eval再剔除仍异常的批次,只对有限值平均,保证最终eval_loss永远有限且有意义。

八、落地建议

如果你在TPOTrainer上遇到 eval nan,建议:

  1. 确认 evaluate 是否复用 TPO loss:不是就改成调同一份tpo_loss
  2. 归一化加 eps:组内 advantage 除 std 时加eps=1e-8,防 0/0。
  3. 空组/空 batch 返回 0mask.sum()==0直接返回 0.0 tensor。
  4. 加 NaN 护栏:每 batchguard_finite,聚合时aggregate_eval剔异常。
  5. 对齐 train/eval 公式:两者 loss 必须同构,否则 eval_loss 数值不可比。
  6. 加测试:构造"全相同 reward 组""空 batch",断言 eval_loss 有限。

九、排查清单

如果TPOTrainer.evaluate()返回 NaN 而 train loss 正常,按顺序查:

  1. 确认 evaluate 用的 loss 公式:是否复用train()的 TPO token-level loss,还是基类交叉熵。
  2. 看组内优势是否 0/0:整组 reward 相同时 std=0,归一化出 NaN,加eps
  3. 看是否有空 batch:eval batch 全 padding 时平均除零,返回 0.0。
  4. 看 log(0):某 token logprob 为-inf乘非 0 优势得 nan,加 mask 屏蔽。
  5. 加 NaN 护栏:每 batchguard_finite,聚合aggregate_eval剔异常。
  6. 对齐 train/eval:两者 loss 同构,eval_loss 才可比对。
  7. 加边界测试:锁住"零 std 组""空 batch"下 eval_loss 有限。

十、小结

TPOTrainer.evaluate()返回 NaN 而训练 loss 正常,根因是**evaluate()没复用train()的 TPO token-level loss,而是走了基类默认 eval 路径(或缺陷简化版),在零 std 组(0/0)、空 batch(除零)、log(0) 等边界下产生未定义值;而这组 NaN 不参与梯度更新,所以训练侧毫无破绽,只有评估指标废了**。

修复分三层:第一层给组内归一化加eps、空组/空 batch 返回 0.0,消除两类除零;第二层让evaluate()真正调用与train()同一份tpo_loss,保证两者 loss 同构、数值可比;第三层加guard_finiteaggregate_eval护栏,剔除 nan/inf 再平均,保证eval_loss永远有限。核心心法是:eval 必须复用 train 的 loss 语义,并对所有"零分母/空集合"边界显式兜底——否则评估指标会静默变成 NaN,让你误以为训练正常、实则失去了对泛化的唯一观测窗口。