跳转至

partial_tool_agent_loop.py — 实现了支持部分回滚的工具调用 Agent 循环

文件路径: verl/experimental/fully_async_policy/agent_loop/partial_tool_agent_loop.py

文件概述

实现了支持部分回滚的工具调用 Agent 循环。在全异步训练中,当参数需要更新时,正在进行的多轮工具调用可能被中断。该类扩展了 ToolAgentLoop,增加了状态保存/恢复和中断检测能力。

关键代码讲解

1. 类定义

@register("async_partial_tool_agent")
class AsyncPartialToolAgentLoop(ToolAgentLoop):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.enable_partial_rollout = self.config.async_training.get("partial_rollout", False)

2. run 方法 - 支持中断/恢复

async def run(self, sampling_params, *, cancellation_event=None, **kwargs):
    # 1. 检查是否是恢复的任务
    output = kwargs.get("output", None)
    if output and output.extra_fields.get("is_cancel", False):
        # 从上次中断的状态恢复
        agent_data, state = self._restore_from_output(output)
    else:
        # 全新开始
        agent_data = await self._init_agent_data(kwargs, param_version)
        state = AgentState.PENDING

    # 2. 运行状态机(带中断检测)
    state = await self._run_state_machine(agent_data, state, sampling_params, cancellation_event)

    # 3. 构建输出
    if state == AgentState.TERMINATED:
        return self._build_completed_output(agent_data, param_version)
    else:
        return self._build_cancelled_output(agent_data, state)  # 保存状态供恢复

3. 状态机 - 带中断检测

async def _run_state_machine(self, agent_data, state, sampling_params, cancellation_event=None):
    while state != AgentState.TERMINATED:
        # 每次状态转换前检查是否被取消
        if cancellation_event and cancellation_event.is_set():
            return state  # 返回当前状态(不是 TERMINATED)

        if state == AgentState.PENDING:
            state = await self._handle_pending_state(agent_data, sampling_params)
        elif state == AgentState.GENERATING:
            state = await self._handle_generating_state_partial(agent_data, sampling_params)
        elif state == AgentState.PROCESSING_TOOLS:
            state = await self._handle_processing_tools_state(agent_data)
        elif state == AgentState.INTERACTING:
            state = await self._handle_interacting_state(agent_data)
    return AgentState.TERMINATED

4. 部分生成状态处理

async def _handle_generating_state_partial(self, agent_data, sampling_params):
    if self.enable_partial_rollout:
        # 使用支持取消的 generate_for_partial 接口
        response_ids, log_probs, is_cancel = await self.server_manager.generate_for_partial(
            request_id=agent_data.request_id,
            prompt_ids=agent_data.prompt_ids, ...
        )

        if is_cancel:
            # 保存已生成的部分
            agent_data.prompt_ids += response_ids
            agent_data.response_mask += [1] * len(response_ids)
            if len(agent_data.response_mask) >= self.response_length:
                return AgentState.TERMINATED  # 已达到长度限制
            return AgentState.GENERATING  # 返回 GENERATING 状态等待恢复
    else:
        # 使用标准 generate 接口
        output = await self.server_manager.generate(...)

5. 状态保存与恢复

保存取消时的状态:

def _build_cancelled_output(self, agent_data, state):
    return AgentLoopOutput(
        prompt_ids=[], response_ids=[], response_mask=[],
        extra_fields={
            "is_cancel": True,
            "agent_data": agent_data,  # 完整的 Agent 状态
            "agent_state": state,       # 当前状态机状态
        },
    )

从保存的状态恢复:

def _restore_from_output(self, output):
    agent_data = output.extra_fields.get("agent_data")
    agent_state = output.extra_fields.get("agent_state")
    return agent_data, agent_state

6. 参数版本跟踪

async def _init_agent_data(self, kwargs, param_version):
    agent_data = AgentData(...)
    agent_data.extra_fields["param_version_start"] = param_version
    agent_data.extra_fields["param_version_end"] = param_version
    return agent_data

def _build_completed_output(self, agent_data, param_version):
    output.extra_fields.update({
        "param_version_start": agent_data.extra_fields["param_version_start"],
        "param_version_end": param_version,  # 结束时的版本可能不同
    })

如果一个 rollout 跨越了参数更新,param_version_start != param_version_end。

核心类/函数列表

名称 类型 说明
AsyncPartialToolAgentLoop 类 支持部分回滚的工具调用 Agent
_run_state_machine() 方法 带中断检测的状态机
_handle_generating_state_partial() 方法 支持取消的生成状态处理
_build_cancelled_output() 方法 构建取消时的输出(含状态快照)
_restore_from_output() 方法 从快照恢复状态

与其他模块的关系

  • 继承自 ToolAgentLoop(agent_loop/tool_agent_loop.py)
  • 使用 FullyAsyncLLMServerManager.generate_for_partial()
  • 被 FullyAsyncAgentLoopWorker 通过 @register 机制实例化

小结

AsyncPartialToolAgentLoop 在 ToolAgentLoop 的状态机基础上增加了中断/恢复能力。核心思想是:中断时将 AgentData(对话历史、已生成 token、状态变量等)和当前 AgentState 完整保存到输出中;恢复时从保存的状态继续执行状态机。这使得参数同步对多轮工具调用是"透明"的——Agent 不知道中间发生了参数更新。