TRL v1.15发布,SFT、DPO、KTO、GRPO、RLOO和蒸馏默认启用融合LM头:通过Triton内核直接计算token级量,避免实例化巨大的[batch, seq, vocab] logits张量。以Gemma 3 1B(262k词表)为例,相同GPU下各算法最大序列长度最高提升6.9倍,8k上下文峰值显存降低52-82%,训练速度最高快约11%。发布方Lysandre表示,此优化无需额外启用,已是v1.15默认行为;该版本还包含SFT选择性激活检查点、视觉数据集assistant-only loss等更新。
ClementDelangue共 1 条来源 · 1 条符合当前筛选
Lysandre介绍TRL v1.15版本,SFT、DPO、KTO、GRPO、RLOO和蒸馏训练默认改用融合LM头,通过Triton内核直接计算所需的token级量,避免物化巨大的[batch, seq, vocab] logits张量。在Gemma 3 1B(262k词表)实测中,最大序列长度最高从28k提升至114k,8k上下文下峰值显存降低52-82%,训练速度最高提升约11%,无需额外配置即默认生效。
abidlabs共 1 条来源 · 1 条符合当前筛选
Quentin Gallouédec发布TRL v1.15,称其为迄今最大的优化。新版本默认启用融合LM head,峰值显存最高降低82%,序列长度提升至7倍;DPO从10k增至59k tokens,GRPO从29k增至115k tokens。
adithya_s_k共 1 条来源 · 1 条符合当前筛选