告别训练与推理不匹配:基于 vLLM 和 TorchTitan 的位运算一致性同策略强化学习
我们展示了一种开源的、具有位运算一致性的同策略(On-Policy)强化学习运行方案,该方案以 TorchTitan 作为训练引擎,以 vLLM 作为推理引擎。基于 vLLM 在批处理不变推理方面的最新工作,我们在开源指南中演示了如何对 Qwen3 1.7B 进行强化学习微调,并实现训练与推理数值的位匹配。

研究表明,强化学习会放大训练器和采样器之间微小的数值差异,导致非确定性和不稳定的训练行为(He 等人,Yao, Liu 等人 以及 Liu, Li 等人)。我们通过实验验证了数值对强化学习结果的影响:当采样器使用与训练器不同的内核时(batch_inv_OFF),在 100 个步骤中观察到奖励有所下降。而在启用位运算精确训练(batch_inv_ON,此时 kl_div 始终等于 0.0)后,模型不仅在更少的步数内完成了训练,还获得了更高的总奖励。

方法
训练和推理框架由于工作负载属性不同,往往使用差异巨大的内核。即使在同一个推理框架内,针对不同场景也会选择不同的内核:大批次尺寸(Batch Size)的内核会在批次维度上进行重度并行化,而小批次尺寸的内核则在单个实例内进行更多并行化,以便在 GPU 的并行核心上实现更好的利用率。所有这些差异会导致训练和推理框架之间的数值偏差,进而损害强化学习的效果。
在这项工作中,我们解决了两个不同框架之间的不变量问题:以 TorchTitan 作为训练框架,vLLM 作为推理框架。我们审计了前向传播过程中每一个内核的调用,确保它们在不同框架间实现位等价。我们利用了 vLLM 最新批处理不变性工作中的前向传播内核,并为这些操作编写了简单的反向传播通道。
vLLM 拥有许多经过高度优化的融合算子,例如 SiLU MLP 和 RMSNorm(带残差连接)。为了保持位等价性,我们引入了完全相同的前向传播算子。这些操作需要注册自定义的反向传播通道,而这可以通过 TorchTitan 编写所使用的原生 PyTorch 轻松实现。
针对强化学习演示,我们编写了一个通用的强化学习脚本,使用了 GSM8K 数据集和正确性奖励。我们利用 TorchTitan 的训练器工具,并编写了一个自定义生成器。我们的生成器 VLLMRolloutEngine 封装了诸如调用生成(generate)和更新权重等简单功能。我们将所有内容同步运行,在单台主机上交替执行训练器和生成器。这展示了精确的同策略执行过程,但在大规模运行中并不常见。
后续工作
我们将继续推动位运算一致性训练和推理的发展。要跟踪这项工作,请查看相关的 RFC:#28326 和 #27433。具体而言,我们将专注于以下方向:
统一的模型定义。 尽管我们已经展示了位等价的训练和推理结果,但目前仍有两份模型代码副本,一份用于训练,一份用于推理。这虽然便于我们进行初步集成,但对于长期维护而言非常脆弱:任何对模型代码的微小改动都可能破坏训练与推理之间的一致性,导致数值不匹配。为训练和推理框架提供一份共享的模型代码,将消除人为疏忽引入错误的可能,并使位匹配特性更易于维护。
编译支持。 目前,我们没有为 TorchTitan 模型使用 torch.compile,因此 vLLM 强制处于即时执行(Eager Mode)模式。移除此限制并不困难,但需要构建一个基于 torch.compile 的 TorchTitan 模型版本。vLLM 深度依赖 torch.compile 且能够借此保持批处理不变性,但要维持跨框架兼容性,需要对训练版本模型进行调整。这将在后续工作中进行探索!
强化学习性能。 我们目前的结果表明,位一致性强化学习的运行速度比非位一致性情况慢 2.4 倍。我们将继续通过优化批处理不变内核以及利用编译等技术,来提升 vLLM 的性能。
更广泛的模型支持。 我们计划将此位一致性强化学习框架扩展到 Qwen3 1.7B 以外,以支持更多开源模型。我们还将推广审计工具和反向传播实现,以覆盖更广泛的算子类型,使“位运算训练-推理一致性”成为一项可扩展且可复用的特性。
如果您感兴趣或希望参与贡献,请加入以下 Slack 频道:
作者:Bram Wasti, Wentao Ye, Teja Rao, Michael Goin, Paul Zhang, Tianyu Liu, Natalia Gimelshein, Woosuk Kwon, Kaichao You, Zhuohan Li