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 记录了跨版本参数的情况,便于后续分析。