ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

AReaL 数据集加载器开发指南:从 areal/dataset 源码看懂四步接入新数据集

AReaL 数据集加载器开发指南:从 areal/dataset 源码看懂四步接入新数据集 AReaL 数据集加载器开发指南从 areal/dataset 源码看懂四步接入新数据集【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL在 AReaL 中接入一个全新的训练数据集无论是数学题、几何题还是自定义语料本质上是实现一对get_name_sft_dataset/get_name_rl_dataset加载函数并注册进统一的分发入口。本文基于仓库内的技能文档 SKILL.md 与areal/dataset/的真实源码编写带你完整走通“新建数据集文件、注册分发、配置可选字段、编写测试”四步流程并深入解释loss_mask语义、messages约定、按路径子串分发的底层机制以及通用回退加载行为让你能独立为一个新数据集产出可复制、可运行、可测试的加载器。一、先理解 AReaL 的数据集加载入口在动手写代码前先确认新数据集最终会被谁调用。AReaL 中所有训练/验证数据都通过统一入口 get_custom_dataset 加载入口函数get_custom_dataset(split, dataset_config, tokenizer, processor, **kwargs)接收_DatasetConfig配置对象定义见 areal/api/cli_args.py根据scheduling_spec决定是否把加载下放到远端>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几个必须遵守的约定均可在仓库源码中找到对应依据返回 HuggingFaceDataset而非List[Dict]。AReaL 的 dataloader、save_to_disk回退加载、SWE 预分词管线都依赖 HF Dataset 的列式存储特性因此模板中统一用dataset.map()做向量化处理、用dataset.filter()做长度过滤避免 Python 循环。SFT 样本必须产出input_idsloss_mask。以 gsm8k.py 为例seq_token是“问题 答案 eos”的完整分词prompt_token是问题部分的分词loss_mask [0] * len(prompt_token) [1] * (len(seq_token) - len(prompt_token))——prompt 段置 0、response 段置 1训练时只对模型回答部分计算损失。这是 SFT 数据集的核心语义不要改动。RL 样本必须产出messages字段role/content 字典列表和answer字段。messages用于 prompt 构建与 rollout 生成answer是奖励函数reward计算所需的 ground truth。max_length是过滤条件而非截断条件超长样本直接丢弃SFT 按input_ids总长过滤RL 按第一条 user message 的 token 数过滤与cli_args.py中max_length的 help 文案 “Longer sequences are filtered out” 一致。三、第二步在areal/dataset/__init__.py中注册这一步是技能文档中的“Step 2”但需要按当前仓库的真实分发方式来落地。注册分两处3.1 加入VALID_DATASETS白名单VALID_DATASETS 是“受支持数据集”的权威列表既用于错误信息提示也供外部校验VALID_DATASETS [ gsm8k, clevr_count_70k, geometry3k, # ... 其他已注册数据集 name, ]3.2 在_get_custom_dataset中追加分发分支注意当前仓库的分发逻辑是按path子串匹配gsm8k in path而不是按显式name参数匹配且采用延迟导入命中才from .xxx import ...以保持模块加载开销最小。因此新分支应写成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, )从源码结构看几个细节值得留意name会出现在path中。分发条件是name in path所以数据集的 HF Hub 名或本地目录名需要包含该标识例如gsm8k对应 OpenAI/gsm8k 类路径。命名过短或过于通用的name如单个字母会造成误命中这也是 SWE 路径匹配 专门用正则(?:^|[/_\-.])swe(?:[/_\-.]|$)做词边界约束的原因——swe只在作为独立路径 tokenswe_data/、swe-bench时才命中answer_swe这类误报会落入通用回退。多模态数据集改传processor。参考 clevr_count_70k 与geometry3k分支视觉数据集的加载函数签名用processor而非tokenizer注册时同样要透传processorprocessor。**kwargs必须透传。_DatasetConfig.dataset_kwargs中的自定义参数最终经get_custom_dataset(**kwargs)进入加载函数见 rdataset 与 worker 调用链丢失**kwargs会导致dataset_kwargs配置静默失效。若数据集只服务某一种训练类型只注册对应分支即可如virl39k仅有rl分支。3.3 未注册路径的兜底行为如果忘记注册_get_custom_dataset会落到 load_from_disk 回退分支能save_to_disk的 HF Dataset 仍可通过本地路径直接加载否则抛出ValueError错误信息会列出VALID_DATASETS全表方便排查——这也是“未注册在__init__.py”被列为常见错误之一的原因。四、第三步可选为数据集增加专属配置字段如果新数据集需要特殊配置例如原始字段名映射、子集选择开关等技能文档建议在配置体系中扩展字段。当前仓库的配置载体是 _DatasetConfig / TrainDatasetConfig位于 areal/api/cli_args.pydataclass class TrainDatasetConfig(_DatasetConfig): # ... 继承 split / path / type / max_length / dataset_kwargs 等字段 name_specific_field: Optional[str] None从源码结构看这里有一个更轻量的替代方案_DatasetConfig.dataset_kwargscli_args.py#L3480-L3486本身就是“透传给get_custom_dataset的额外关键字参数字典”并且会经由RDataset或_get_custom_dataset原样传入加载函数**kwargs。对于一两个可选参数优先复用dataset_kwargs并在加载函数签名中以def get_name_sft_dataset(..., sub_set: str | None None)形式接收可以减少配置类改动只有当字段需要进入 YAML 文档、CLI 校验或跨 train/valid 共享语义时才在TrainDatasetConfig上声明正式字段。注意_DatasetConfig.__post_init__会校验sources与path/type互斥cli_args.py#L3505-L3512新增字段不要与path、type、sources语义冲突。五、第四步编写测试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仓库中已有可对照的测试命名与组织方式例如 tests/test_swe_sft_dataset.py、tests/test_swe_dataset.py 与 tests/test_mopd_dataset.py。建议在骨架之外再补两类断言loss_mask 语义对 SFT 数据集取一条样本断言sum(loss_mask) len(input_ids) - prompt 长度且 mask 前缀全 0防止“整条序列都算损失”这类隐蔽 bug对照 gsm8k.py#L14-L20 的实现max_length过滤生效传入一个极小的max_length断言len(dataset) 0或所有样本长度均小于阈值。六、必备字段规范SFT 与 RL 样本结构技能文档对“加载完成后每个样本必须长什么样”给出了明确契约这里完整保留SFT 数据集经模板处理后模板函数实际输出的是 token 级列input_ids/loss_mask。技能文档同时给出了 messages 形式的逻辑样本结构用于描述 SFT 对话本身{ messages: [ {role: user, content: ...}, {role: assistant, content: ...}, ] }RL 数据集{ messages: [ {role: user, content: ...}, ], answer: ground_truth_for_reward, # Optional metadata for reward function }两条硬性约束messages必须是含role与content键的字典列表OpenAI 风格RL 样本必须带可参与 reward 计算的answer字段可选的额外元数据列会随样本一路传到奖励函数。对照实现gsm8k.py 的 RL 分支 在 user content 中拼接了“请把最终答案放进\boxed{}”的指令注意其真实实现未保留answer列reward 侧通过解析生成文本中的 boxed 结果判分——如果你的数据集奖励函数直接比对 ground truth则应按技能文档模板保留answer列。七、参考实现速查表仓库内可直接对读的参考实现对应技能文档 Reference Implementations 一表数据集文件说明GSM8Kareal/dataset/gsm8k.py数学应用题最贴近模板的标准 SFT/RL 双实现Geometry3Kareal/dataset/geometry3k.py几何题多模态processor传参参考CLEVRareal/dataset/clevr_count_70k.py视觉计数任务HH-RLHFareal/dataset/hhrlhf.py有用性/无害性偏好数据rw/dpo类型专用加载器TORLareal/dataset/torl_data.py工具使用 RL其中 hhrlhf.py 还展示了另一种训练类型的扩展方式同一数据集可以按type提供多个变体加载器get_hhrlhf_rw_dataset产出chosen_ids/rejected_idsget_hhrlhf_dpo_dataset额外通过“逐 token 前缀比对”推导chosen_loss_mask/rejected_loss_mask只对分歧后的回答段计损失见 hhrlhf.py#L51-L68。新数据集若同时服务多种算法如 GRPO DPO可参照这种“一个模块、多个get_name_type_dataset函数”的组织方式。八、常见错误清单技能文档总结的五个高频坑结合当前仓库源码逐一说明后果返回List[Dict]而非 HFDatasetload_from_disk回退、dataset_kwargs透传、data-service 远端加载rdataset.py 中RDataset仅存元数据、由 worker 端再调_get_custom_dataset重建都假设返回值是 HF Dataset。用 Python 循环代替dataset.map()/dataset.filter()丢失向量化与分片能力大数据集加载性能显著退化。RL 数据集缺少messages字段prompt 构建无处取材rollout 阶段直接失败。message 格式错误必须是[{role: ..., content: ...}, ...]嵌套或字符串形式无法被对话模板渲染。未在areal/dataset/__init__.py注册路径不含已注册子串时落入load_from_disk回退HF Hub 名非本地目录加载必失败抛出带VALID_DATASETS列表的ValueError。九、落地检查清单完成开发后可按以下顺序自查areal/dataset/name.py存在get_name_sft_dataset/get_name_rl_dataset签名与 gsm8k.py 一致max_length走过滤而非截断VALID_DATASETS 含name_get_custom_dataset新增name in path and type sft / rl分支且透传**kwargs多模态改传processor数据集的path命名包含name标识避免与既有子串如swe冲突可选配置优先走dataset_kwargs必要时扩展TrainDatasetConfigtests/test_name_dataset.py通过覆盖列名、loss_mask语义与max_length过滤。走完以上清单新数据集即可像gsm8k、geometry3k一样通过 YAML 中的dataset.pathdataset.type配置被 AReaL 的训练与验证流程直接消费。【免费下载链接】AReaLThe RL Bridge for LLM-based Agent Applications. Made Simple Flexible.项目地址: https://gitcode.com/GitHub_Trending/are/AReaL创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表