SWIFT(swift)GRPO 训练推理不一致(Training-Inference-Mismatch):重要性采样校正、诊断指标与 Off-Policy 序列掩码
【免费下载链接】swiftUse PEFT or Full-parameter to CPT/SFT/DPO/GRPO 600+ LLMs (Qwen3.6, DeepSeek-V4, GLM-5.1, InternLM3, Llama4, ...) and 300+ MLLMs (Qwen3-VL, Qwen3-Omni, InternVL3.5, Ovis2.5, GLM4.5v, Gemma4, Llava, Phi4, ...) (AAAI 2025).项目地址: https://gitcode.com/GitHub_Trending/swift1/swift
本文围绕 SWIFT(swift)仓库中 GRPO 算法引入 vLLM 加速采样后产生的"训练-推理不一致"(Training-Inference Mismatch)问题展开:先说明该问题如何破坏 GRPO 的 on-policy 假设,再完整讲解 SWIFT 提供的四类重要性采样(IS)校正模式、五组训练期诊断指标的实现原理与命令行参数用法,并介绍源自 DeepSeek-V3.2 的 Off-Policy 序列掩码技术。读完后你将能够在swift rlhf的 GRPO 训练中以正确参数开启/仅监控该机制,并通过rollout_correction/前缀指标判断当前训练是否受推理引擎偏差影响。
背景:GRPO 的 on-policy 假设与 vLLM 引入的分布偏差
GRPO(Group Relative Policy Optimization)的训练目标可以表示为:
$$ \mathcal{L}{\text{GRPO}} = - \mathbb{E}{y \sim \pi_\theta} \left[ \min \left( r_t(\theta) \hat{A}_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) \hat{A}_t \right) \right] $$
其中:
- $r_t(\theta) = \frac{\pi_\theta(y_t|x, y_{<t})}{\pi_{\theta_{\text{old}}}(y_t|x, y_{<t})}$ 是重要性采样比(importance sampling ratio);
- $\hat{A}_t$ 是基于奖励与组内 baseline 计算的优势函数(advantage);
- $\epsilon$ 是裁剪参数(SWIFT 中对应
--epsilon,默认 0.2,见 args_mixin.py)。
核心假设:样本 $y$ 必须采自策略 $\pi_\theta$。落到工程上即两点:
- 采样(rollout)模型与训练(policy)模型必须是同一个模型$\pi_\theta$;
- 两者输出的概率分布必须完全一致,即 $\pi_{\text{rollout}} = \pi_\theta$。
而 GRPO 的训练速度在很大程度上受采样过程(rollout)制约。为加速采样,训练框架会引入 vLLM 等高性能推理引擎,理想假设是通过权重同步使 vLLM 与训练模型保持一致,即 $\pi_{\text{vLLM}} \equiv \pi_\theta$。
但实践中,即使权重完全同步,由于算子实现(kernel 实现、数值精度、attention 后端等)差异,两个引擎给出的概率分布仍然存在偏差:
$$ \pi_{\text{vLLM}}(y|x) \neq \pi_\theta(y|x) $$
此时真实的训练目标变为:
$$ \mathcal{L} = - \mathbb{E}{y \sim \pi{\text{vLLM}}} \left[ \min \left( r_t(\theta) \hat{A}_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) \hat{A}_t \right) \right] $$
即样本来自 $\pi_{\text{vLLM}}$,而梯度却按 $\pi_\theta$ 计算。这违反了算法的 on-policy 假设,引入训练-推理不一致(training-inference mismatch),可能导致训练不稳定甚至性能退化(官方文档称之为 RL collapse 的一类诱因)。
SWIFT 针对该问题提供两条工程路线:
- 重要性采样校正(Importance Sampling Correction):对 loss 乘以 IS 权重,把期望从 rollout 分布"拉回"训练分布;
- Off-Policy 序列掩码(Off-Policy Sequence Masking,源自 DeepSeek-V3.2):对偏差过大且优势为负的整条序列直接弃用。
两者对应的参数与实现均位于 grpo_trainer.py(HF 训练路径)与 megatron grpo_trainer.py(Megatron 训练路径),参数定义见 args_mixin.py。
解决方案一:重要性采样(IS)校正
基本思想
重要性采样的基本公式是:当样本实际来自分布 $q$ 而非目标分布 $p$ 时,可引入权重修正期望计算:
$$ \mathbb{E}{x \sim p} [f(x)] = \mathbb{E}{x \sim q} \left[ \frac{p(x)}{q(x)} \cdot f(x) \right] $$
映射到 GRPO 场景,校正后的损失函数为:
$$ \mathcal{L}{\text{corrected}} = - \mathbb{E}{y \sim \pi_{\text{vLLM}}} \left[ w(x, y) \cdot \min \left( r_t(\theta) \hat{A}_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) \hat{A}_t \right) \right] $$
其中 $w(x, y)$ 是用于校正 vLLM 与训练模型之间分布偏差的 IS 权重。
校正粒度:Token 级与序列级
IS 权重可以在两种粒度上计算:
- Token 级(Token-Level):逐 token 计算 IS 比:
$$ w_{i,t}^{\text{token}} = \frac{\pi_\theta(y_{i,t}|x, y_{i,<t})}{\pi_{\text{vLLM}}(y_{i,t}|x, y_{i,<t})} $$
- 序列级(Sequence-Level):先计算序列级 IS 比,再广播到每个 token:
$$ w_i^{\text{seq}} = \left[ \frac{\pi_\theta(y_i|x)}{\pi_{\text{vLLM}}(y_i|x)} \right]^{\frac{1}{|y_i|}} = \exp\left( \frac{1}{|y_i|} \sum_{t=1}^{|y_i|} \log \frac{\pi_\theta(y_{i,t}|x, y_{i,<t})}{\pi_{\text{vLLM}}(y_{i,t}|x, y_{i,<t})} \right) $$
即序列级权重是 token 级比值的几何平均(对 log 比取 completion token 上的均值再取指数)。
稳定性控制:Truncate 与 Mask
过大的 IS 权重会引发梯度爆炸、 destabilize 训练,因此需要控制权重:
1. Truncate(截断):将权重截断到 $[0, \tau]$ 区间:
$$ w_{\text{truncate}} = \min(w, \tau) $$
保留所有样本,但限制其最大影响力。
2. Mask(掩码):权重超过阈值的 token/序列直接置零丢弃:
$$ w_{\text{mask}} = \begin{cases} w & \text{if } w \leq \tau \ 0 & \text{otherwise} \end{cases} $$
四种校正模式
组合"粒度 × 控制策略",得到四种校正模式,通过--rollout_importance_sampling_mode选择:
| 模式 | 说明 |
|---|---|
token_truncate | Token 级截断 |
token_mask | Token 级掩码 |
sequence_truncate | 序列级截断 |
sequence_mask | 序列级掩码 |
阈值由--rollout_importance_sampling_threshold设置,默认值为 2.0(源码注释中标记为论文中的常数 $C$,见 args_mixin.py)。
源码实现:数值安全与权重应用位置
从源码结构看,四种模式的统一实现是_apply_rollout_importance_sampling(grpo_trainer.py),有几个工程细节值得注意:
- log 比安全钳制:计算 $\exp(\text{log_ratio})$ 之前,先把 log 比 clamp 到 $[-20, 20]$(
SAFETY_BOUND = 20.0)。注释解释了原因:log 比为 20 时 $\exp(20) \approx 4.85$ 亿,这已经是极端值;该钳制同时防止 padding 位置(logprobs 常填 -1e10)造成数值溢出。 - 序列级比值:
_compute_sequence_level_ratios在 token 级比值上先取log,再按completion_mask求均值后取exp,与上文几何平均公式一致(grpo_trainer.py)。 - 权重与 loss 的乘积位置:IS 权重在 policy loss 计算完成之后、求和之前逐 token 相乘:
if rollout_is_weights is not None and self.rollout_importance_sampling_mode is not None: per_token_loss = per_token_loss * rollout_is_weights见 grpo_trainer.py。也就是说,IS 校正只作用于策略 loss 项,而 KL 惩罚(beta项)是在乘权重之前已并入per_token_loss的,两者一并被 IS 权重加权。
此外还有两个前提条件:
- vLLM 版本约束:SWIFT 通过
check_vllm_version_ge('0.10.2')判断版本;若 vLLM 低于 0.10.2,会自动置disable_rollout_importance_sampling=True,此时若显式设置了rollout_importance_sampling_mode会直接抛出ValueError(rollout_mixin.py)。原因是较新版本的 vLLM 支持processed_logprobs,能返回与训练侧对齐的 logprobs,IS 校正才有可靠的 $\pi_{\text{vLLM}}$ 估计。 - IS 比值的定义方向:
_get_rollout_is_correction中rollout_log_ratio = old_per_token_logps - rollout_per_token_logps,即 $\log(\pi_\theta/\pi_{\text{rollout}})$,其中old_per_token_logps是训练侧(当前策略或上一步策略)的 per-token logprobs,rollout_per_token_logps是 vLLM 采样时回传的 logprobs(grpo_trainer.py)。使用 Liger 融合 loss 的路径同样支持传入vllm_is_ratio(grpo_trainer.py),Megatron 路径有对应的同名字段实现(megatron grpo_trainer.py)。
训练期诊断指标:量化"不一致程度"
SWIFT 在日志中追加一组以rollout_correction/为前缀的指标(写入self._metrics[mode][f'rollout_correction/{key}'],见 grpo_trainer.py),用于监控训练-推理不一致的严重程度。指标实现集中在_compute_rollout_offpolicy_metrics(grpo_trainer.py)与_compute_is_correction_metrics(grpo_trainer.py)。
1. KL 散度
KL 散度度量 rollout 策略与训练策略的偏差,两个估计量都估计 $\text{KL}(\pi_{\text{vLLM}} | \pi_\theta)$:
直接估计量kl:
$$ \text{KL}(\pi_{\text{vLLM}} | \pi_\theta) = \mathbb{E}{\pi{\text{vLLM}}}\left[ \log \frac{\pi_{\text{vLLM}}}{\pi_\theta} \right] $$
K3 估计量k3_kl:
$$ \text{KL}(\pi_{\text{vLLM}} | \pi_\theta) \approx \mathbb{E}{\pi{\text{vLLM}}}\left[ \rho - \log \rho - 1 \right], \quad \rho = \frac{\pi_\theta}{\pi_{\text{vLLM}}} $$
K3 估计量在 KL 值较小时数值更稳定,且恒为非负(实现上对ρ − log ρ − 1再做了[-10, 10]的 clamp,见 grpo_trainer.py)。
2. 困惑度(PPL)
困惑度度量模型对一条序列的预测不确定性:
$$ \text{PPL} = \exp\left( -\frac{1}{|y|} \sum_{t=1}^{|y|} \log p(y_t) \right) $$
相关指标:
training_ppl/training_log_ppl:训练策略的 PPL 及其对数;rollout_ppl/rollout_log_ppl:rollout 策略的 PPL 及其对数;log_ppl_diff:log PPL 差值,正值表示训练策略给该序列分配了更低的概率(对应更高的 PPL);log_ppl_abs_diff:log PPL 差值的绝对值均值;log_ppl_diff_max/log_ppl_diff_min:log PPL 差值的最大/最小值;ppl_ratio:PPL 比值 $\frac{\text{PPL}{\text{training}}}{\text{PPL}{\text{rollout}}}$。
实现上ppl_ratio是在 log 空间用exp(log_ppl_diff)逐序列计算后再取均值,以避免数值不稳定(grpo_trainer.py)。
3. χ² 散度(Chi-squared Divergence)
χ² 散度度量 IS 权重的方差:
$$ \chi^2(\pi_\theta | \pi_{\text{vLLM}}) = \mathbb{E}{\pi{\text{vLLM}}}\left[ \rho^2 \right] - 1, \quad \rho = \frac{\pi_\theta}{\pi_{\text{vLLM}}} $$
chi2_token:Token 级 χ² 散度,$\mathbb{E}[\rho_t^2] - 1$;chi2_seq:序列级 χ² 散度(基于几何平均),$\mathbb{E}[\rho_{\text{geo}}^2] - 1$,其中 $\rho_{\text{geo}} = \exp(\frac{1}{T}\sum_t \log \rho_t)$。
χ² 散度越大,说明 IS 权重方差越大、训练越不稳定。chi2_seq采用几何平均而非连乘,使其量级与chi2_token可比。
4. 有效样本量(ESS)
ESS 度量重要性采样之后"真正有效"的样本数量:
$$ \text{ESS} = \frac{1}{\mathbb{E}\left[\left(\frac{w}{\mathbb{E}[w]}\right)^2\right]} $$
ESS 越接近 1,说明 IS 权重分布越均匀、样本利用率越高;权重完全相等(严格 on-policy)时 ESS = 1,权重差异悬殊(严重 off-policy)时 ESS 显著变小。实现上计算 ESS 前会先把权重 clamp 到 $[1/\tau, \tau]$ 以保证稳定性(grpo_trainer.py)。
5. IS 权重统计
is_weight_mean:IS 权重均值,理想值为 1.0;clipped_frac:被截断或掩码的样本比例。Token 级 truncate 统计 $\mathbb{E}[\mathbb{1}(\rho_t > \tau)]$;Token 级 mask 统计权重为 0 的 token 比例;序列级两种模式都统计序列级比值超阈值的序列比例(grpo_trainer.py)。
使用方式
仅记录诊断指标(不开启校正)
若只想监控训练-推理不一致程度、而不启用 IS 校正:
swift rlhf \ --rlhf_type grpo \ --model Qwen/Qwen2.5-7B-Instruct \ --log_rollout_offpolicy_metrics true \ ...该开关(log_rollout_offpolicy_metrics,默认False)会记录全部诊断指标(KL、PPL、χ²、ESS 等),但不修改 loss 函数(args_mixin.py)。
启用重要性采样校正
swift rlhf \ --rlhf_type grpo \ --model Qwen/Qwen2.5-7B-Instruct \ --rollout_importance_sampling_mode token_truncate \ --rollout_importance_sampling_threshold 2 \ ...--rollout_importance_sampling_mode:默认None(禁用),可选token_truncate/token_mask/sequence_truncate/sequence_mask;--rollout_importance_sampling_threshold:截断/掩码阈值,默认 2。
当设置了rollout_importance_sampling_mode时,诊断指标会自动记录,无需再单独设置log_rollout_offpolicy_metrics(触发逻辑见 grpo_trainer.py)。
适用前提:需 vLLM >= 0.10.2(低版本会抛错并提示);且训练侧需能拿到 rollout 引擎回传的rollout_per_token_logps——若 batch 内任意 rank 缺失该字段,指标与校正都会跳过。
解决方案二:Off-Policy 序列掩码(DeepSeek-V3.2)
除 IS 校正外,SWIFT 还提供Off-Policy 序列掩码,技术来自 DeepSeek-V3.2 论文。
原理
其核心思想是:当当前策略与旧策略(rollout/old policy)偏差过大时,直接从 loss 中丢弃(mask)该条序列。该策略专门针对优势为负的序列,因为策略偏移大时这类序列最容易引发训练不稳定。
对每条序列计算:
$$ \delta_i = \frac{1}{|y_i|} \sum_{t=1}^{|y_i|} \bigl( \log \pi_{\text{old}}(y_{i,t}|x, y_{i,<t}) - \log \pi_\theta(y_{i,t}|x, y_{i,<t}) \bigr) $$
当同时满足以下两个条件时,序列 $i$ 被掩码(均值均在completion_mask=1的 token 上计算):
- $\delta_i > \tau$
- 且$\hat{A}_i < 0$
其中:
- $\pi_{\text{old}}$ 优先使用
rollout_per_token_logps(rollout/行为策略回传的 logprobs),不可用时回退到old_per_token_logps(实现见 grpo_trainer.py); - $\tau$ 由
--off_policy_sequence_mask_delta设置,默认None表示禁用。
实现上,掩码通过扩展成 token 级后与completion_mask相与来完成,被掩序列在整个 loss 求和中不再贡献梯度(grpo_trainer.py;掩码判定逻辑见_compute_off_policy_sequence_mask,grpo_trainer.py)。日志中会以offpolicy_sequence_mask: enable/disable记录开关状态。
兼容性限制:在启用 OPD-RL(teacher 蒸馏,即 GRPO 配置了teacher_model)时,off_policy_sequence_mask_delta不允许使用,参数校验与训练器内部都会抛出ValueError(args_mixin.py、grpo_trainer.py)。
用法
swift rlhf \ --rlhf_type grpo \ --model Qwen/Qwen2.5-7B-Instruct \ --off_policy_sequence_mask_delta 0.05 \ ...IS 校正与序列掩码解决的是问题的两个侧面:IS 校正对所有样本做分布偏差加权,序列掩码则对偏差大且为负优势的样本直接弃用,二者可以独立开启,也可结合使用。
小结
- GRPO 的数学推导建立在"采样分布 = 训练策略分布"的 on-policy 假设上;vLLM 加速采样虽通过权重同步保持一致,但算子实现差异仍使 $\pi_{\text{vLLM}} \neq \pi_\theta$,从而引入训练-推理不一致。
- SWIFT 的应对手段分三层:
- 监控层:
--log_rollout_offpolicy_metrics true记录rollout_correction/前缀的 KL(kl、k3_kl)、PPL、χ²、ESS、is_weight_mean、clipped_frac指标; - 校正层:
--rollout_importance_sampling_mode(四种模式)+--rollout_importance_sampling_threshold(默认 2),在 loss 上乘以经过截断/掩码的 IS 权重; - 弃用层:
--off_policy_sequence_mask_delta对"策略偏移大 + 负优势"的序列整体掩码(DeepSeek-V3.2 方案)。
- 监控层:
- 核心实现位于 swift/rlhf_trainers/grpo_trainer.py,参数定义在 swift/rlhf_trainers/args_mixin.py,Megatron 路径有对应实现(swift/megatron/trainers/grpo_trainer.py)。
- 实际训练时建议先只开监控指标,观察
kl、chi2_token、ess是否处于健康区间,再决定是否需要开启 IS 校正或序列掩码,并注意 vLLM 版本(>= 0.10.2)这一硬性前提。
【免费下载链接】swiftUse PEFT or Full-parameter to CPT/SFT/DPO/GRPO 600+ LLMs (Qwen3.6, DeepSeek-V4, GLM-5.1, InternLM3, Llama4, ...) and 300+ MLLMs (Qwen3-VL, Qwen3-Omni, InternVL3.5, Ovis2.5, GLM4.5v, Gemma4, Llava, Phi4, ...) (AAAI 2025).项目地址: https://gitcode.com/GitHub_Trending/swift1/swift
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考