tool_agent_loop.py — 实现了支持工具调用的多轮 Agent 循环¶
文件路径: verl/experimental/agent_loop/tool_agent_loop.py
文件概述¶
实现了支持工具调用的多轮 Agent 循环。这是 verl 中最核心的 Agent 循环实现之一,支持 LLM 自主调用外部工具(如代码执行器、搜索引擎等),并基于工具返回结果继续对话。采用状态机(State Machine)模式管理多轮对话流程。
关键代码讲解¶
1. 状态枚举 AgentState¶
class AgentState(Enum):
PENDING = "pending" # 等待处理
GENERATING = "generating" # LLM 正在生成
PROCESSING_TOOLS = "processing_tools" # 正在执行工具调用
TERMINATED = "terminated" # 已结束
INTERACTING = "interacting" # 用户交互中
2. AgentData - 封装所有状态变量¶
class AgentData:
def __init__(self, messages, image_data, video_data, metrics, request_id, tools_kwargs, ...):
self.messages = messages # 对话历史
self.image_data = image_data # 图片数据
self.prompt_ids: list[int] = [] # 持续增长的 token 序列
self.response_ids: list[int] = [] # 最新一轮的响应 token
self.response_mask: list[int] = [] # 累积的响应掩码
self.response_logprobs: list[float] = []
self.turn_scores: list[float] = [] # 每轮交互得分
self.tool_rewards: list[float] = [] # 工具奖励
self.user_turns = 0
self.assistant_turns = 0
self.tool_calls: list[FunctionCall] = [] # 当前轮的工具调用
3. ToolAgentLoop - 核心类¶
@register("tool_agent")
class ToolAgentLoop(AgentLoopBase):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.max_user_turns = self.rollout_config.multi_turn.max_user_turns
self.max_assistant_turns = self.rollout_config.multi_turn.max_assistant_turns
self.max_parallel_calls = self.rollout_config.multi_turn.max_parallel_calls
self.max_tool_response_length = self.rollout_config.multi_turn.max_tool_response_length
# 加载工具配置
tool_list = initialize_tools_from_config(tool_config_path)
self.tools = {tool.name: tool for tool in tool_list}
self.tool_schemas = [tool.tool_schema.model_dump(...) for tool in tool_list]
# 工具调用解析器
self.tool_parser = ToolParser.get_tool_parser(format_name, self.tokenizer)
4. 状态机主循环 - run 方法¶
async def run(self, sampling_params, **kwargs) -> AgentLoopOutput:
# 初始化 AgentData
agent_data = AgentData(messages=messages, ...)
# 状态机循环
state = AgentState.PENDING
while state != AgentState.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(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 AgentLoopOutput(...)
5. 各状态处理器¶
PENDING -> GENERATING:应用聊天模板,准备 prompt
async def _handle_pending_state(self, agent_data, sampling_params):
prompt_ids = await self.apply_chat_template(
agent_data.messages, tools=self.tool_schemas, images=..., videos=...
)
agent_data.prompt_ids = prompt_ids
return AgentState.GENERATING
GENERATING -> PROCESSING_TOOLS / INTERACTING / TERMINATED:调用 LLM 生成,检查是否包含工具调用
async def _handle_generating_state(self, agent_data, sampling_params):
output = await self.server_manager.generate(
request_id=agent_data.request_id,
prompt_ids=agent_data.prompt_ids,
sampling_params=sampling_params,
)
agent_data.prompt_ids += output.token_ids # 累加到 prompt
agent_data.response_mask += [1] * len(output.token_ids) # 标记为 LLM 生成
# 检查终止条件
if len(agent_data.response_mask) >= self.response_length:
return AgentState.TERMINATED
# 提取工具调用
_, agent_data.tool_calls = await self.tool_parser.extract_tool_calls(output.token_ids)
if agent_data.tool_calls:
return AgentState.PROCESSING_TOOLS
else:
return AgentState.TERMINATED
PROCESSING_TOOLS -> GENERATING:并行执行工具调用,将结果编码回 token
async def _handle_processing_tools_state(self, agent_data):
# 并行执行多个工具调用
tasks = [self._call_tool(tc, ...) for tc in agent_data.tool_calls[:self.max_parallel_calls]]
responses = await asyncio.gather(*tasks)
# 将工具响应编码为 token
response_ids = await self.apply_chat_template(add_messages, ...)
agent_data.prompt_ids += response_ids
agent_data.response_mask += [0] * len(response_ids) # 标记为非 LLM 生成
agent_data.user_turns += 1
return AgentState.GENERATING # 继续生成
6. 状态机流程图¶
┌──────────┐
│ PENDING │
└────┬─────┘
│ apply_chat_template
▼
┌────────────────┐
┌─────│ GENERATING │─────┐
│ └────────────────┘ │
│ 有工具调用 无工具调用/超限 │
▼ ▼
┌──────────────────┐ ┌────────────┐
│ PROCESSING_TOOLS │ │ TERMINATED │
└────────┬─────────┘ └────────────┘
│ 执行工具,编码结果 ▲
│ │
└───────────────────────────┘
再次生成
7. 工具调用执行¶
async def _call_tool(self, tool_call, tools_kwargs, agent_data):
tool_name = tool_call.name
tool_args = json.loads(tool_call.arguments)
tool = self.tools[tool_name]
# 创建工具实例 -> 执行 -> 释放
instance_id, _ = await tool.create(create_kwargs=kwargs.get("create_kwargs", {}))
response, reward, res = await tool.execute(instance_id, tool_args, agent_data=agent_data)
await tool.release(instance_id)
# 截断过长的工具响应
if len(response.text) > self.max_tool_response_length:
response.text = response.text[:self.max_tool_response_length] + "...(truncated)"
return ToolResponse(text=response.text), reward, res
核心类/函数列表¶
| 名称 | 类型 | 说明 |
|---|---|---|
AgentState |
枚举 | 状态机的 5 种状态 |
AgentData |
类 | 封装 Agent 循环的全部状态 |
ToolAgentLoop |
类 | 工具调用 Agent 循环(核心) |
_handle_pending_state() |
方法 | 处理 PENDING 状态 |
_handle_generating_state() |
方法 | 处理 GENERATING 状态 |
_handle_processing_tools_state() |
方法 | 处理工具调用状态 |
_handle_interacting_state() |
方法 | 处理用户交互状态 |
_call_tool() |
方法 | 执行单个工具调用 |
与其他模块的关系¶
- 继承自
AgentLoopBase(agent_loop.py) - 使用
ToolParser(tool_parser.py)解析工具调用 - 使用
build_gpt_oss_tool_response_text(utils.py)格式化特定模型的工具响应 - 被
AsyncPartialToolAgentLoop(fully_async_policy)继承并扩展,支持部分回滚
小结¶
ToolAgentLoop 通过状态机模式优雅地实现了 LLM 的多轮工具调用循环。每一轮:LLM 生成响应 -> 解析工具调用 -> 并行执行工具 -> 将结果编码为 token 继续生成。response_mask 精确区分了 LLM 生成的 token(mask=1)和工具响应 token(mask=0),确保策略梯度只作用于模型的决策部分。