【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 的点:
- 0/0:组内 advantage 归一化时
std=0(整组 reward 相同),除以零得 nan; - 空 batch:某 eval batch 全是 padding 样本,
num_tokens=0,平均时除零; - log(0):某 token 的 logprob 为
-inf(概率 0),乘上非 0 优势得-inf*有限 = nan; - 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()eps让std=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_loss与train()的 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,建议:
- 确认 evaluate 是否复用 TPO loss:不是就改成调同一份
tpo_loss。 - 归一化加 eps:组内 advantage 除 std 时加
eps=1e-8,防 0/0。 - 空组/空 batch 返回 0:
mask.sum()==0直接返回 0.0 tensor。 - 加 NaN 护栏:每 batch
guard_finite,聚合时aggregate_eval剔异常。 - 对齐 train/eval 公式:两者 loss 必须同构,否则 eval_loss 数值不可比。
- 加测试:构造"全相同 reward 组""空 batch",断言 eval_loss 有限。
九、排查清单
如果TPOTrainer.evaluate()返回 NaN 而 train loss 正常,按顺序查:
- 确认 evaluate 用的 loss 公式:是否复用
train()的 TPO token-level loss,还是基类交叉熵。 - 看组内优势是否 0/0:整组 reward 相同时 std=0,归一化出 NaN,加
eps。 - 看是否有空 batch:eval batch 全 padding 时平均除零,返回 0.0。
- 看 log(0):某 token logprob 为
-inf乘非 0 优势得 nan,加 mask 屏蔽。 - 加 NaN 护栏:每 batch
guard_finite,聚合aggregate_eval剔异常。 - 对齐 train/eval:两者 loss 同构,eval_loss 才可比对。
- 加边界测试:锁住"零 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_finite与aggregate_eval护栏,剔除 nan/inf 再平均,保证eval_loss永远有限。核心心法是:eval 必须复用 train 的 loss 语义,并对所有"零分母/空集合"边界显式兜底——否则评估指标会静默变成 NaN,让你误以为训练正常、实则失去了对泛化的唯一观测窗口。