执行摘要
- 一句话:修复轨迹合并逻辑,简化 middleware 损失掩码计算
- 推荐动作:建议精读
slime/agent/trajectory.py 中 merge_turns 的 common_prefix 对齐设计,该模式可推广至其他序列拼接场景。注意检查下游代码对 output_loss_mask 的依赖,并及时适配 output_log_probs。新增测试用例可作为后续修改的回归保障。
功能与动机
原始 middleware 中的 raw-splice 和 TITO 验证逻辑过于复杂且难以维护,prompt drift 处理不精确导致训练 target 污染。本次重构将序列对齐与 mask 计算职责转移到 trajectory 模块,通过最长公共前缀匹配精确处理 drift,并移除冗余的 output_loss_mask 字段,使数据流更清晰。
实现拆解
-
重构 slime/agent/trajectory.py:移除 TurnRecord.output_loss_mask 字段,新增 output_log_probs 和 TokenSegment.rollout_log_probs;新增 TurnSegment 数据类及 make_turn_segment 工厂函数;重写 merge_turns,将之前基于严格前后缀匹配的逻辑改为基于 _common_prefix_len 的最长公共前缀匹配,在 prompt drift 时仅保留匹配前缀并对应的 loss_mask 置 0,同时记录 output_spans 以处理跨 turn 的掩码。
-
简化 examples/coding_agent_rl/middleware.py:移除 _RAW_PH_PREFIX/SUFFIX、_common_prefix_len、verify_tito_for_turn、_record_turn 等函数,不再在 middleware 中维护 raw splice 和 per-turn loss mask;导入 TurnSegment、make_turn_segment、merge_turn_segments,将损失掩码决策全部交给 trajectory.merge_turns;新增 _build_tools_schema、_anthropic_tools_to_chat_tools 等工具函数以支持动态工具注入。
-
新增测试文件 tests/test_agent_trajectory.py:用 6 个 @pytest.mark.unit 用例覆盖匹配前缀保留、跳过中间 turn、连续 drift、输出被拆分、token 数变化、prompt base 改变等场景,验证 merge_turns 的 loss_mask 和 rollout_log_probs 输出。
-
更新 CI 注册:将 test_agent_trajectory.py 加入 .github/workflows/pr-test.yml 的 CPU 测试矩阵,并同步更新 Jinja2 模板。
-
辅助变更:scripts/run-qwen3.5-27B.sh 新增 rollout 配置参数(TP/DP/内存利用率);examples/coding_agent_rl/run_qwen36_35b_a3b_swe_8nodes.sh 新增 sglang mamba 调度策略;Skil 文档补充 CI 兼容性要点。
关键文件:
slime/agent/trajectory.py(模块 轨迹;类别 source;类型 core-logic;符号 _prompt_matches_current, TurnSegment, _output_mask, make_turn_segment): 核心文件:重写了轨迹合并逻辑,引入 common_prefix 精确对齐,移除 loss_mask 字段,新增 logprobs 支持。
examples/coding_agent_rl/middleware.py(模块 中间件;类别 source;类型 core-logic;符号 _make_segment, _build_tools_schema, _anthropic_tools_to_chat_tools, _common_prefix_len): 大幅精简:移除 raw-splice 和 per-turn 验证逻辑,将损失掩码责任转移至 trajectory,新增工具注入函数。
tests/test_agent_trajectory.py(模块 测试;类别 test;类型 test-coverage;符号 test_merge_turns_preserves_matched_prefix_on_prompt_drift, test_merge_turns_drops_middle_turn_when_next_prompt_skips_it, test_merge_turns_handles_consecutive_prompt_drifts, test_merge_turns_masks_whole_output_when_prompt_drift_splits_it): 新增单元测试,全面覆盖 merge_turns 的各种 drift 场景,确保重构后逻辑正确。
.github/workflows/pr-test.yml(模块 CI;类别 infra;类型 infrastructure): 将新测试文件注册到 CI CPU 测试矩阵,确保 CI 执行。
.github/workflows/pr-test.yml.j2(模块 CI;类别 infra;类型 infrastructure): 同步更新 Jinja2 模板以生成 CI 配置文件。
关键符号:merge_turns, _common_prefix_len, _output_log_probs, make_turn_segment, merge_turn_segments, _build_tools_schema, _anthropic_tools_to_chat_tools
关键源码片段
slime/agent/trajectory.py
核心文件:重写了轨迹合并逻辑,引入 common_prefix 精确对齐,移除 loss_mask 字段,新增 logprobs 支持。
# slime/agent/trajectory.py - 核心合并函数
def merge_turns(turns: list[TurnRecord], *, metadata: dict[str, Any] | None = None) -> TokenSegment | None:
if not turns:
return None
prompt_ids = list(turns[0].prompt_ids)
response_ids: list[int] = []
loss_mask: list[int] = []
rollout_log_probs: list[float] = []
output_spans: list[tuple[int, int]] = [] # 记录每个 turn 输出在 response_ids 中的 [start, end)
for i, turn in enumerate(turns):
if i > 0:
# 如果新的 prompt 不再包含之前的 base_prompt,说明上下文完全变了,直接重启 segment
if turn.prompt_ids[: len(prompt_ids)] != prompt_ids:
logger.warning("[trajectory] merge prompt base changed; starting segment from drifted prompt")
prompt_ids = list(turn.prompt_ids)
response_ids = []
loss_mask = []
rollout_log_probs = []
output_spans = []
else:
# 取得新 prompt 中超出之前 response 的部分
prompt_suffix = turn.prompt_ids[len(prompt_ids) :]
matched_len = _common_prefix_len(response_ids, prompt_suffix)
if matched_len < len(response_ids):
# 只匹配了部分历史输出,未匹配部分必须从响应中截掉,并将对应的 loss 置 0
logger.warning(
"[trajectory] merge prefix drift; truncating %d unstitched response tokens",
len(response_ids) - matched_len,
)
for start, end in output_spans:
if start < matched_len < end:
# 如果这个 turn 的输出被部分截断,截断部分的 loss_mask 和 logprob 必须清零
loss_mask[start:matched_len] = [0] * (matched_len - start)
rollout_log_probs[start:matched_len] = [0.0] * (matched_len - start)
response_ids = response_ids[:matched_len]
loss_mask = loss_mask[:matched_len]
rollout_log_probs = rollout_log_probs[:matched_len]
output_spans = [
(start, min(end, matched_len)) for start, end in output_spans if start < matched_len
]
# 新的 context tail(未匹配部分)加入响应,loss 为 0
context_tail = prompt_suffix[matched_len:]
response_ids.extend(context_tail)
loss_mask.extend([0] * len(context_tail))
rollout_log_probs.extend([0.0] * len(context_tail))
# 追加当前 turn 的生成输出
start_pos = len(response_ids)
response_ids.extend(turn.output_ids)
loss_mask.extend([1] * len(turn.output_ids)) # 默认全为 1,后续可能会被部分清零
rollout_log_probs.extend(_output_log_probs(turn))
output_spans.append((start_pos, len(response_ids)))
return TokenSegment(
prompt_ids=prompt_ids,
response_ids=response_ids,
loss_mask=loss_mask,
rollout_log_probs=rollout_log_probs,
metadata=dict(metadata or {}),
)
评论区精华
本 PR 无公开 review 评论。
风险与影响
- 风险:主要风险:① 移除
TurnRecord.output_loss_mask 字段会破坏现有直接引用该字段的代码(如自定义 rollout 后处理),需确认无外部依赖;② output_log_probs 要求上游引擎提供 logprobs,若未提供将默认为空列表,可能导致下游计算出错;③ 新的 merge_turns 逻辑对极长序列对齐性能需关注;④ 测试覆盖了常见 drift 场景,但未测试 finish_reason 异常、空输出等边界。
- 影响:影响 coding_agent_rl 示例的全部 rollout 流程,以及任何直接使用
slime.agent.trajectory 模块的组件。轨迹合并逻辑是 agent RL 训练数据的核心管道,正确性直接影响训练效果。简化 middleware 可降低后续维护成本,但向下游暴露了新字段 output_log_probs 和 rollout_log_probs。
- 风险标记:核心路径变更, 数据格式迁移(移除 output_loss_mask), 测试覆盖有限
关联脉络
- PR #1956 Add slime/agent/ and move sandbox impl inside: 建立了 slime/agent 模块,为 trajectory.py 提供了组织归属。
- PR #1960 Extract more util code from coding_agent_rl example: 将 trajectory 和 parsing 等工具从示例抽取到核心库,与本 PR 形成一致方向。
- PR #1923 [examples] add coding_agent_rl: agent-in-sandbox RL minimal demo: 首次引入 coding_agent_rl 示例,其中包含了此处重构前的 raw-splice 逻辑。
参与讨论