news 2026/9/17 19:55:33

AReaL 自定义 RolloutWorkflow 开发指南:从编写、注册到接入训练的完整四步流程

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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技能文档(.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 可以看出几个契约要点,编写自定义工作流时务必遵守:

  1. 必须异步arun_episodeasync def,内部所有 I/O(生成、奖励、文件读写)都必须非阻塞,否则会拖垮并发的 rollout 调度;
  2. 返回None表示拒绝:返回None意味着该轨迹被拒收,不进入训练;
  3. 截断标记:若工作流能判断模型响应是否因达到长度上限而停止,应在张量结果中提供is_truncated布尔张量(每条轨迹一个值)。PPO 用它做 reward masking、value bootstrapping 和截断指标统计;
  4. 奖励归一化场景:如果行级奖励不同且启用了奖励归一化,应在张量结果中提供一个有限的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):除了模板中的ridinput_idsgconfigtokenizer外,还带有metadata(透传自定义信息)、image_data/processor(VLM 图像输入)等字段。多模态工作流可直接复用;
  • ModelResponse(areal/api/io_struct.py#L64-L92):生成结果包含input_tokensoutput_tokensoutput_logprobsoutput_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=15max_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实现可以看到完整的判定逻辑:

  1. workflowRolloutWorkflow实例或子类 → 直接作为 rollout 工作流执行,与 rollout worker 同机调度;
  2. workflow是字符串 → 用import_from_string尝试导入,导入结果若是RolloutWorkflow同样按 rollout 工作流处理;导入失败时按 fail-safe 策略当作需要 proxy worker 的 agent 工作流
  3. 其他任意带兼容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_truncatedstop_reason == "length"时为True

关键要求与常见错误(含源码佐证)

技能文档列出的五条硬性要求,均可在源码中找到对应实现依据:

要求说明源码佐证
必须 asyncarun_episode必须是async def且非阻塞areal/api/workflow_api.py#L15-L18 中async def arun_episode抽象签名
禁止同步 I/O文件操作改用aiofilesRLVRWorkflow全流程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、未定义的闭包对象都不能作为奖励函数传给工作流。

参考实现深度解析

技能文档给出的参考实现表与当前仓库一一对应:

工作流文件适用场景
MultiTurnWorkflowareal/workflow/multi_turn.py多轮对话 / 错误重试
RLVRWorkflowareal/workflow/rlvr.py单轮 RL with verifiable rewards
VisionRLVRWorkflowareal/workflow/vision_rlvr.py视觉 + RLVR

RLVRWorkflow为例(areal/workflow/rlvr.py#L49-L183),它是自定义工作流最贴近生产形态的范本:

  • 可插拔钩子get_input_ids_fndata_extract_prompt_fn均支持传字符串路径,内部用import_from_string动态加载,默认分别走apply_chat_templatedata["messages"]提取;reward_fn同样支持字符串路径,在arun_episode首次调用时惰性加载;
  • 完整的轨迹张量集合:返回input_idsloss_mask(输入侧 0 / 输出侧 1)、logprobs(输入侧填 0.0)、versions(输入侧填 -1)、turn_ids(输入侧 -1 / 输出侧 0)、attention_maskrewardsis_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_logprobsloss_maskversionsturn_ids增量拼接,turn_ids按轮次编号;构造器中通过"差集截取"(s2[len(s1):])预计算追加轮次的 prompt token,消除了 encode-decode 不一致问题;最终奖励按turn_discount的指数折现累计。若你的任务需要"多次尝试、逐步折现"的奖励结构,这个实现是最直接的参照。

结语:新增工作流检查清单

综合技能文档与源码事实,一个新RolloutWorkflow交付前建议逐项核对:

  1. arun_episodeasync def,生成与奖励均为await调用,无同步 I/O;
  2. gconfig经过new_with_stop_and_pad_token_ids(tokenizer)预处理,ModelRequest.gconfig上显式new(n_samples=1)
  3. 奖励函数经AsyncRewardWrapper包装且可 pickle;
  4. 返回字典的 key 与参考实现保持一致,张量带批次维,含loss_maskis_truncated等 PPO 所需元数据;
  5. 已在 areal/workflow/init.py 的__all___LAZY_IMPORTS中注册;
  6. 训练脚本以"areal.workflow.<name>.MyWorkflow"字符串路径引用;
  7. 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/17 19:52:43

医学影像AI落地三重关:DICOM预处理、临床验证与PACS集成

简介&#xff1a;本资源是一篇聚焦深度学习在医学影像合成领域前沿进展的综述论文&#xff0c;面向医学AI方向的本科生毕业设计、研究生科研入门及临床工程技术人员&#xff0c;旨在系统梳理伪CT、合成MRI与合成PET三大核心任务的技术路径与挑战。全文基于2018–2023年主流研究…

作者头像 李华
网站建设 2026/9/17 19:47:20

Python判断语句if-else详解与实战应用

1. Python判断语句基础入门判断语句是编程中最基础也最重要的逻辑控制结构之一。作为Python入门者&#xff0c;掌握if-else的使用方法是写出实用代码的第一步。判断语句的本质是让程序具备"思考"能力&#xff0c;根据不同的条件执行不同的代码块。在实际开发中&#…

作者头像 李华
网站建设 2026/9/17 19:45:43

KubeEdge 依赖剖析:go-digest 内容寻址摘要包的原理与工程实践

KubeEdge 依赖剖析:go-digest 内容寻址摘要包的原理与工程实践 【免费下载链接】kubeedge Kubernetes Native Edge Computing Framework (project under CNCF) 项目地址: https://gitcode.com/GitHub_Trending/ku/kubeedge KubeEdge 作为 Kubernetes 原生边缘计算框架,其…

作者头像 李华
网站建设 2026/9/17 19:44:45

Spring Boot在线批改作业系统:从源码到部署的完整实战指南

Spring Boot在线批改作业系统&#xff0c;光是这个名字就能让不少正在做毕设或者课程设计的同学眼睛一亮。每年这时候后台总有人问有没有适合练手的Java后端项目&#xff0c;我基本都会推荐这种带完整业务闭环的管理系统——因为它既有用户角色区分&#xff0c;又有核心业务流&…

作者头像 李华