Prhub

#1960 Extract more util code from coding_agent_rl example

原始 PR 作者 zhuzilin 合并时间 2026-05-27 15:13 文件变更 6 提交数 2 评论 0 代码增减 +461 / -391

执行摘要

抽取 coding_agent_rl 通用工具到核心 agent 模块

延续近期将示例代码迁移到 slime/agent 核心库的系列工作(PR #1956、#1958、#1957),进一步抽离通用工具函数和数据结构,避免重复实现,降低维护成本。

建议代码审查者重点关注 trajectory.py 中 merge_turns 的多轮对齐逻辑和 parsing.py 的工具调用解析路径,确保与原有行为一致。值得阅读的设计决策包括将配置从环境变量迁移到 args 的统一管理方式。

讨论亮点

无 review 评论或讨论。

实现拆解

  1. 创建 slime/agent/trajectory.py:定义 TurnRecordTokenSegment 不可变数据类,实现 merge_turns 将多轮对话线性化为单段训练样本、write_segment_to_sample 填充 Sample 字段、fan_out_sample_segments 将多段样本展平并均分奖励。
  2. 创建 slime/agent/parsing.py:定义 ParsedModelOutput 数据类,提供 parse_model_output 统一调用 SGLang 的 reasoning 和 function-call 解析器,并保留 XML 回退解析。
  3. 修改 examples/coding_agent_rl/middleware.py:移除内联的 _update_prompt_and_maskverify_tito_cross_turn,改为使用 slime.agent.trajectory 中的 TurnRecordmerge_turns;简化 _build_prompt_generate 接口,将 prompt/output 记录为 TurnRecord 列表,后续由 pop_session_split 线性化。
  4. 修改 examples/coding_agent_rl/generate.py:删除内联的 _write_segment_to_sample_fan_out_to_samples,改从 slime.agent.trajectory 导入 fan_out_sample_segments;移除 SWE_MAX_RESPONSE_TOKENSSWE_TOOL_PARSER 等环境变量,改为从 args 获取 rollout_max_context_lensglang_tool_call_parser 等。
  5. 更新启动脚本和文档:删除 SWE_MAX_RESPONSE_TOKENSSWE_MAX_SEGMENT_TOKENSSWE_SAVE_TRAJECTORY_TREE 等环境变量导出,在 README 中说明新增的命令行参数用法。
  6. 无对应测试文件改动,风险在于新模块缺少测试覆盖。
文件 模块 状态 重要度
slime/agent/trajectory.py Agent 核心 added 9.0
slime/agent/parsing.py Agent 核心 added 8.88
examples/coding_agent_rl/middleware.py Agent 编排 modified 8.74
examples/coding_agent_rl/generate.py Agent 编排 modified 8.14
examples/coding_agent_rl/run_qwen36_35b_a3b_swe_8nodes.sh 脚本配置 modified 3.57
examples/coding_agent_rl/README.md 文档 modified 2.39

关键符号

TurnRecord TokenSegment merge_turns write_segment_to_sample fan_out_sample_segments ParsedModelOutput parse_model_output parse_tool_uses parse_xml_tool_uses _render_token_ids _common_prefix_len verify_tito_cross_turn _update_prompt_and_mask _generate _record_turn

关键源码片段

slime/agent/trajectory.py core-logic

核心新增文件,定义 agent 轨迹转训练样本的数据结构与算法,包括 TurnRecord、TokenSegment、merge_turns、fan_out_sample_segments 等关键符号。

"""Token-level trajectory helpers for agent rollouts."""
from __future__ import annotations
import copy
import dataclasses
import logging
from typing import Any
from slime.utils.types import Samplelogger = logging.getLogger(__name__)@dataclasses.dataclass(frozen=True)
class TurnRecord:
    """一次模型生成回合的精确 token 快照。    ``prompt_ids`` 是发送给生成器的完整 token 化 prompt。
    ``output_ids`` 是原始生成的输出。
    ``output_loss_mask`` 默认全 1,可被每回合验证清零。
    """
    prompt_ids: list[int]
    output_ids: list[int]
    output_loss_mask: list[int]
    finish_reason: str@dataclasses.dataclass(frozen=True)
class TokenSegment:
    """从 agent 轨迹组装的一个训练片段。"""
    prompt_ids: list[int]
    response_ids: list[int]
    loss_mask: list[int]
    metadata: dict[str, Any] = dataclasses.field(default_factory=dict)def merge_turns(turns: list[TurnRecord], *, metadata: dict[str, Any] | None = None) -> TokenSegment | None:
    """将 TurnRecord 列表重放为单个线性 TokenSegment。    第一轮的 prompt 作为片段 prompt。后续轮的 prompt 必须以
    ``prompt + response_so_far`` 开头;其后缀是新非模型上下文(loss_mask=0),
    接着是模型输出及其每回合输出 mask。
    """
    if not turns:
        return None
    prompt_ids = list(turns[0].prompt_ids)
    response_ids: list[int] = []
    loss_mask: list[int] = []
    for i, turn in enumerate(turns):
        if i > 0:
            expected_len = len(prompt_ids) + len(response_ids)
            if turn.prompt_ids[:len(prompt_ids)] == prompt_ids and \
               turn.prompt_ids[len(prompt_ids):expected_len] == response_ids:
                # 正常追加:prompt 是已有 prompt+response,多出的部分是新的非模型上下文
                context_tail = turn.prompt_ids[expected_len:]
                response_ids.extend(context_tail)
                loss_mask.extend([0] * len(context_tail))
            elif turn.prompt_ids[:len(prompt_ids)] == prompt_ids:
                # prompt 前缀匹配但后续偏移,重新基线化
                logger.warning("[trajectory] merge prefix drift; rebaselining segment")
                response_ids = list(turn.prompt_ids[len(prompt_ids):])
                loss_mask = [0] * len(response_ids)
            else:
                # prompt 基础改变,从当前 prompt 重新开始
                logger.warning("[trajectory] merge prompt base changed; starting fresh")
                prompt_ids = list(turn.prompt_ids)
                response_ids = []
                loss_mask = []
        response_ids.extend(turn.output_ids)
        loss_mask.extend(turn.output_loss_mask if len(turn.output_loss_mask) == len(turn.output_ids)
                         else [0]*len(turn.output_ids))
    return TokenSegment(prompt_ids=prompt_ids, response_ids=response_ids,
                        loss_mask=loss_mask, metadata=dict(metadata or {}))

评论区精华

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

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

风险与影响

  1. 回归风险:middleware.py 和 generate.py 重构幅度大(+161/-260、+34/-103),轨迹记录和样本组装逻辑改为调用新模块,可能出现边界情况处理不一致。
  2. 缺少测试覆盖:trajectory.py 和 parsing.py 均为新增且无单元测试,merge_turns 的多轮线性化逻辑容易出错。
  3. 配置兼容性:generate.py 从环境变量改为 args 获取 parser 参数,若调用方未通过 args 传递正确值,可能导致解析失败。
  1. 用户:示例代码配置变简洁,需改用 --rollout-max-context-len 等统一参数。
  2. 系统:核心 agent 模块增加 trajectory 和 parsing 基础设施,为更多 agent 工作流提供复用基础。
  3. 团队:维护点从示例内联代码转移到核心模块,降低重复。
核心路径变更 缺少测试覆盖 配置接口变更

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论