跳转至

十八:训推不一致

来源:http://mp.weixin.qq.com/s?__biz=MzYyNTk3Njg1NA==&mid=2247484127&idx=1&sn=2455ecd7785cc9925cb59d281bec50d7&chksm=f01eb1a6c76938b0bf4937b174089d1cc082d126cf10dc56c46e6b403f4d4905aae2f27689d2#rd

1. 学习定位

训推不一致(training-inference mismatch)指模型在训练阶段和推理/采样阶段所处的计算条件、数据分布、解码方式、数值精度、系统实现或策略定义不一致,导致训练优化的目标与真实推理行为之间出现偏差。 在大模型强化学习中,训推不一致尤其重要。RL 训练依赖 rollout 采样、log probability、importance ratio、KL、reward 和 advantage。如果 rollout 端和 training 端对同一段 token 计算出的概率不一致,那么 PPO/GRPO/GSPO/SAPO 里的 ratio 会失真,进而导致策略更新方向和强度错误。 本日知识链路:

训练目标和推理行为不一致
-> 数据分布不一致
-> 解码策略不一致
-> 数值精度不一致
-> rollout engine 和 training engine 不一致
-> logprob / KL / ratio 不一致
-> RL policy update 偏离真实 on-policy
-> 通过精度对齐、recompute logprob、importance sampling、GSPO/R3 等方法缓解

2. 训推不一致的基本定义

广义训推不一致包括: - 训练时使用 teacher forcing,推理时自回归生成。

  • 训练数据分布和线上用户分布不同。

  • 训练时使用固定 prompt 模板,推理时模板变化。

  • 训练时使用高精度计算,推理时使用低精度或量化。

  • 训练时使用一种 attention/kernel,推理时使用另一种 kernel。

  • RL rollout 时使用推理引擎,training 时使用训练框架重新计算 logprob。

对 LLM RL 来说,最关键的是:

生成 token 的行为策略概率
和训练更新时使用的 old_logprob / ref_logprob / new_logprob
必须对应同一个策略定义和同一个输入上下文。
如果不一致,importance ratio:

r_t = exp(logp_new_t - logp_old_t)
就不再表示“新策略相对采样旧策略的概率变化”。

3. 传统 NLP 中的 Exposure Bias

经典序列建模中的训推不一致通常称为 exposure bias。 训练时:

P(y_t | x, ground_truth y_<t)
模型每一步的前文y<t全部是真实标注的正确 token,全程处于标准正确分布中学习。

推理时:

P(y_t | x, generated y_&lt;t)
模型前文y<t是自己上一步生成的 token,并非真实标签。训练时模型看到的是正确历史,推理时模型要面对自己之前生成的 token。一旦前面生成错误,后续状态会偏离训练分布,错误可能累积。 在 RLHF 和推理 RL 中,问题更复杂,因为训练数据往往来自当前 policy 的自回归 rollout,但 rollout 引擎与训练引擎仍可能不一致。

4. 大模型 RL 中的训推不一致

大模型 RL 通常分为两个阶段:

rollout:
&nbsp; 用推理引擎生成 responses,记录 tokens、logprobs、rewards。

training:
&nbsp; 用训练框架重新前向,计算 new_logprobs、old_logprobs、ref_logprobs、loss。
常见系统组合: - rollout 使用 vLLM/SGLang 等高吞吐推理引擎。

  • training 使用 PyTorch/DeepSpeed/Megatron/FSDP。

  • rollout 使用 FP8/FP16/BF16 KV cache 或特定 attention kernel。

  • training 使用 BF16/FP32 accumulation 和不同 kernel。

如果同一模型权重、同一 prompt 和同一 response,在 rollout 端和 training 端得到的 logprob 不同,就产生 logprob mismatch。

5. Logprob Mismatch

logprob mismatch 是 LLM RL 中最直接的训推不一致形式。 对同一个 token:

logp_rollout = log pi_rollout(a_t | s_t)
logp_train &nbsp; = log pi_train(a_t | s_t)
理想情况下:

logp_rollout ≈ logp_train
如果差异很大:

delta_logp = logp_train - logp_rollout
PPO ratio 会被扭曲:

r_t = exp(logp_new - logp_old)

img

old_logp 来自 rollout engine,而 new_logp 来自 training engine,两者存在系统偏差,那么 ratio 不只反映策略更新,还混入了系统差异。

6. True On-policy、TIS 与 MIS

在 LLM RL 系统讨论中,经常出现三类处理方式: True on-policy:

rollout 和 training 使用完全一致的模型、精度、kernel、logprob 计算。使用同一引擎做生成和训练。
old_logp 与真实采样策略一致。不需要任何重要性采样修正。
TIS(Trajectory-level Importance Sampling):

对整条 response 计算重要性权重,修正 rollout/training 分布差异。

整条句子生成时的策略,和训练时的策略有点不一样→&nbsp;给整句话乘一个权重,整体修正。适用于引擎不一致(vLLM + PyTorch)、用旧策略数据训练新策略的情况
MIS(Mini-batch Importance Sampling 或 token/minibatch 级修正,具体命名依系统而异):

