跳转至

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),确保策略梯度只作用于模型的决策部分。