Prhub

#2110 Extract append_response_tokens to Sample

原始 PR 作者 zhuzilin 合并时间 2026-06-21 10:12 文件变更 2 提交数 1 评论 0 代码增减 +98 / -25

执行摘要

增强 append_response_tokens 的校验与元数据一致性

原有实现允许 trainable=True 时不传 log_probs 静默忽略,以及非训练 token 传入 log_probs 被忽略,这些情况可能导致训练数据不一致且难以调试。本 PR 通过主动抛异常的方式强制约束,并统一元数据长度校验,保证训练流程的可靠性。

值得精读,特别是其防御式编程风格和统一元数据校验的设计思路。对比 #2108 的抽取,本 PR 进一步完善了接口契约。对于维护训练管道的工程师,建议检查所有 append_response_tokens 调用点是否满足新校验要求。

讨论亮点

该 PR 无 review 评论。

实现拆解

  1. types.py 中新增 _numel 辅助函数,统一计算张量元素数。

  2. Sample 类中新增 multimodal_train_input_id 字段,支持多模态训练输入标识。

  3. 重构 append_response_tokens

    • 新增参数校验:trainable=True 必须提供 log_probs,否则抛异常;trainable=False 禁止传入 log_probs,否则抛异常。
    • 调整 loss_mask 初始化:当 loss_maskNone 时,根据已有 response_length 初始化为全 1 列表。
    • 统一 rollout_log_probs 初始化:当前面有 response 但无 rollout_log_probs 且后续追加 trainable token 时抛异常。
    • 将 top-p 填充逻辑统一交由 _apply_meta_info 处理,不再单独调用 _pad_rollout_top_p_offsets
    • 结尾调用 _validate_response_metadata_lengths 进行一致性校验。
  4. 在测试文件中为现有用例补充 rollout_log_probsloss_mask 的初始值以匹配新行为,新增两个测试验证新增校验。

文件 模块 状态 重要度
slime/utils/types.py 数据模型 modified 7.88
tests/test_rollout_metrics.py 测试 modified 5.56

关键符号

_numel _validate_response_metadata_lengths append_response_tokens

关键源码片段

tests/test_rollout_metrics.py test-coverage

配合核心变更调整测试数据并新增边界测试

@pytest.mark.unit
def test_append_response_tokens_requires_trainable_log_probs():
    sample = Sample()
    # 当 trainable=True 但未提供 log_probs 时,应抛出 ValueError
    with pytest.raises(ValueError, match="trainable response tokens require rollout log probabilities"):
        sample.append_response_tokens(tokens=[10], trainable=True)
​
​
@pytest.mark.unit
def test_append_response_tokens_rejects_non_trainable_log_probs():
    sample = Sample()
    # 当 trainable=False 却传了 log_probs 时,应抛出 ValueError
    with pytest.raises(ValueError, match="non-trainable response tokens should not pass rollout log probabilities"):
        sample.append_response_tokens(tokens=[10], log_probs=[-0.1], trainable=False)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

主要风险在于新增的参数校验可能暴露调用侧未遵守约定而导致的运行时中断,但这是预期行为,有助于提前发现错误。_validate_response_metadata_lengths 的实现细节需要确认是否覆盖所有边界场景,若校验逻辑有漏洞可能导致 false positive 或 false negative。此外,loss_mask 初始化策略变更(从空列表改为全 1 列表)可能影响依赖此属性的下游逻辑,需确认所有调用点都已适配。由于 PR 无 review 讨论,潜在风险还需通过后续测试覆盖。

直接影响所有通过 Sample.append_response_tokens 追加 token 的代码路径,包括 rollout 引擎输出处理和 agent 轨迹生成。对于遵循正确用法的调用者无感知;对于未提供 log_probstrainable=True 调用会立即抛出 ValueError,需要在调用侧补传。影响范围中等,但因涉及训练数据管道,错误用法被暴露后需及时修复。

核心路径变更 新增参数校验可能暴露调用侧问题

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论