在更细粒度上修正 logprob mismatch 或 off-policy 数据偏差。
- 按 token 修正

  • 或按 mini-batch 修正 名字不固定:

  • Token-level Importance Sampling

  • Mini-batch IS

  • Step-wise IS

核心思想都是:如果采样分布和训练分布不完全一致,要么让它们一致,要么用重要性采样修正。 方式 粒度 解决什么 是否需要一致引擎 训练稳定性 速度 True On-policy完全一致 所有偏差 必须一致 ⭐⭐⭐⭐⭐ 最慢 TIS整段轨迹 分布偏移 不需要 ⭐⭐⭐ 快 MIStoken / 小批量 细粒度 logprob 偏差 不需要 ⭐⭐⭐⭐ 较快

7. 数值精度导致的不一致

LLM 训练和推理常用不同数值格式: - FP32:精度高,成本高。

  • FP16:范围较窄,容易 overflow/underflow。

  • BF16:指数范围接近 FP32,尾数更少。

  • FP8/INT8/INT4:推理或训练加速,误差更大。

在 softmax/logprob 计算中,小的 logits 差异会影响概率:

logp = logits[action] - logsumexp(logits)

img

如果 logits 因精度、kernel 或量化出现微小差异,top token 的选择、logprob、KL、ratio 都可能变化。对于长序列,token-level 小误差会累积成 sequence-level 大偏差。

8. BF16 与 FP16 的差异

FP16 和 BF16 都是 16-bit 格式,但分配不同:

FP16:
&nbsp; 1 sign bit, 5 exponent bits指数位, 10 mantissa bits尾数位

BF16:
&nbsp; 1 sign bit, 8 exponent bits, 7 mantissa bits
直觉: - FP16 尾数更多,局部精度更高。

  • BF16 指数更多,动态范围更大,更不容易 overflow/underflow。

在大模型训练中,BF16 常更稳,因为梯度、激活、logits 的动态范围大。FP16 可能需要 loss scaling。 在训推不一致语境中,关键不是“谁绝对更好”,而是:

rollout 和 training 是否使用相同或可控的数值路径。
如果 rollout 用 FP16,training 用 BF16,或二者 kernel/accumulation 不同,同一 token logprob 可能不同。

9. Kernel 与 Attention 实现不一致

即使权重和 dtype 相同,不同 kernel 也可能产生不同结果: - FlashAttention 版本不同。

  • fused/unfused RMSNorm。

  • fused logits processor。

  • rotary embedding 实现不同。

  • tensor parallel 切分不同。

  • padding 和 position id 处理不同。

  • KV cache 精度不同。

这些差异通常很小,但在 RL 中可能放大,因为 policy gradient 依赖概率比:

ratio = exp(delta_logp)
delta_logp 的系统偏差会指数放大成 ratio 偏差。

10. 解码策略不一致

训练和推理的解码设置也可能不同: - temperature。

  • top-k / top-p。

  • repetition penalty。

  • max tokens。

  • stop words。

  • chat template。

  • system prompt。

  • special tokens。

在 RL rollout 中,如果生成时使用 top-p 截断,但训练时按完整词表概率计算 logprob,需要明确:

old_logp 是截断采样分布下的 logprob
还是原始模型分布下的 logprob
如果定义不清,importance ratio 的语义就会混乱。

11. Mask、Position ID 与模板不一致

LLM RL 中常见工程错误: - prompt token 被错误计入 response loss。

  • padding token 参与 logprob/KL。

  • left padding/right padding 导致 position id 不一致。

  • rollout 使用一种 chat template,training 使用另一种。

  • EOS/stop token 处理不一致。

  • response 截断后 reward 和 mask 没有对齐。

这些问题会直接造成:

同一个 token 序列在训练端不是同一个条件概率问题。
例如 position id 错位会让模型看到不同的位置编码,logprob 自然不同。

12. R3 与 Router 训推不一致

在 MoE 模型中,router 决定 token 被分配到哪些 experts。RL rollout 和 training 如果 router 行为不一致,会带来更复杂的训推不一致。 R3 相关工作关注训练与推理中 router 行为、专家选择或负载之间的不一致,并尝试让 MoE 在 RL 训练中更稳定。 MoE 训推不一致的典型来源: - rollout/training 的精度不同导致 router logits 改变。

  • top-k expert 选择边界敏感,小扰动改变 expert。

  • 训练时并行切分和推理时并行切分不同。

  • auxiliary load balancing loss 和 RL objective 之间张力。

一旦 expert 路由变了,同一 token 的前向路径就变了,logprob mismatch 可能比 dense model 更大。

13. GSPO 与训推不一致

GSPO 使用 sequence-level importance ratio,而不是逐 token ratio:

r_seq = exp(1/|y| * sum_t (logp_new_t - logp_old_t))
它能缓解部分 token-level ratio 高方差问题。 从训推不一致角度看,GSPO 的意义是: - 不让单个 token 的 logprob mismatch 过度支配更新。

  • 让整条 response 作为更一致的优化单元。

  • 对 MoE 或长序列场景更稳定。

但 GSPO 不能自动消除 logprob mismatch。底层 rollout/training logprob 仍需尽量对齐。

