【Bug已解决】DPOTrainer: ref_model is None during eval when gradient_checkpointing=True 解决方案 【Bug已解决】DPOTrainer: ref_model is None during eval when gradient_checkpointingTrue 解决方案原始报错DPOTrainer: ref_model is None during eval when gradient_checkpointingTrue 场景用DPOTrainer做偏好训练开了gradient_checkpointingTrue省显存。训练阶段 ref_model参考模型正常存在但进入评估eval阶段时ref_model变成了None。eval 需要 ref_model 算参考对数概率于是报错或 ref log probs 缺失评估结果不可信。 关键词DPOTrainer、ref_model、gradient_checkpointing、评估阶段、模型状态管理、训练/评估切换。一、现象长什么样配置gradient_checkpointingTrue后训练阶段正常ref_model 在损失能算训练若干步触发 evaluation_loop评估eval 里取self.ref_model发现是Noneeval 直接报错或静默用None算 ref log probs 得到错误结果关掉gradient_checkpointing后 eval 正常ref_model 在。用户困惑为什么训练时好好的一评估 ref_model 就没了根因是梯度检查点相关的模型管理逻辑在切换到 eval 时把 ref_model 置空或没保留而训练路径恰好不依赖那个被清的状态。二、背景DPO 为什么 eval 也需要 ref_modelDPO 的损失同时依赖策略模型policy和参考模型ref_modelref_model 固定提供锚点对数概率。这不只是训练时需要——评估时同样需要 ref_model来算 eval 损失/准确率否则评估指标如 eval accuracy无法正确衡量相对 ref 的偏好提升。gradient_checkpointingTrue是为了省训练显存训练时在反向传播时重新计算中间激活不全程保存。这个优化只应作用于训练主路径且只影响显存/计算不应改变ref_model 对象是否存在。bug 在于某个实现把 gradient_checkpointing 的副作用扩散到了 ref_model 的生命周期导致 eval 时它没了。三、根因ref_model 生命周期被 eval 切换破坏根因拆解eval 切换清模型进入 eval 时某段逻辑可能为省显存把ref_model置None或卸载却没在 eval 结束/开始时恢复。gradient_checkpointing 副作用gc 配置在模型管理里被当成训练专属eval 路径误以为不需要 ref_model 而跳过保留。未校验即用eval 直接用self.ref_model没有是否存在的前置检查None 直接传进前向。训练掩盖训练路径走的是另一条已确保 ref_model 的代码所以训练正常、eval 才爆。无重建ref_model 为 None 时没有从已保存的权重重建直接废掉 eval。下面用最小模型复现eval 切换把 ref_model 清成 None再给修复。四、最小可运行复现class DPOTrainer: def __init__(self, gradient_checkpointing): self.gradient_checkpointing gradient_checkpointing self.ref_model ref_weights # 训练时存在 self.mode train def enter_eval(self): self.mode eval if self.gradient_checkpointing: # 错误为省显存把 ref_model 清掉且 eval 也需要它 self.ref_model None def evaluate(self): if self.ref_model is None: raise RuntimeError(ref_model is None during eval) return eval_ok if __name__ __main__: t DPOTrainer(gradient_checkpointingTrue) t.enter_eval() try: t.evaluate() except RuntimeError as e: print(eval 失败:, e) # ref_model is None during eval运行抛错——enter_eval 把 ref_model 清 None而 eval 真的需要它。五、方案eval 前确保 ref_model 就绪第一层进入 eval 前显式保证 ref_model 存在gradient_checkpointing 不应清它class DPOTrainerSafe: def __init__(self, gradient_checkpointing, ref_weights_path): self.gradient_checkpointing gradient_checkpointing self.ref_weights_path ref_weights_path self.ref_model self._load_ref() # 始终持有 def _load_ref(self): # 真实环境从权重加载冻结的 ref 模型 return {weights: self.ref_weights_path, frozen: True} def enter_eval(self): self.mode eval # 关键gradient_checkpointing 只影响训练主路径显存不清 ref_model if self.ref_model is None: self.ref_model self._load_ref() # 缺失则重建 def evaluate(self): if self.ref_model is None: raise RuntimeError(ref_model 仍为空eval 无法进行) return eval_ok if __name__ __main__: t DPOTrainerSafe(gradient_checkpointingTrue, ref_weights_path/m/ref) t.enter_eval() print(t.evaluate()) # eval_okref_model 在eval 切换不再清 ref_model缺失则重建gradient_checkpointing 的显存优化与 ref_model 生命周期解耦。六、方案gradient_checkpointing 作用域隔离第二层把 gradient_checkpointing 的启用严格限定在主模型policy的训练路径不波及 ref_modelref 本就冻结、不反向class ModelManager: def __init__(self, gc): self.gc gc def configure(self, policy, ref): # 主模型训练时按需开 gc policy[gradient_checkpointing] self.gc # ref 冻结、不训永远不开 gc也不因 gc 被卸载 ref[gradient_checkpointing] False ref[requires_grad] False return policy, ref if __name__ __main__: mm ModelManager(gcTrue) policy, ref mm.configure({name: policy}, {name: ref}) print(policy gc:, policy[gradient_checkpointing]) # True训练用 print(ref gc:, ref[gradient_checkpointing]) # False冻结不用 print(ref 仍存在:, ref[name]) # ref 不被清作用域隔离让 gc 的显存优化只服务于训练主模型ref_model 作为冻结组件一直可用eval 阶段自然有它。七、方案缺失时优雅报错/重建 显式校验第三层eval 入口加前置校验ref_model 为 None 时要么从权重重建、要么给清晰错误绝不拿 None 进前向def ensure_ref_for_eval(trainer): if trainer.ref_model is None: if getattr(trainer, ref_weights_path, None): trainer.ref_model trainer._load_ref() # 重建 else: raise RuntimeError( eval 需要 ref_model但它为 None 且无权重可重建 请检查 gradient_checkpointing 切换逻辑是否误清了 ref_model) return trainer.ref_model if __name__ __main__: t DPOTrainerSafe(gradient_checkpointingTrue, ref_weights_path/m/ref) t.ref_model None # 模拟被误清 t.enter_eval() ensure_ref_for_eval(t) # eval 前重建 print(eval 可用 ref:, t.ref_model is not None)前置校验 重建 清晰错误把eval 时 ref_model 为 None从崩溃/静默变成可恢复或明确报错。八、验证把eval 时 ref_model 就绪锁进测试def test_ref_model_present_in_eval(): t DPOTrainerSafe(gcTrue, ref_weights_path/m/ref) t.enter_eval() assert t.ref_model is not None assert t.evaluate() eval_ok def test_ref_rebuilt_when_cleared(): t DPOTrainerSafe(gcTrue, ref_weights_path/m/ref) t.ref_model None t.enter_eval() ensure_ref_for_eval(t) assert t.ref_model is not None def test_clear_error_without_weights(): t DPOTrainerSafe(gcTrue, ref_weights_pathNone) t.ref_model None t.enter_eval() try: ensure_ref_for_eval(t) assert False except RuntimeError as e: assert ref_model in str(e) if __name__ __main__: test_ref_model_present_in_eval() test_ref_rebuilt_when_cleared() test_clear_error_without_weights() print(DPO eval ref_model 就绪测试通过。)九、排查清单eval 时 ref_model 为 None按顺序查阶段对比训练正常、eval 才 None说明是 eval 切换清了 ref_model。gc 关联关掉 gradient_checkpointing 是否恢复是则 gc 逻辑误清 ref。作用域gc 是否只应用于训练主模型还是波及冻结的 ref_model前置校验eval 入口是否检查 ref_model 是否存在None 直接进前向会崩。重建ref_model 为 None 时能否从权重重建还是直接废掉 eval生命周期ref_model 对象在 train/eval 切换间是否被保活错误信息若真的缺报错是否说明gc 切换误清 如何补权重十、小结DPOTrainer 在 gradient_checkpointingTrue 时 eval 阶段 ref_model 为 None是梯度检查点的显存优化副作用扩散到 ref_model 生命周期eval 切换把它清掉却没恢复。修复三层eval 前确保就绪进入 eval 显式保证 ref_model 存在gc 不清它作用域隔离gradient_checkpointing 只服务于训练主模型冻结的 ref 永远可用校验 重建eval 入口前置校验None 时从权重重建或清晰报错绝不拿 None 进前向。核心原则ref_model 是 DPO 训练和评估的共同依赖它的生命周期不能受训练专属优化如 gradient_checkpointing影响。把 gc 的作用域严格限定在训练主路径并在 eval 切换时保活/重建 ref_model评估才能正确算出相对参考模型的偏好指标。

本月热点