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 不知道中间发生了参数更新。