14. 训推不一致的度量

常见度量: Token-level logprob difference:

delta_t = logp_train_t - logp_rollout_t

img

统计: - mean absolute delta。

  • max delta。

  • percentile。

  • ratio distribution。

  • token-level KL。

  • sequence-level logprob delta。

  • top-1 token 一致率。

  • top-k distribution drift。

Sequence-level difference:

Delta_seq = sum_t (logp_train_t - logp_rollout_t)

img

长度归一化:

Delta_seq_avg = 1/T * Delta_seq

img

在 RL 中还应监控: - clip fraction。

  • approx KL。

  • reward 与 ratio 的相关性。

  • dropped/filtered samples 比例。

15. 训推不一致的影响

训推不一致可能导致: - PPO/GRPO ratio 偏离真实值。

  • clip fraction 异常升高。

  • KL 估计失真。

  • advantage 与 policy update 不匹配。

  • reward 上升但真实推理效果下降。

  • 训练不稳定或发散。

  • MoE expert 路由不稳定。

  • 线上推理质量与离线评估不一致。

一个常见现象:

训练指标看起来正常
但部署推理时回答质量、长度、格式或成功率明显不同
这通常提示训练环境和推理环境存在未对齐因素。

16. 缓解策略

常用缓解方法:

1. 统一 rollout 和 training 的模型权重、tokenizer、chat template。
2. 对齐 dtype、attention kernel、position id、padding 方式。
3. rollout 后在 training engine 中 recompute old_logp。
4. 使用 true on-policy 数据。
5. 使用 importance sampling 修正 off-policy/mismatch。
6. 监控 token/sequence logprob delta。
7. 对 ratio 做 clipping、sequence-level aggregation 或 soft gate。
8. 对 MoE router 做额外稳定化或一致性约束。
9. 部署前用真实 inference stack 做评估。
实践中最稳的方案通常是:让训练时用于 loss 的 old_logp 由同一个 training graph 重新计算,而不是完全信任 rollout engine 返回的 logprob。

17. 排错流程

排查训推不一致可以按如下顺序:

1. 固定同一模型 checkpoint。
2. 固定同一 prompt 和 response。
3. 分别用 rollout engine 和 training engine 计算 token logprobs。
4. 比较 token id、attention mask、position id、logits、logprobs。
5. 逐项关闭 sampling processor、量化、flash attention、并行切分差异。
6. 检查 chat template、EOS、padding、stop words。
7. 比较 dense model 和 MoE model 的差异。
8. 记录 mismatch 分布并设报警阈值。
不要只看最终 reward。训推不一致往往要从 token 级别开始定位。

18. 与前几天算法的关系

PPO:

ratio = exp(logp_new - logp_old)
对 old_logp 正确性高度敏感。
GRPO:

同样依赖 ratio,且多回答组内 advantage 会放大有效样本差异。
GSPO:

将 ratio 聚合到 sequence level,缓解 token-level 高方差。
DAPO:

通过 dynamic sampling、token-level loss、clip-higher 等提高训练有效性,但仍需 logprob 对齐。
SAPO:

用 soft adaptive gate 平滑处理 off-policy ratio,但不能替代底层一致性检查。

19. 面试中的表达框架

回答“什么是训推不一致”时,可以用四层框架:

定义:
&nbsp; 训练时优化的分布/计算路径和推理时真实使用的分布/计算路径不一致。

来源:
&nbsp; 数据分布、解码方式、数值精度、kernel、模板、mask、MoE router、rollout/training engine。

影响:
&nbsp; logprob/KL/ratio/advantage 失真,导致 RL 更新偏差和部署效果下降。

解决:
&nbsp; 对齐系统栈,recompute old_logp,监控 mismatch,使用 IS、GSPO/R3 等稳定化方法。

20. 核心总结

第十八天需要掌握的最小闭环:

训推不一致:
&nbsp; training objective / computation / data
&nbsp; != inference behavior / computation / data

LLM RL 关键点:
&nbsp; rollout logprob must match training logprob semantics

Mismatch sources:
&nbsp; dtype, kernel, quantization, sampling, template,
&nbsp; mask, position id, KV cache, MoE router

Impact:
&nbsp; ratio = exp(logp_new - logp_old) distorted
&nbsp; KL distorted
&nbsp; policy update biased

Mitigation:
&nbsp; align stack
&nbsp; recompute old_logp
&nbsp; true on-policy
&nbsp; importance sampling
&nbsp; sequence-level ratio
&nbsp; router/precision stability
&nbsp; real inference evaluation

21. 参考资料

  • BF16 比 FP16 更优?:https://arxiv.org/abs/2510.26788

  • R3 原论文:https://arxiv.org/abs/2510.11370

  • GSPO 原论文:https://arxiv.org/abs/2507.18071

  • 训推不一致总结:https://github.com/zhaochenyang20/Awesome-ML-SYS-Tutorial/blob/main/rlhf/slime/mismatch/blog-cn.md

  • vLLM 文档:https://docs.vllm.ai/

  • SGLang 文档:https://docs.sglang.ai/

            预览时标签不可点
    

    <div class="