跳转至

agent_loop.py — 实现了全异步策略专用的 Agent 循环管理器和 Worker

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

文件概述

实现了全异步策略专用的 Agent 循环管理器和 Worker。核心特性是支持部分回滚(Partial Rollout)——在参数同步时中断正在进行的推理,同步完成后从中断点恢复。

关键代码讲解

1. FullyAsyncLLMServerManager - 支持断点续传的服务管理器

class FullyAsyncLLMServerManager(AsyncLLMServerManager):
    async def generate(self, request_id, *, prompt_ids, sampling_params, ...):
        """支持 abort 后自动恢复的 generate"""
        final_output = TokenOutput(token_ids=[], log_probs=[], num_preempted=0)

        while True:
            # 1. 生成 token
            output = await super().generate(
                request_id=request_id,
                prompt_ids=prompt_ids + final_output.token_ids,  # 加上已生成的部分
                sampling_params=sampling_params, ...
            )

            # 2. 合并输出
            final_output.token_ids.extend(output.token_ids)
            if output.log_probs is not None:
                final_output.log_probs.extend(output.log_probs)

            # 3. 更新 max_new_tokens
            if original_max_tokens is not None:
                sampling_params[limit_key] = original_max_tokens - len(final_output.token_ids)

            # 4. 检查是否因 abort 中断
            if output.stop_reason not in ("aborted", "abort") or not partial_rollout_resume:
                break  # 正常结束或不支持恢复

        return final_output

关键点:当 stop_reason 是 "aborted" 时,已生成的 token 会被保留,下次调用时作为 prompt 的一部分继续生成。

2. generate_for_partial - 支持取消的生成接口

async def generate_for_partial(self, request_id, *, prompt_ids, sampling_params, ...):
    """返回 (token_ids, log_probs, is_cancel)"""
    server = self._choose_server(request_id)
    output = await server.generate_for_partial.remote(
        request_id=request_id,
        prompt_ids=prompt_ids,
        sampling_params=sampling_params, ...
    )
    return output  # (token_ids, log_probs, is_cancel)

3. FullyAsyncAgentLoopWorker

@ray.remote
class FullyAsyncAgentLoopWorker(AgentLoopWorker):
    def __init__(self, ...):
        self.server_manager = FullyAsyncLLMServerManager(config, server_handles)
        self.cancellation_event = asyncio.Event()  # 共享取消事件

    async def generate_sequences_no_post(self, batch, partial_output_list):
        """生成序列,支持部分回滚"""
        tasks = []
        for i in range(len(batch)):
            kwargs = {k: v[i] for k, v in batch.non_tensor_batch.items()}
            kwargs["output"] = partial_output_list[i]  # 传入之前的部分结果
            tasks.append(asyncio.create_task(
                self._partial_run_agent_loop(sampling_params, trajectory, **kwargs)
            ))
        outputs = await asyncio.gather(*tasks)

        is_cancel = any(o.extra_fields.get("is_cancel", False) for o in outputs)
        if not is_cancel:
            output = self._postprocess(outputs)
            return output, False
        return outputs, True  # 返回部分结果,等待恢复

    async def cancel_agent_loops(self):
        """设置取消事件"""
        self.cancellation_event.set()

    async def resume_agent_loops(self):
        """清除取消事件"""
        self.cancellation_event.clear()

4. FullyAsyncAgentLoopManager

class FullyAsyncAgentLoopManager(AgentLoopManager):
    def __init__(self, config, ...):
        # 根据 rollout 后端选择 Replica 类
        if rollout_name == "sglang":
            self.rollout_replica_class = FullyAsyncSGLangReplica
        elif rollout_name == "vllm":
            self.rollout_replica_class = FullyAsyncvLLMReplica

    async def generate_single_sample_async(self, sample, partial_output_list):
        """异步处理单个样本"""
        worker = self._select_best_worker()  # 轮询选择 worker
        output_future = worker.generate_sequences_no_post.remote(sample, partial_output_list)
        return await asyncio.wrap_future(output_future.future())

    async def cancel(self):
        """取消所有 worker 和 rollout replica 的进行中任务"""
        worker_cancel_tasks = [w.cancel_agent_loops.remote() for w in self.agent_loop_workers]
        rollout_cancel_tasks = [r.cancel() for r in self.rollout_replicas]
        await asyncio.gather(*rollout_cancel_tasks, *worker_cancel_tasks)

    async def resume(self):
        """恢复所有 worker 和 replica"""
        rollout_resume_tasks = [r.resume() for r in self.rollout_replicas]
        worker_resume_tasks = [w.resume_agent_loops.remote() for w in self.agent_loop_workers]
        await asyncio.gather(*rollout_resume_tasks, *worker_resume_tasks)

部分回滚流程图

正常生成中 ──────────────────────────────────────────────────
    │
    │ ParameterSynchronizer.sync_weights() 触发
    ▼
cancel() ──► 所有 Worker + Replica 收到取消信号
    │
    │ Agent Loop 检测到 cancellation_event
    │ 保存当前状态(AgentData + AgentState)到 AgentLoopOutput
    │ 返回 is_cancel=True
    ▼
pause() ──► 等待所有活跃任务完成
    │
    │ 权重同步...
    ▼
resume() ──► 清除取消事件
    │
    │ cancel_queue 中的任务重新提交
    │ Agent Loop 从保存的状态恢复
    ▼
继续生成 ──────────────────────────────────────────────────

核心类/函数列表

名称 类型 说明
FullyAsyncLLMServerManager 类 支持 abort 恢复的服务管理器
FullyAsyncAgentLoopWorker Ray Remote 类 支持取消/恢复的 Worker
FullyAsyncAgentLoopManager 类 全异步 Agent 循环管理器

与其他模块的关系

  • 继承自 agent_loop/agent_loop.py 中的 AsyncLLMServerManager、AgentLoopWorker、AgentLoopManager
  • 被 FullyAsyncRollouter 使用
  • 与 partial_single_turn_agent_loop.py 和 partial_tool_agent_loop.py 配合工作

小结

该文件是全异步训练的推理调度核心。通过 cancellation_event + cancel_queue 机制,实现了推理任务的优雅中断和恢复。FullyAsyncLLMServerManager 的 abort 恢复机制使得中断对 Agent Loop 透明,已生成的 token 不会丢失。