ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

AReaL 自定义 RolloutWorkflow 完全指南:从零实现并接入训练流水线

AReaL 自定义 RolloutWorkflow 完全指南:从零实现并接入训练流水线 AReaL 自定义 RolloutWorkflow 完全指南从零实现并接入训练流水线【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL本文基于 AReaL 仓库内的add-workflow技能文档.agents/skills/add-workflow/SKILL.md系统讲解如何在 AReaL 中新增一个RolloutWorkflow实现从理解 workflow 在整体架构中的位置、编写arun_episode核心逻辑、注册导出到在训练脚本中按字符串路径接入与编写测试并逐条对照源码说明各接口的真实契约。读完本文你将能够独立完成一个可运行、可测试的自定义 rollout 工作流并将其无缝接入 GRPO/PPO 训练。在上图所示的 AReaL 架构中Generation 区域里的 Rollout Worker 正是 workflow 的执行载体训练控制器按批次分发 promptworkflow 负责生成 打分 组织轨迹结果经 Replay Buffer 进入训练侧。因此RolloutWorkflow是 AReaL 中策略产生行为的抽象边界——换一个任务形态单轮数学、多轮对话、工具调用、视觉推理通常只需要新增一个 workflow 实现而不必改动训练引擎。一、RolloutWorkflow 接口契约源码里到底要求什么1.1 基类定义与返回值语义所有 rollout 工作流继承自 areal/api/workflow_api.py 中的RolloutWorkflowclass RolloutWorkflow(ABC): abstractmethod async def arun_episode( self, engine: InferenceEngine, data: dict[str, Any] ) - dict[str, Any] | None | dict[str, InteractionWithTokenLogpReward]: Run a single episode of the workflow. raise NotImplementedError()基类只约束一个抽象方法arun_episode但 该方法的 docstring 定义了几条关键契约写自定义 workflow 前必须理解arun_episode必须是async def且不可阻塞。整个 rollout 阶段以并发协程方式运行任何同步 I/O如open()读文件都会拖垮吞吐——这是技能文档Key Requirements第一条的底层原因。返回None表示拒绝该轨迹这条轨迹不会进入训练可用于过滤低质量样本框架会相应地重新组织 batch。返回张量轨迹字典dict[str, torch.Tensor]是最常见的形态。基类 docstring 明确建议若 workflow 能判断最终响应是否因达到长度上限而停止应附带一个is_truncated布尔张量每条轨迹一个值。RLVRWorkflow 的实现 正是这样做的is_truncated: torch.tensor(resp.stop_reason length, dtypetorch.bool)。PPO 会用该元数据做奖励屏蔽、value bootstrapping 与截断指标统计。每次arun_episode调用对应一条逻辑 rollout即使上下文压缩导出多行 tensor奖励统计也只按 rollout 计数advantage 统计则保持 token 粒度。对于行级奖励不同且启用奖励归一化的场景需要提供一个有限的标量rollout_reward或在其导出的 interactions 上给出且所有值一致作为组/批奖励统计的参考值。1.2 引擎入参InferenceEngine与ModelRequest/ModelResponsearun_episode收到的engine是 areal/api/engine_api.py 中InferenceEngine抽象类的实例workflow 只需使用其异步生成接口async def agenerate(self, req: ModelRequest) - ModelResponse: ...请求与响应结构体定义在 areal/api/io_struct.pyModelRequestL28-L45rid请求 ID默认uuid4生成、input_idstoken 列表、gconfigGenerationHyperparameters、metadata以及 VLM 场景的image_data/processor字段。ModelResponseL63-L84input_tokens、output_tokens、output_logprobs、output_versions以及stop_reason取值length | stop | tool_calls | abortis_truncated判断即依赖它。生成参数由 areal/api/cli_args.py 的GenerationHyperparametersdataclass 承载常用字段及默认值包括n_samples1、max_new_tokens16384、max_tokens32768、temperature1.0、top_p1.0、top_k1e8、greedyFalse、stop、seed等。workflow 构造时约定调用gconfig.new_with_stop_and_pad_token_ids(tokenizer)将 EOS/pad token ID 合入停止条件再通过self.gconfig.new(n_samples1)派生出单样本请求级配置——这一点在 RLVRWorkflow 的__init__与技能文档模板中完全一致。二、按四步流程新增一个 Workflow技能文档核心步骤逐条落地Step 1创建 workflow 文件areal/workflow/name.py技能文档给出的最小模板可直接作为起点复制修改import uuid from typing import Any, Callable import torch from areal.api.cli_args import GenerationHyperparameters from areal.api.engine_api import InferenceEngine from areal.api.io_struct import ModelRequest, ModelResponse from areal.api.reward_api import AsyncRewardWrapper from areal.api.workflow_api import RolloutWorkflow from areal.utils import logging logger logging.getLogger(MyWorkflow) class MyWorkflow(RolloutWorkflow): Description of your workflow. def __init__( self, gconfig: GenerationHyperparameters, tokenizer, reward_fn: Callable, ): self.gconfig gconfig.new_with_stop_and_pad_token_ids(tokenizer) self.tokenizer tokenizer self.async_reward_fn AsyncRewardWrapper(reward_fn) async def arun_episode( self, engine: InferenceEngine, data: dict[str, Any], ) - dict[str, Any] | None | dict[str, InteractionWithTokenLogpReward]: Run a single episode. MUST be async and non-blocking. # 1. Prepare input_ids from data input_ids self.tokenizer.apply_chat_template( data[messages], tokenizeTrue, add_generation_promptTrue, ) # 2. Build ModelRequest req ModelRequest( riduuid.uuid4().hex, input_idslist(input_ids), gconfigself.gconfig.new(n_samples1), tokenizerself.tokenizer, ) # 3. Generate completion (async) resp: ModelResponse await engine.agenerate(req) # 4. Compute reward (async) prompt_str self.tokenizer.decode(input_ids) completion_str self.tokenizer.decode(resp.output_tokens) reward await self.async_reward_fn( prompt_str, completion_str, resp.input_tokens, resp.output_tokens, **data, ) # 5. Return results in expected format return { input_ids: torch.tensor(resp.input_tokens), output_ids: torch.tensor(resp.output_tokens), reward: torch.tensor(reward), }模板五步流程与源码的对应关系模板步骤源码依据要点1. 构造input_idsrlvr.py 的get_input_ids_fn仓库参考实现通过apply_chat_template工具函数areal/utils/hf_utils.py做模板渲染并支持enable_thinking开关2. 构造ModelRequestio_struct.py L28-L45rid用uuid.uuid4().hexgconfig必须派生为n_samples1的请求级配置3. 异步生成engine_api.pyagenerate必须await禁止轮询或同步阻塞4. 异步算奖励reward_api.pyAsyncRewardWrapper奖励函数按约定签名reward_fn(prompt, completions, prompt_ids, completion_ids, **kwargs)传入**data会把数据集的其余字段如solution透传进去5. 组织轨迹张量rlvr.pyarun_episode返回带 batch 维1的张量字典2.1 奖励函数封装AsyncRewardWrapper的真实行为模板中AsyncRewardWrapper(reward_fn)不是简单的一层糖areal/api/reward_api.py 里的实现值得注意底层使用ProcessPoolExecutor把同步奖励函数丢到进程池执行因此奖励函数及其参数必须可 pickle默认timeout_seconds15、max_retries3超时先重试重试耗尽后记 warning 并返回 0.0而不是抛异常保证单条奖励失败不会炸掉整个 rollout 批次针对BrokenProcessPool有自动重建 executor 的恢复逻辑_recreate_executor进程池损坏后下次调用自动恢复max_workers默认按cpu_count // device_count // 2推导可按需显式传入。这也是技能文档Common Mistakes中未用AsyncRewardWrapper包裹奖励函数被列为常见错误的原因裸调同步奖励函数会在事件循环里阻塞。2.2 输出张量的完整约定以 RLVRWorkflow 为基准技能文档模板返回了最小三件套input_ids / output_ids / reward但仓库中的参考实现 areal/workflow/rlvr.py 展示了 PPO/GRPO 训练实际消费的完整字段集单轮 workflow 建议对齐seq resp.input_tokens resp.output_tokens logprobs [0.0] * resp.input_len resp.output_logprobs loss_mask [0] * resp.input_len [1] * resp.output_len versions [-1] * resp.input_len resp.output_versions turn_ids [-1] * resp.input_len [0] * resp.output_len res { input_ids: torch.tensor(seq, dtypetorch.int32), loss_mask: torch.tensor(loss_mask, dtypetorch.int32), logprobs: torch.tensor(logprobs, dtypetorch.float32), versions: torch.tensor(versions, dtypetorch.int32), turn_ids: torch.tensor(turn_ids, dtypetorch.int32), attention_mask: torch.ones(len(seq), dtypetorch.bool), rewards: torch.tensor(reward, dtypetorch.float32), is_truncated: torch.tensor(resp.stop_reason length, dtypetorch.bool), } return {k: v.unsqueeze(0) for k, v in res.items()}这里体现了技能文档 Key Requirements 中张量格式应为[batch, seq_len, ...]的约定arun_episode处理单条轨迹因此每个张量先构造[seq_len]再unsqueeze(0)补齐 batch 维多轨迹场景则用 areal/utils/data.py 的concat_padded_tensors合并——它会把非 batch 维右侧 padding 到逐维最大值后沿 dim 0 拼接attention_mask恒以 0 padding并且要求所有输入的字典键一致否则直接抛ValueError。Step 2在__init__.py中注册导出技能文档要求把类加入 areal/workflow/init.py 的__all__。需要注意当前仓库实际采用的是惰性导入模式L9-L24__all__ [ RLVRWorkflow, MultiTurnWorkflow, VisionRLVRWorkflow, ] _LAZY_IMPORTS { RLVRWorkflow: areal.workflow.rlvr, MultiTurnWorkflow: areal.workflow.multi_turn, VisionRLVRWorkflow: areal.workflow.vision_rlvr, } def __getattr__(name: str): if name in _LAZY_IMPORTS: import importlib module importlib.import_module(_LAZY_IMPORTS[name]) val getattr(module, name) globals()[name] val return val raise AttributeError(fmodule {__name__!r} has no attribute {name!r})因此新增MyWorkflow时比技能文档中的直接 import 写法更贴合仓库现状的做法是在__all__中追加MyWorkflow并在_LAZY_IMPORTS中映射MyWorkflow: areal.workflow.name。惰性导入的意义在于避免包导入时拉起重量级依赖如各 workflow 模块各自的第三方 SDKareal.workflow.MyWorkflow的字符串路径解析后续 Step 3 会用到依然可以命中__getattr__正常工作。Step 3在训练入口脚本中按字符串路径接入workflow 在整个训练栈中以WorkflowLike形式流转其类型定义见 workflow_api.py L124-L132既接受RolloutWorkflow实例/类也接受字符串导入路径或任何具有兼容async run()方法的对象。字符串路径方式正是技能文档 Step 3 推荐的用法仓库示例 examples/math/gsm8k_rl.py 展示了真实调用with PPOTrainer( config, train_datasettrain_dataset, valid_datasetvalid_dataset, ) as trainer: trainer.train( workflowareal.workflow.openai.math_agent.MathAgent, workflow_kwargsworkflow_kwargs, eval_workflowareal.workflow.openai.math_agent.MathAgent, eval_workflow_kwargseval_workflow_kwargs, )换成自定义 workflow 时把字符串替换为areal.workflow.name.MyWorkflow并可通过workflow_kwargs透传构造参数如temperature、top_p等trainer.train( workflowareal.workflow.name.MyWorkflow, workflow_kwargsdict( temperatureconfig.gconfig.temperature, top_pconfig.gconfig.top_p, ), )rollout 侧实际由 areal/infra/workflow_executor.py 的WorkflowExecutor按 worker 并行调度每个 episode测试侧可以参考 tests/test_workflow_executor_filtering.py 中对 executor 行为包括被 workflow 拒绝/过滤的轨迹处理的验证方式。Step 4编写测试技能文档建议创建tests/test_name_workflow.py最小骨架import pytest from areal.workflow.name import MyWorkflow pytest.mark.asyncio async def test_workflow_basic(): # Test basic functionality pass实际编写时建议至少覆盖三类断言均有仓库内的成熟参照arun_episode张量契约mock 一个InferenceEngine其agenerate返回一个手工构造的ModelResponse断言返回字典的键集、batch 维 1、loss_mask前input_len位全 0 等奖励链路验证同步奖励函数经AsyncRewardWrapper后能被正确await以及超时返回 0.0 的降级路径reward_api.py L154-L161拒绝路径构造数据触发return None确认该轨迹被丢弃而批次仍可推进——可对照 tests/test_grouped_rollout_workflow.py 中分组 rollout 的完整/不完整场景处理。三、参考实现对照表技能文档 Reference Implementations 原文继承Workflow文件适用场景MultiTurnWorkflowareal/workflow/multi_turn.py多轮对话 rolloutRLVRWorkflowareal/workflow/rlvr.py可验证奖励的 RL单轮支持 thinking token 开关VisionRLVRWorkflowareal/workflow/vision_rlvr.py视觉 RLVR涉及ModelRequest的image_data/processor字段选型建议写单轮任务以RLVRWorkflow为第一参照它同时演示了reward_fn字符串动态加载、get_input_ids_fn可注入、stats_tracker打点与 session tracing需要多轮/工具调用时研究MultiTurnWorkflow如何在一个 episode 内循环agenerate并组织turn_ids涉及图像输入时以VisionRLVRWorkflow为模板填充 VLM 相关字段。四、关键要求与常见错误清单原文继承 源码印证技能文档 Key Requirements 五条逐条给出仓库层面的解释Asyncarun_episode必须async def且非阻塞——rollout worker 以并发协程方式并发执行大量 episode阻塞调用会串行化整批生成No sync I/O文件操作一律aiofilesWrap rewards用AsyncRewardWrapper包裹奖励函数获得进程池隔离 超时默认 15s 重试默认 3 次 进程池自愈详见第二节 2.1Tensor format输出张量为[batch, seq_len, ...]单 episode 结果记得unsqueeze(0)见 rlvr.py L183Use helpers多字典/多变长序列合并用 concat_padded_tensors避免手写 padding 逻辑出错。Common Mistakes 四条对应的失败模式用open()而非aiofiles.open()同步文件 I/O 冻结事件循环忘记awaitagenerate、奖励调用漏await会把协程对象/未完成 future 当结果用类型错误往往在 batch 聚合阶段才爆发未包裹奖励函数裸调同步奖励在 CPU 密集打分时阻塞同进程其他 episode张量 shape 约定错误忘记 batch 维、loss_mask与input_len对齐错误、attention_mask未按concat_padded_tensors的 0-padding 约定构造都会让训练侧 batch 校验失败。五、动手前检查清单Prerequisites 自检开始编码前按技能文档 Prerequisites 先确认三件事workflow 的目的与要求单轮/多轮/工具调用/视觉、输入输出数据格式数据集字段 →data字典 → 轨迹张量字典的映射、要用的奖励函数签名与是否可 pickle。完成开发后建议按以下清单自检arun_episode是async def内部所有 I/O生成、奖励、文件均可等待点返回None的路径语义明确 拒绝轨迹不进入训练轨迹字典包含训练侧必需字段input_ids、loss_mask、logprobs、rewards等可截停场景附带is_truncated已在 areal/workflow/init.py 的__all__与_LAZY_IMPORTS中注册入口脚本以areal.workflow.name.MyWorkflow字符串路径接入trainer.train测试覆盖正常路径、奖励降级路径与轨迹拒绝路径。掌握以上流程后为 AReaL 增加任何形态的 rollout 行为新的任务环境、新的交互模式、新的轨迹组织方式都只是一件新增一个文件 两处注册 一个测试的模块化操作训练引擎、权重同步与调度基础设施均无需改动。【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表