跳转至

partial_single_turn_agent_loop.py — 实现了支持部分回滚的单轮 Agent 循环

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

文件概述

实现了支持部分回滚的单轮 Agent 循环。当单轮生成被中断时,已生成的 token 会被保存,恢复后从断点继续生成,最终将多段结果拼接为完整的响应。

关键代码讲解

核心逻辑

@register("partial_single_turn_agent")
class PartialSingleTurnAgentLoop(AgentLoopBase):
    async def run(self, sampling_params, **kwargs):
        output = kwargs.get("output", None)
        messages = list(kwargs["raw_prompt"])
        param_version = kwargs.get("param_version", 0)

        if not output:
            # 全新开始:应用聊天模板获取 prompt_ids
            prompt_ids = await self.apply_chat_template(messages, images=images, videos=videos)
        else:
            if output.extra_fields.get("is_cancel", False):
                # 恢复:将之前的 prompt + 已生成的 response 作为新的 prompt
                prompt_ids = output.prompt_ids + output.response_ids
                param_version_start = output.extra_fields.get("param_version_start", param_version)
            else:
                return output  # 已完成的样本直接返回

        # 调用支持取消的生成接口
        response_ids, response_logprobs, is_cancel = await self.server_manager.generate_for_partial(
            request_id=request_id,
            prompt_ids=prompt_ids, ...
        )

        if output:
            # 拼接之前的结果和新生成的结果
            prompt_ids = output.prompt_ids
            response_logprobs = output.response_logprobs + response_logprobs
            response_ids = output.response_ids + response_ids
            response_mask = [1] * len(response_ids)

        # 如果已达到最大长度,标记为完成
        if len(response_ids) >= self.response_length:
            is_cancel = False

        return AgentLoopOutput(
            prompt_ids=prompt_ids,
            response_ids=response_ids[:self.response_length],
            response_mask=response_mask[:self.response_length],
            response_logprobs=response_logprobs[:self.response_length],
            extra_fields={
                "is_cancel": is_cancel,
                "param_version_start": param_version_start,
                "param_version_end": param_version_end,
            }, ...
        )

恢复流程

第一次运行(使用参数 v1):
  prompt_ids = [1, 2, 3, 4]
  generate_for_partial() → response_ids = [5, 6, 7], is_cancel=True
  输出: prompt=[1,2,3,4], response=[5,6,7], is_cancel=True

参数同步(v1 → v2)...

第二次运行(使用参数 v2):
  prompt_ids = [1, 2, 3, 4, 5, 6, 7]  ← 拼接了上次的结果
  generate_for_partial() → response_ids = [8, 9, 10], is_cancel=False
  输出: prompt=[1,2,3,4], response=[5,6,7,8,9,10], is_cancel=False
                                         ↑ v1 生成  ↑ v2 生成

核心类/函数列表

名称 类型 说明
PartialSingleTurnAgentLoop 类 支持部分回滚的单轮 Agent

与其他模块的关系

  • 继承自 AgentLoopBase(agent_loop/agent_loop.py)
  • 注册为 "partial_single_turn_agent"
  • 使用 FullyAsyncLLMServerManager.generate_for_partial()

小结

PartialSingleTurnAgentLoop 通过简单的 prompt 拼接实现了单轮生成的断点续传。相比 AsyncPartialToolAgentLoop,它不需要保存复杂的状态机状态——只需将已生成的 token 追加到 prompt 中继续生成即可。param_version_start/end 记录了跨版本参数的情况,便于后续分析。