news 2026/9/17 4:53:54

AReaL 新数据集加载器开发指南:从 Loader 实现到注册、配置与测试的完整流程

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
AReaL 新数据集加载器开发指南:从 Loader 实现到注册、配置与测试的完整流程

AReaL 新数据集加载器开发指南:从 Loader 实现到注册、配置与测试的完整流程

【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple & Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL

本文基于 AReaL 仓库内置的add-dataset技能文档(.agents/skills/add-dataset/SKILL.md),系统讲解如何在 AReaL 中为 RL / SFT / 对齐训练接入一个全新的数据集加载器:如何编写 SFT 与 RL 两种 loader、如何在数据集注册表中完成分派、如何按需扩展配置项以及如何补齐测试。读完本文,你可以独立完成一个新数据集从代码到训练入口的完整接入,并理解 AReaL 数据集框架的底层分派机制。

一、何时使用本技能

add-dataset技能在以下场景触发(摘自原文档):

  • 用户询问“如何给 AReaL 添加数据集?”
  • 用户希望集成一个新的训练数据集
  • 用户提到要创建 dataset loader

AReaL 的数据集模块位于 areal/dataset/ 目录,内置了 GSM8K、Geometry3K、CLEVR、HH-RLHF、ToRL 等多个参考 loader(完整列表见 areal/dataset/init.py 中的VALID_DATASETS)。接入新数据集的标准路径是:创建 loader 文件 → 注册分派 → 可选配置扩展 → 补充测试,下文逐步展开。

二、AReaL 数据集加载架构:先理解你要接入的框架

在动手写 loader 之前,有必要先看清 AReaL 是如何把配置转化为数据集对象的。整条调用链为:

训练入口 (如 examples/math/gsm8k_rl.py) └─ get_custom_dataset(split, dataset_config, tokenizer, ...) # areal/dataset/__init__.py ├─ 多数据源 sources → get_routed_dataset(...) # MOPD 路由混采 ├─ 单控制器 + scheduling_spec → RDataset(远程数据服务) └─ _get_custom_dataset(path, type, split, ...) # 按路径/类型分派到具体 loader └─ 分派失败时回退 load_from_disk(通用 HF 磁盘数据集)

2.1 入口函数get_custom_dataset的真实签名

训练入口通过 get_custom_dataset 获取数据集。以 examples/math/gsm8k_rl.py 为例:

train_dataset = get_custom_dataset( split="train", dataset_config=config.train_dataset, tokenizer=tokenizer, )

