TRL v1.15发布:默认启用融合LM头,最长训练序列最高提升6.9倍
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等更新。