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

发布时间:2026/9/17 4:54:00
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 loaderAReaL 的数据集模块位于 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( splittrain, dataset_configconfig.train_dataset, tokenizertokenizer, )其内部有三条分支源码见 areal/dataset/init.py多数据源模式若dataset_config.sources非空例如 MOPD 教师路由混采直接走get_routed_dataset不经过单个 loader单控制器 scheduling_spec模式返回RDataset把数据加载放到远程>if gsm8k in path and type sft: from .gsm8k import get_gsm8k_sft_dataset return get_gsm8k_sft_dataset(pathpath, splitsplit, tokenizertokenizer, max_lengthmax_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目前实际使用的取值包括sft、rl、rw、dpo对应 SFT、强化学习、Reward Weighted 与 DPO 对齐训练新 loader 应明确自己支持哪一种path中必须能稳定命中你选定的关键字子串——例如openai/gsm8k命中gsm8k这也是为什么 examples/math/gsm8k_grpo.yaml 中train_dataset.path直接写openai/gsm8k部分数据集对路径匹配做了更精细的处理。从源码结构看SWE 数据使用正则(?:^|[/_\-.])swe(?:[/_\-.]|$)精确匹配路径 tokenareal/dataset/init.py以避免answer_sft、/home/swetha/这类仅包含swe三字母片段的路径被误派发到 SWE 轨迹管线所有分支都未命中时会回退到datasets.load_from_disk(path)尝试按“通过dataset.save_to_disk()保存的通用 HuggingFace 数据集”加载再失败则抛出包含VALID_DATASETS列表的ValueErrorareal/dataset/init.py。也就是说如果你的数据本身已是标准 HF 格式且无需预处理甚至可以不写 loader 直接走回退路径。2.3 数据集配置_DatasetConfig训练 YAML 中的train_dataset:/valid_dataset:段由 _DatasetConfig位于 areal/api/cli_args.py解析常用字段包括字段默认值说明splittrain使用的数据集 splittrain / validation / testpathNone数据集路径HF Hub 名或本地路径与sources互斥typeNone训练数据类型如rl、sft与sources互斥sources[]多数据源混合列表MOPD 场景每个 source 需声明teacher_groupmixture_sampling_policyproportional混合采样策略proportional按源规模比例uniform循环补齐较短源batch_size1dataloader 批大小shuffleTrue是否打乱pin_memoryFalse是否 pin memoryGPU 训练建议开启num_workers0数据加载 worker 进程数其中max_length会透传给你的 loader用于过滤超长样本。一个真实的最小配置来自 examples/math/gsm8k_grpo.yamltrain_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(pathpath, splitsplit) 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(pathpath, splitsplit) 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_ids与loss_mask两列。RL loader输出messagesOpenAI 风格对话列表供 rollout 生成与answerground truth供 reward 函数判分。RL 阶段不需要完整序列因此长度过滤只对 user 内容做 tokenization 检查成本更低。内置 get_gsm8k_rl_dataset 在 user 消息中还追加了输出格式约束Please put your final answer within \\boxed{}.这是把“任务指令”与“原始问题”融合进 prompt 的典型做法新数据集若有固定作答格式要求可参照此模式。四、Step 2在areal/dataset/__init__.py中注册注册分两处原文档要求把数据集名加入VALID_DATASETS列表areal/dataset/init.py。该列表同时用于报错提示——回退加载失败时异常信息会打印Supported datasets are: {VALID_DATASETS}方便使用者快速定位。在_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( pathpath, splitsplit, tokenizertokenizer, max_lengthmax_length, **kwargs, ) elif name in path and type rl: from .name import get_name_rl_dataset return get_name_rl_dataset( pathpath, splitsplit, tokenizertokenizer, max_lengthmax_length, **kwargs, )两个实践提示子串关键字要选得有辨识度避免与现有路径误撞如果关键字过于通用如三字母片段可参考 SWE 分支的做法用正则限定为完整路径 tokenareal/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_filespath, ...)以及“SFT 不支持时显式raise NotImplementedError”的写法。五、Step 3可选为数据集添加专属配置技能文档建议如果数据集需要特殊配置在 areal/api/cli_args.py 中扩展配置 dataclassdataclass 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, splittrain, tokenizertokenizer) 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, splittrain, tokenizertokenizer) assert len(dataset) 0 assert messages in dataset.column_names assert answer in dataset.column_names要点是断言列名契约而非具体内容SFT 数据集必须提供input_ids与loss_maskRL 数据集必须提供messages与answer。仓库中已有大量数据集测试可作参照例如 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_maskRL 输出messagesGeometry3Kareal/dataset/geometry3k.py几何题图文多模态processor加载图像裁剪/RGB 化跨模态 loss maskCLEVR-Count-70Kareal/dataset/clevr_count_70k.py视觉计数多模态计数任务HH-RLHFareal/dataset/hhrlhf.py有用性/无害性偏好数据支持rw与dpo两种对齐类型ToRLareal/dataset/torl_data.py工具调用 RL本地 parquet 加载、rank0 下载 成功标志同步SFT 显式不支持其中 GSM8K 是最值得精读的最小范本约 60 行它同时示范了load_dataset参数namemain、dataset.map(process)remove_columns的向量化处理、以及max_length过滤。八、数据集字段契约与常见错误技能文档明确了两类数据集的必备字段Required FieldsSFT 数据集面向 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_source、ability可以原样保留给 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 侧无法构造生成请求消息格式错误应为带role和content的字典列表与 AReaL workflow 层的 OpenAI 风格消息约定不一致忘记在__init__.py注册训练时报Dataset ... is not supported. Supported datasets are: [...]回退load_from_disk也失败时抛出见 areal/dataset/init.py九、端到端接入检查清单把以上步骤串起来一个新数据集在 AReaL 中的完整落地路径为写 loader新建areal/dataset/name.py实现get_name_sft_dataset/get_name_rl_dataset返回 HFDatasetSFT 带input_idsloss_maskRL 带messagesanswer注册分派更新 areal/dataset/init.py 的VALID_DATASETS并在_get_custom_dataset中新增“路径子串 type”分支延迟导入配置训练在实验 YAML 中通过train_dataset.path包含你的关键字子串与train_dataset.type触发分派并用max_length、batch_size、num_workers等 _DatasetConfig 字段控制加载行为接入训练入口入口脚本调用get_custom_dataset(split..., dataset_config..., tokenizer...)参考 examples/math/gsm8k_rl.py补测试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),仅供参考