SWIFT(swift)GRPO 训练推理不一致(Training-Inference-Mismatch):重要性采样校正、诊断指标与 Off-Policy 序列掩码
SWIFTswiftGRPO 训练推理不一致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本文围绕 SWIFTswift仓库中 GRPO 算法引入 vLLM 加速采样后产生的训练-推理不一致Training-Inference Mismatch问题展开先说明该问题如何破坏 GRPO 的 on-policy 假设再完整讲解 SWIFT 提供的四类重要性采样IS校正模式、五组训练期诊断指标的实现原理与命令行参数用法并介绍源自 DeepSeek-V3.2 的 Off-Policy 序列掩码技术。读完后你将能够在swift rlhf的 GRPO 训练中以正确参数开启/仅监控该机制并通过rollout_correction/前缀指标判断当前训练是否受推理引擎偏差影响。背景GRPO 的 on-policy 假设与 vLLM 引入的分布偏差GRPOGroup 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.pyHF 训练路径与 megatron grpo_trainer.pyMegatron 训练路径参数定义见 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_{t1}^{|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_truncateToken 级截断token_maskToken 级掩码sequence_truncate序列级截断sequence_mask序列级掩码阈值由--rollout_importance_sampling_threshold设置默认值为 2.0源码注释中标记为论文中的常数 $C$见 args_mixin.py。源码实现数值安全与权重应用位置从源码结构看四种模式的统一实现是_apply_rollout_importance_samplinggrpo_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_samplingTrue此时若显式设置了rollout_importance_sampling_mode会直接抛出ValueErrorrollout_mixin.py。原因是较新版本的 vLLM 支持processed_logprobs能返回与训练侧对齐的 logprobsIS 校正才有可靠的 $\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 logprobsrollout_per_token_logps是 vLLM 采样时回传的 logprobsgrpo_trainer.py。使用 Liger 融合 loss 的路径同样支持传入vllm_is_ratiogrpo_trainer.pyMegatron 路径有对应的同名字段实现megatron grpo_trainer.py。训练期诊断指标量化不一致程度SWIFT 在日志中追加一组以rollout_correction/为前缀的指标写入self._metrics[mode][frollout_correction/{key}]见 grpo_trainer.py用于监控训练-推理不一致的严重程度。指标实现集中在_compute_rollout_offpolicy_metricsgrpo_trainer.py与_compute_is_correction_metricsgrpo_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_{t1}^{|y|} \log p(y_t) \right) $$相关指标training_ppl/training_log_ppl训练策略的 PPL 及其对数rollout_ppl/rollout_log_pplrollout 策略的 PPL 及其对数log_ppl_difflog PPL 差值正值表示训练策略给该序列分配了更低的概率对应更高的 PPLlog_ppl_abs_difflog PPL 差值的绝对值均值log_ppl_diff_max/log_ppl_diff_minlog PPL 差值的最大/最小值ppl_ratioPPL 比值 $\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_tokenToken 级 χ² 散度$\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. 有效样本量ESSESS 度量重要性采样之后真正有效的样本数量$$ \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_meanIS 权重均值理想值为 1.0clipped_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_{t1}^{|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_mask1的 token 上计算$\delta_i \tau$且$\hat{A}_i 0$其中$\pi_{\text{old}}$ 优先使用rollout_per_token_logpsrollout/行为策略回传的 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_maskgrpo_trainer.py。日志中会以offpolicy_sequence_mask: enable/disable记录开关状态。兼容性限制在启用 OPD-RLteacher 蒸馏即 GRPO 配置了teacher_model时off_policy_sequence_mask_delta不允许使用参数校验与训练器内部都会抛出ValueErrorargs_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/前缀的 KLkl、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.pyMegatron 路径有对应实现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),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →