[错误] 长度相等的 3D mRoPE 位置 ID 产生不一致的锯齿布局 → V1 训练器中的 split_with_sizes 发生崩溃
作者: iceflysnow创建于 2026年9月17日更新于 2026年9月17日
总结 在 V1 训练器中,每个样本的 3D(mRoPE)position_ids 的形状为 (num_components, seq_len),在管道中传递为具有不规则的嵌套张量。在批量中所有序列长度相同的情况下(包括 batch_size == 1 的简单情况),两个相互作用的问题会导致训练崩溃:
- 写入端。
torch.nested.as_nested_tensor(list_of_2d, layout=torch.jagged)当所有输入样本长度相等时,会将组件维度误认为是不规则维度:而不是标准的不规则@2 格式(lengths=[L_1..L_B],values=(C, ΣL_i),_ragged_idx=2),它会生成不规则@1 格式(lengths=[C]×B,values=(B·C, L),_ragged_idx=1)。 - 消费端。
maybe_fix_3d_position_ids()(verl/utils/tensordict_utils.py)在任何 3D 嵌套position_ids上无条件地设置_ragged_idx = 2,而不验证实际的布局。应用于上述不规则@1 张量,这会生成一个元数据不一致的张量,其下一个unbind()会引发错误:
RuntimeError: split_with_sizes expects split_sizes to sum exactly to 34135
(input tensor's size at dimension 1), but got split_sizes=[4, 4, 4, ...×128]我们在自己的 35B MoE 多回合 RL 运行中正好遇到了这种情况 — 运行了 81 个健康步骤,然后出现了崩溃。下面有详细信息。这个特性并不是特定于 torch-2.9 — 在 CPU 上使用 torch 2.13.0 也可以重现这个问题。
内容来源: verl-project/verl