其内部有三条分支(源码见 areal/dataset/init.py):

  1. 多数据源模式:若dataset_config.sources非空(例如 MOPD 教师路由混采),直接走get_routed_dataset,不经过单个 loader;
  2. 单控制器 +scheduling_spec模式:返回RDataset,把数据加载放到远程>if "gsm8k" in path and type == "sft": from .gsm8k import get_gsm8k_sft_dataset return get_gsm8k_sft_dataset(path=path, split=split, tokenizer=tokenizer, max_length=max_length, **kwargs) elif "gsm8k" in path and type == "rl": from .gsm8k import get_gsm8k_rl_dataset ... elif "hh-rlhf" in path and type == "rw": from .hhrlhf import get_hhrlhf_rw_dataset ...

    由此可以归纳出三条接入要点:

    • type目前实际使用的取值包括sftrlrwdpo(对应 SFT、强化学习、Reward Weighted 与 DPO 对齐训练),新 loader 应明确自己支持哪一种;
    • path中必须能稳定命中你选定的关键字子串——例如openai/gsm8k命中"gsm8k",这也是为什么 examples/math/gsm8k_grpo.yaml 中train_dataset.path直接写openai/gsm8k
    • 部分数据集对路径匹配做了更精细的处理。从源码结构看,SWE 数据使用正则(?:^|[/_\-.])swe(?:[/_\-.]|$)精确匹配路径 token(areal/dataset/init.py),以避免answer_sft/home/swetha/这类仅包含swe三字母片段的路径被误派发到 SWE 轨迹管线;
    • 所有分支都未命中时会回退datasets.load_from_disk(path),尝试按“通过dataset.save_to_disk()保存的通用 HuggingFace 数据集”加载;再失败则抛出包含VALID_DATASETS列表的ValueError(areal/dataset/init.py)。也就是说,如果你的数据本身已是标准 HF 格式且无需预处理,甚至可以不写 loader 直接走回退路径。

    2.3 数据集配置_DatasetConfig

    训练 YAML 中的train_dataset:/valid_dataset:段由 _DatasetConfig(位于 areal/api/cli_args.py)解析,常用字段包括:

    字段默认值说明
    split"train"使用的数据集 split(train / validation / test)
    pathNone数据集路径(HF Hub 名或本地路径),与sources互斥
    typeNone训练数据类型(如rlsft),与sources互斥
    sources[]多数据源混合列表(MOPD 场景,每个 source 需声明teacher_group
    mixture_sampling_policy"proportional"混合采样策略:proportional按源规模比例,uniform循环补齐较短源
    batch_size1dataloader 批大小
    shuffleTrue是否打乱
    pin_memoryFalse是否 pin memory(GPU 训练建议开启)
    num_workers0数据加载 worker 进程数

    其中max_length会透传给你的 loader,用于过滤超长样本。一个真实的最小配置来自 examples/math/gsm8k_grpo.yaml:

    train_dataset: batch_size: 256 shuffle: true pin_memory: true num_workers: 4 path: openai/gsm8k type: rl max_length: 1024

    三、Step 1:创建数据集文件areal/dataset/<name>.py

    技能文档给出的标准模板包含一对函数:get_<name>_sft_dataset(面向 SFT 的完整序列 tokenization)与get_<name>_rl_dataset(面向 RL 的 prompt + 答案结构)。完整模板如下(继承自原文档):

    from datasets import Dataset, load_dataset def get_<name>_sft_dataset( path: str, split: str, tokenizer, max_length: int | None = None, ) -> Dataset: """Load dataset for SFT training. Args: path: Path to dataset (HuggingFace hub or local path) split: Dataset split (train/validation/test) tokenizer: Tokenizer for processing max_length: Maximum sequence length (optional) Returns: HuggingFace Dataset with processed samples """ dataset = load_dataset(path=path, split=split) def process(sample): # Tokenize the full sequence (prompt + response) seq_token = tokenizer.encode( sample["question"] + sample["answer"] + tokenizer.eos_token ) prompt_token = tokenizer.encode(sample["question"]) # Loss mask: 0 for prompt, 1 for response loss_mask = [0] * len(prompt_token) + [1] * (len(seq_token) - len(prompt_token)) return {"input_ids": seq_token, "loss_mask": loss_mask} dataset = dataset.map(process).remove_columns(["question", "answer"]) if max_length is not None: dataset = dataset.filter(lambda x: len(x["input_ids"]) <= max_length) return dataset def get_<name>_rl_dataset( path: str, split: str, tokenizer, max_length: int | None = None, ) -> Dataset: """Load dataset for RL training. Args: path: Path to dataset split: Dataset split tokenizer: Tokenizer for length filtering max_length: Maximum sequence length Returns: HuggingFace Dataset with prompts and answers for reward computation """ dataset = load_dataset(path=path, split=split) def process(sample): messages = [ { "role": "user", "content": sample["question"], } ] return {"messages": messages, "answer": sample["answer"]} dataset = dataset.map(process).remove_columns(["question"]) if max_length is not None: def filter_length(sample): content = sample["messages"][0]["content"] tokens = tokenizer.encode(content) return len(tokens) <= max_length dataset = dataset.filter(filter_length) return dataset

    两种 loader 的关键设计点:

    SFT loader:对“问题 + 答案 + EOS”整体做 tokenization,并生成loss_mask——prompt 部分置 0、response 部分置 1,保证 SFT 只在回答段计算损失。这正是内置 get_gsm8k_sft_dataset 的做法:loss_mask = [0] * len(prompt_token) + [1] * (len(seq_token) - len(prompt_token)),处理完成后只保留input_idsloss_mask两列。

    RL loader:输出messages(OpenAI 风格对话列表,供 rollout 生成)与answer(ground truth,供 reward 函数判分)。RL 阶段不需要完整序列,因此长度过滤只对 user 内容做 tokenization 检查,成本更低。内置 get_gsm8k_rl_dataset 在 user 消息中还追加了输出格式约束("Please put your final answer within \\boxed{}."),这是把“任务指令”与“原始问题”融合进 prompt 的典型做法,新数据集若有固定作答格式要求可参照此模式。

    四、Step 2:在areal/dataset/__init__.py中注册

    注册分两处(原文档要求):

    1. 把数据集名加入VALID_DATASETS列表(areal/dataset/init.py)。该列表同时用于报错提示——回退加载失败时,异常信息会打印Supported datasets are: {VALID_DATASETS},方便使用者快速定位。
    2. _get_custom_dataset的分派链中为你的数据源增加分支。技能文档的写法是:
    # Add to VALID_DATASETS VALID_DATASETS = [ # ... existing datasets "<name>", ] # Add to _get_custom_dataset function def _get_custom_dataset(name: str, ...): # ... existing code elif name == "<name>": from areal.dataset.<name> import get_<name>_sft_dataset, get_<name>_rl_dataset if dataset_type == "sft": return get_<name>_sft_dataset(path, split, max_length, tokenizer) else: return get_<name>_rl_dataset(path, split, max_length, tokenizer)

    结合 2.2 节的源码事实,落笔时的实际形式应为“路径子串 + type”的双重条件分支,并保持函数内延迟导入(与现有分支一致,避免未使用数据集时产生依赖副作用):

    elif "<name>" in path and type == "sft": from .<name> import get_<name>_sft_dataset return get_<name>_sft_dataset( path=path, split=split, tokenizer=tokenizer, max_length=max_length, **kwargs, ) elif "<name>" in path and type == "rl": from .<name> import get_<name>_rl_dataset return get_<name>_rl_dataset( path=path, split=split, tokenizer=tokenizer, max_length=max_length, **kwargs, )

    两个实践提示:

    • 子串关键字要选得有辨识度,避免与现有路径误撞;如果关键字过于通用(如三字母片段),可参考 SWE 分支的做法,用正则限定为完整路径 token(areal/dataset/init.py),相关行为可由 tests/test_dataset_swe_path_dispatch.py 这类分派测试验证。
    • 多模态数据集用processor而不是tokenizer:参考 get_geometry3k_sft_dataset,它通过processor完成图文 tokenization、生成pixel_values/image_grid_thw等多模态输入,并用get_multimodal_sft_loss_mask计算跨模态 loss mask;而 get_torl_data_rl_dataset 则演示了本地 parquet 文件的加载方式(load_dataset("parquet", data_files=path, ...))以及“SFT 不支持时显式raise NotImplementedError”的写法。

    五、Step 3(可选):为数据集添加专属配置

    技能文档建议:如果数据集需要特殊配置,在 areal/api/cli_args.py 中扩展配置 dataclass:

    @dataclass class TrainDatasetConfig: # ... existing fields <name>_specific_field: Optional[str] = None

    对照源码现状:AReaL 中该配置类的实际名称是 _DatasetConfig,它采用“通用字段 +**kwargs透传”的开放设计——get_custom_dataset会把额外的**kwargs原样转发给_get_custom_dataset,最终由你的 loader 通过**kwargs接收(如 get_gsm8k_sft_dataset 的签名末尾)。因此对于轻量场景,你也可以不修改_DatasetConfig,而是在 YAML 的train_dataset:段直接传自定义键(经由dataset_kwargs/kwargs透传),仅在字段需要参与 CLI 校验或默认值治理时才扩展到_DatasetConfig

    六、Step 4:补充测试

    为每个新 loader 创建tests/test_<name>_dataset.py。原文档给出的最小测试模板:

    import pytest from areal.dataset.<name> import get_<name>_sft_dataset, get_<name>_rl_dataset def test_sft_dataset_loads(tokenizer): dataset = get_<name>_sft_dataset("path/to/data", split="train", tokenizer=tokenizer) assert len(dataset) > 0 assert "input_ids" in dataset.column_names assert "loss_mask" in dataset.column_names def test_rl_dataset_loads(tokenizer): dataset = get_<name>_rl_dataset("path/to/data", split="train", tokenizer=tokenizer) assert len(dataset) > 0 assert "messages" in dataset.column_names assert "answer" in dataset.column_names

    要点是断言列名契约而非具体内容:SFT 数据集必须提供input_idsloss_mask,RL 数据集必须提供messagesanswer。仓库中已有大量数据集测试可作参照,例如 tests/test_mopd_dataset.py、tests/test_swe_sft_dataset.py 覆盖 loader 行为,tests/test_dataset_swe_path_dispatch.py 覆盖路径分派逻辑——新增数据集时建议同时覆盖这两层。

    七、参考实现对照表

    技能文档给出的参考实现一览(已对照仓库确认路径):

    数据集文件说明技术特点(源自实现)
    GSM8Kareal/dataset/gsm8k.py数学应用题文本双 loader 的标准范式:SFT 输出input_ids/loss_mask,RL 输出messages
    Geometry3Kareal/dataset/geometry3k.py几何题(图文)多模态processor加载,图像裁剪/RGB 化,跨模态 loss mask
    CLEVR-Count-70Kareal/dataset/clevr_count_70k.py视觉计数多模态计数任务
    HH-RLHFareal/dataset/hhrlhf.py有用性/无害性偏好数据支持rwdpo两种对齐类型
    ToRLareal/dataset/torl_data.py工具调用 RL本地 parquet 加载、rank0 下载 + 成功标志同步,SFT 显式不支持

    其中 GSM8K 是最值得精读的最小范本(约 60 行):它同时示范了load_dataset参数(name="main")、dataset.map(process)+remove_columns的向量化处理、以及max_length过滤。

    八、数据集字段契约与常见错误

    技能文档明确了两类数据集的必备字段(Required Fields):

    SFT 数据集(面向 pipeline 的消息形态描述):

    { "messages": [ {"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}, ] }

    RL 数据集

    { "messages": [ {"role": "user", "content": "..."}, ], "answer": "ground_truth_for_reward", # Optional metadata for reward function }

    对应到 loader 实现层面,即:RL 数据集必须携带messagesrole/content的字典列表,供 rollout 直接消费)和answer(供 reward 函数判分);SFT 数据经 loader 处理后应提供训练可消费的 token 序列与 loss mask。answer之外的可选元数据列(如data_sourceability)可以原样保留给 reward 函数使用——areal/dataset/torl_data.py 的列注释即为这种“prompt + reward 元数据”结构的实例。

    原文档最后列出的高频错误,全部保留如下,并附一条排查线索:

    常见错误后果 / 排查线索
    返回List[Dict]而非 HuggingFaceDataset下游dataset.map/filter、DataLoader 均依赖 HF Dataset API,列表会直接报错
    用 Python for 循环逐条处理而非dataset.map()/filter()丧失 HF datasets 的向量化与缓存能力,大规模数据下极慢
    RL 数据集缺少"messages"字段rollout 侧无法构造生成请求
    消息格式错误(应为带rolecontent的字典列表)与 AReaL workflow 层的 OpenAI 风格消息约定不一致
    忘记在__init__.py注册训练时报Dataset ... is not supported. Supported datasets are: [...](回退load_from_disk也失败时抛出,见 areal/dataset/init.py)

    九、端到端接入检查清单

    把以上步骤串起来,一个新数据集在 AReaL 中的完整落地路径为:

    1. 写 loader:新建areal/dataset/<name>.py,实现get_<name>_sft_dataset/get_<name>_rl_dataset,返回 HFDataset,SFT 带input_ids+loss_mask,RL 带messages+answer
    2. 注册分派:更新 areal/dataset/init.py 的VALID_DATASETS,并在_get_custom_dataset中新增“路径子串 + type”分支(延迟导入);
    3. 配置训练:在实验 YAML 中通过train_dataset.path(包含你的关键字子串)与train_dataset.type触发分派,并用max_lengthbatch_sizenum_workers等 _DatasetConfig 字段控制加载行为;
    4. 接入训练入口:入口脚本调用get_custom_dataset(split=..., dataset_config=..., tokenizer=...)(参考 examples/math/gsm8k_rl.py);
    5. 补测试tests/test_<name>_dataset.py覆盖列名契约,必要时补路径分派测试。

    需要说明的适用前提:本文所有签名与字段均基于当前仓库版本,type取值、scheduling_spec远程数据服务(RDataset)等分支仅在单控制器 + 数据服务部署下生效;多数据源sources模式属于 MOPD 路由混采场景,与常规单源接入互斥。新数据集若为纯文本且已保存为标准 HF 格式,也可利用load_from_disk回退路径直接接入而不编写专属 loader——但从可维护性与过滤控制角度看,仍建议按上述流程实现显式 loader。

    【免费下载链接】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 4:53:06

libcurl实战指南:从HTTPS传输到多线程的C/C++网络编程

简介&#xff1a;这份PDF教程是给C/C开发者的libcurl中文入门与进阶资料&#xff0c;基于官方教程翻译整理&#xff0c;并加入译者基于7.19.6版本的C示例代码。内容覆盖全局初始化与清理、编译链接选项、SSL支持检测、easy接口的handle创建与属性设置、multi接口使用要点、错误…

作者头像 李华
网站建设 2026/9/17 4:52:34

基于Matlab的风能资源评估实战:气象塔测风数据处理与指标计算

风电项目前期&#xff0c;一群人扛着设备在山上待几个月&#xff0c;图的是什么&#xff1f;就是那几十米高的气象塔上&#xff0c;几个风速仪和风向标记录下来的每一秒数据。这些从气象塔实测来的历史风力数据&#xff0c;是整个风能资源评估最原始、也最可靠的依据。后面无论…

作者头像 李华
网站建设 2026/9/17 4:52:26

Java GC优化实战:从内存生命周期到ZGC/G1选型

1. GC优化&#xff1a;不是调几个参数就完事&#xff0c;而是理解内存生命周期的实战工程“GC优化”这四个字在Java、Go、Python甚至前端JavaScript圈子里&#xff0c;几乎天天被提起&#xff0c;但真正能说清楚“我在优化什么”“为什么这个参数有效”“线上卡顿到底是不是GC惹…

作者头像 李华
网站建设 2026/9/17 4:51:28

上下文节流实战:如何将Agent的8万Token压缩至1600

最近在调一个多轮客服Agent&#xff0c;碰到一个很典型的问题&#xff1a;对话才跑了一上午&#xff0c;上下文就从几千token膨胀到8万多&#xff0c;账单肉眼可见地涨&#xff0c;响应还越来越慢。后来我把Context Mode接进去&#xff0c;同样的场景token直接压到1600左右&…

作者头像 李华
网站建设 2026/9/17 4:51:03

Agent 长链路评测:利用虚拟场景模拟多轮工具调用

Agent 长链路评测&#xff1a;利用虚拟场景模拟多轮工具调用在自主智能体&#xff08;Autonomous Agents&#xff09;从单步原型迈向能够独立处理复杂业务&#xff08;如自动化故障排障、多系统数据对账、端到端自动化测试&#xff09;的工业化落地阶段&#xff0c;算法团队面临…

作者头像 李华
网站建设 2026/9/17 4:50:14

MATLAB路标识别完整流程:HSV分割、形态学与模板匹配实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华