AReaL 自定义 RolloutWorkflow 开发指南:从编写、注册到接入训练的完整四步流程
【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple & Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL
本文基于 AReaL 仓库内置的add-workflow技能文档(.claude/skills/add-workflow/SKILL.md),系统讲解如何在 AReaL 中新增一个RolloutWorkflow实现:包括工作流抽象契约的源码解析、标准代码模板、包注册方式、训练脚本接入路径与测试编写。读完本文,你可以独立为自己的强化学习任务(数学推理、多轮对话、视觉 RLVR 等)编写一个异步、非阻塞、符合张量输出规范的自定义 rollout 工作流,并让它无缝接入 AReaL 的 GRPO/PPO 训练链路。
何时需要新增一个 Workflow
AReaL 将"从一条样本数据出发,调用推理引擎生成响应、计算奖励、组装成训练用张量轨迹"这一整套逻辑抽象为RolloutWorkflow。当你的任务属于以下情况时,就需要新增一个工作流实现:
- 现有的
RLVRWorkflow(单轮可验证奖励)、MultiTurnWorkflow(多轮重试)、VisionRLVRWorkflow(视觉 + RLVR)无法满足你的交互逻辑(例如需要工具调用循环、沙箱执行、多阶段推理等); - 你需要自定义输入数据的解析方式(
data字典到input_ids的转换); - 你需要自定义奖励计算与轨迹组装逻辑(例如按轮次折现奖励、拼接多轮 token 序列)。
开始编写前,技能文档要求先明确三件事:工作流的目的与需求、输入/输出数据格式、要使用的奖励函数。
理解抽象契约:arun_episode 的签名与返回值
新工作流必须继承 RolloutWorkflow 并实现唯一的抽象方法arun_episode:
class RolloutWorkflow(ABC): @abstractmethod async def arun_episode( self, engine: InferenceEngine, data: dict[str, Any] ) -> dict[str, Any] | None | dict[str, InteractionWithTokenLogpReward]: ...从 areal/api/workflow_api.py 的 docstring 可以看出几个契约要点,编写自定义工作流时务必遵守:
- 必须异步:
arun_episode是async def,内部所有 I/O(生成、奖励、文件读写)都必须非阻塞,否则会拖垮并发的 rollout 调度; - 返回
None表示拒绝:返回None意味着该轨迹被拒收,不进入训练; - 截断标记:若工作流能判断模型响应是否因达到长度上限而停止,应在张量结果中提供
is_truncated布尔张量(每条轨迹一个值)。PPO 用它做 reward masking、value bootstrapping 和截断指标统计; - 奖励归一化场景:如果行级奖励不同且启用了奖励归一化,应在张量结果中提供一个有限的
rollout_reward标量作为组/批奖励统计的参考值,各行保留各自奖励。
此外,WorkflowLike类型别名(areal/api/workflow_api.py#L127-L132)表明workflow参数既接受RolloutWorkflow实例/类,也接受字符串导入路径,这决定了后文训练脚本中workflow="areal.workflow.<name>.MyWorkflow"的字符串写法是官方支持的。
第一步:创建工作流文件
按技能文档,新建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, torch.Tensor]: """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"], tokenize=True, add_generation_prompt=True, ) # 2. Build ModelRequest req = ModelRequest( rid=uuid.uuid4().hex, input_ids=list(input_ids), gconfig=self.gconfig.new(n_samples=1), tokenizer=self.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), }模板涉及的核心数据结构,可以从源码进一步确认:
ModelRequest(areal/api/io_struct.py#L29-L60):除了模板中的rid、input_ids、gconfig、tokenizer外,还带有metadata(透传自定义信息)、image_data/processor(VLM 图像输入)等字段。多模态工作流可直接复用;ModelResponse(areal/api/io_struct.py#L64-L92):生成结果包含input_tokens、output_tokens、output_logprobs、output_versions,以及stop_reason(取值为"length" / "stop" / "tool_calls" / "abort")。判断is_truncated的标准写法就是resp.stop_reason == "length";gconfig的预处理:gconfig.new_with_stop_and_pad_token_ids(tokenizer)(定义于 areal/api/cli_args.py 的GenerationHyperparameters类)会在生成超参上注入 tokenizer 的 stop/pad token id,这是保证引擎正确终止生成的前提,不要省略;AsyncRewardWrapper(areal/api/reward_api.py#L62-L100):它将同步奖励函数包装为异步调用,底层使用ProcessPoolExecutor进程池,默认timeout_seconds=15、max_retries=3,并具备 broken pool 自动重建能力。注意源码注释明确指出:奖励函数及其参数必须可 pickle(会被分发到 worker 进程),且奖励计算不会阻塞事件循环。
第二步:在 areal/workflow/init.py 中注册
技能文档给出的注册写法是直接from areal.workflow.<name> import MyWorkflow并加入__all__。而当前仓库中 areal/workflow/init.py 实际采用的是惰性导入模式:
__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(f"module {__name__!r} has no attribute {name!r}")因此在当前代码库中新增导出,推荐按现有模式操作:在__all__中加入"MyWorkflow",并在_LAZY_IMPORTS中加入"MyWorkflow": "areal.workflow.<name>"。这样import areal.workflow不会连带加载torch/transformers等重依赖,只有真正访问MyWorkflow时才导入对应模块。如果你的工作流文件依赖较轻,直接 eager import 也能工作(__getattr__只兜底未直接导入的名字),但跟随惰性模式与现有代码风格保持一致。
第三步:在训练脚本中引用新工作流
注册完成后,在训练入口脚本中通过字符串导入路径引用:
trainer.train( workflow="areal.workflow.<name>.MyWorkflow", # ... other args )这个字符串会在训练器内部经import_from_string动态导入。从 areal/trainer/rl_trainer.py 的_requires_proxy_workflow实现可以看到完整的判定逻辑:
workflow是RolloutWorkflow实例或子类 → 直接作为 rollout 工作流执行,与 rollout worker 同机调度;workflow是字符串 → 用import_from_string尝试导入,导入结果若是RolloutWorkflow同样按 rollout 工作流处理;导入失败时按 fail-safe 策略当作需要 proxy worker 的 agent 工作流;- 其他任意带兼容
run()方法的对象 → 走 OpenAI 兼容 proxy worker 路径(RolloutController.start_proxy())。
所以只要你继承RolloutWorkflow,字符串路径即可被正确识别,无需额外配置。仓库中的真实用法可参考 examples/math/gsm8k_rl.py 中的workflow="areal.workflow.openai.math_agent.MathAgent"与 examples/vlm/geometry3k_grpo.py 中的workflow="areal.workflow.vision_rlvr.VisionRLVRWorkflow";docs/en/tutorial/gsm8k_grpo.md 中的workflow="areal.workflow.rlvr.RLVRWorkflow"也是同一机制。
第四步:编写测试
技能文档建议在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仓库中已有可直接对照的测试先例:
- tests/test_workflow_detection.py 验证
workflow = "areal.workflow.rlvr.RLVRWorkflow"这类字符串路径能被工作流检测逻辑正确分类; - tests/test_rollout_controller.py 中以
workflow="areal.workflow.rlvr.RLVRWorkflow"驱动完整的 rollout controller 端到端流程。
自定义工作流的测试建议覆盖:arun_episode返回键的完整性(input_ids/rewards等)、批次维unsqueeze(0)是否到位、奖励函数被AsyncRewardWrapper正确包装、以及is_truncated在stop_reason == "length"时为True。
关键要求与常见错误(含源码佐证)
技能文档列出的五条硬性要求,均可在源码中找到对应实现依据:
| 要求 | 说明 | 源码佐证 |
|---|---|---|
| 必须 async | arun_episode必须是async def且非阻塞 | areal/api/workflow_api.py#L15-L18 中async def arun_episode抽象签名 |
| 禁止同步 I/O | 文件操作改用aiofiles | RLVRWorkflow全流程await engine.agenerate/await self.async_reward_fn,无同步阻塞调用,见 areal/workflow/rlvr.py |
| 奖励必须包装 | 用AsyncRewardWrapper包装奖励函数 | 基于ProcessPoolExecutor,带超时/重试/pool 重建,见 areal/api/reward_api.py#L62-L100 |
张量格式[batch, seq_len, ...] | 输出张量带批次维 | 参考实现末尾统一return {k: v.unsqueeze(0) for k, v in res.items()},见 areal/workflow/rlvr.py#L183 |
使用concat_padded_tensors | 多路输出合并时用统一工具函数 | 定义于 areal/utils/data.py#L245-L261:对非批次维做右侧 padding 后沿 dim 0 拼接,attention_mask恒为 0 填充,且要求所有输入字典 key 完全一致 |
常见错误清单(技能文档原文继承):
- 用
open()代替aiofiles.open()做文件读写,阻塞事件循环; - 忘记
await异步调用,导致拿到协程对象而非结果; - 奖励函数未用
AsyncRewardWrapper包装,同步阻塞 rollout 并丢失超时/重试保护; - 张量 shape 约定错误(缺少批次维、
loss_mask/attention_mask类型不对)。
补充一点来自源码的隐性约束:由于奖励经ProcessPoolExecutor分发,奖励函数及其入参必须可 pickle,本地变量、lambda、未定义的闭包对象都不能作为奖励函数传给工作流。
参考实现深度解析
技能文档给出的参考实现表与当前仓库一一对应:
| 工作流 | 文件 | 适用场景 |
|---|---|---|
MultiTurnWorkflow | areal/workflow/multi_turn.py | 多轮对话 / 错误重试 |
RLVRWorkflow | areal/workflow/rlvr.py | 单轮 RL with verifiable rewards |
VisionRLVRWorkflow | areal/workflow/vision_rlvr.py | 视觉 + RLVR |
以RLVRWorkflow为例(areal/workflow/rlvr.py#L49-L183),它是自定义工作流最贴近生产形态的范本:
- 可插拔钩子:
get_input_ids_fn与data_extract_prompt_fn均支持传字符串路径,内部用import_from_string动态加载,默认分别走apply_chat_template和data["messages"]提取;reward_fn同样支持字符串路径,在arun_episode首次调用时惰性加载; - 完整的轨迹张量集合:返回
input_ids、loss_mask(输入侧 0 / 输出侧 1)、logprobs(输入侧填 0.0)、versions(输入侧填 -1)、turn_ids(输入侧 -1 / 输出侧 0)、attention_mask、rewards、is_truncated九种张量,全部int32/float32/bool且带批次维; - 指标与追踪:
stats_tracker.get(workflow_context.stat_scope()).scalar(reward=reward)上报奖励指标,并用@trace_session("reward")/atrace_session_phase("generate")装饰器为 SessionTracer 提供生成与奖励两阶段的性能追踪。
MultiTurnWorkflow(areal/workflow/multi_turn.py#L19-L142)则演示了更复杂的轨迹组装:在max_turns内循环调用engine.agenerate,每轮把前缀 token、output_logprobs、loss_mask、versions、turn_ids增量拼接,turn_ids按轮次编号;构造器中通过"差集截取"(s2[len(s1):])预计算追加轮次的 prompt token,消除了 encode-decode 不一致问题;最终奖励按turn_discount的指数折现累计。若你的任务需要"多次尝试、逐步折现"的奖励结构,这个实现是最直接的参照。
结语:新增工作流检查清单
综合技能文档与源码事实,一个新RolloutWorkflow交付前建议逐项核对:
arun_episode为async def,生成与奖励均为await调用,无同步 I/O;gconfig经过new_with_stop_and_pad_token_ids(tokenizer)预处理,ModelRequest.gconfig上显式new(n_samples=1);- 奖励函数经
AsyncRewardWrapper包装且可 pickle; - 返回字典的 key 与参考实现保持一致,张量带批次维,含
loss_mask、is_truncated等 PPO 所需元数据; - 已在 areal/workflow/init.py 的
__all__与_LAZY_IMPORTS中注册; - 训练脚本以
"areal.workflow.<name>.MyWorkflow"字符串路径引用; tests/test_<name>_workflow.py覆盖基础功能,可对照 tests/test_workflow_detection.py 与 tests/test_rollout_controller.py 的既有模式。
完成以上步骤后,自定义工作流即可与 AReaL 的 GRPO/PPO 训练链路、rollout controller 及统计追踪体系直接协作,无需修改框架代码。
【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple & Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考