torchtune 自定义数据集实战:用自定义 Message Transform 与 SFTDataset 构建端到端数据管线
【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune
当你的数据 schema 无法套用 torchtune 内置的数据集构建函数(builder)时,可以通过「自定义消息转换 +SFTDataset」两个组件组合出一条完整的自定义数据集管线:先把原始样本转换成 torchtune 标准的Message对话格式,再交给SFTDataset完成分词与训练样本准备。本文基于官方文档 custom_datasets.rst 的完整流程展开,并结合 torchtune/datasets/_sft.py 等源码,讲清每一步的参数细节、底层处理逻辑与在 recipe 配置中的接线方式,读完即可上手为任意 JSON/CSV/本地文件数据集编写可运行的自定义数据集代码。
一、整体流程:两级 Transform 的数据管线
torchtune 把所有微调数据统一抽象为「与模型的对话」:每条样本最终都会变成一组Message对象,再由模型专属的 tokenizer 转成 token。因此自定义数据集不需要你从零实现 Dataset 逻辑,只需完成两件事(这也是原文档给出的核心结论):
- 自定义 message transform:把原始样本(一行 JSON 数据)转换成
Message列表; - 用
SFTDataset包装该 transform:由它负责数据加载、分词、mask 与 label 生成。
对应到源码,数据在SFTDataset中的流转路径是:
__init__中调用load_dataset(source, **load_dataset_kwargs)加载原始数据(torchtune/datasets/_sft.py);__getitem__取出单行样本后交给内部的SFTTransform执行「message_transform → model_transform」两步处理(torchtune/datasets/_sft.py)。
也就是说,你写的那个小 builder 函数只是配置入口,真正干活的SFTDataset类已经实现了加载、校验、过滤、分词和 label 构建的完整闭环。
二、第一步:创建自定义 message transform
message transform 的职责是把你的原始样本字典转成 torchtune 的Message对象列表,并要求结果存放在"messages"键下。做法是继承 torchtune/modules/transforms/_transforms.py 中的Transform协议类(它本质上只要求实现一个__call__(sample) -> sample接口),把转换逻辑写进__call__方法。
原文档给出的完整示例如下(可直接照抄):
from typing import Any, Mapping from torchtune.data import Message from torchtune.modules.transforms import Transform class MyMessageTransform(Transform): def __call__(self, sample: Mapping[str, Any]) -> Mapping[str, Any]: return { "messages": [ Message(role="user", content=sample["input"], masked=True, eot=True), Message(role="assistant", content=sample["output"], masked=False, eot=True), ] }几个关键点需要结合Message类源码(torchtune/data/_messages.py)来理解:
| 参数 | 取值 | 含义 |
|---|---|---|
role | "system"/"user"/"assistant"/"ipython"/"tool" | 消息角色。system是系统提示词,user是输入 prompt,assistant是模型回答(也是实际计算 loss 的部分),ipython/tool是工具调用返回 |
content | str或list[dict] | 纯文本可直接传字符串,内部会自动包成[{"type": "text", "content": ...}];多模态内容需传{"type": "image", "content": torch.Tensor}与文本片段交替的列表 |
masked | bool,默认False | 是否对 loss 屏蔽。示例中 user 消息masked=True,即只在 assistant 回答上训练 |
eot | bool,默认True | 是否在该消息末尾追加 end-of-turn 特殊 token,表示对话轮次交接 |
示例中 user 消息masked=True、assistant 消息masked=False,等价于内置 transform 的masking_strategy="train_on_assistant"行为(见 torchtune/data/_messages.py 中mask_messages的实现)。如果你希望同时学习用户输入(例如纯文本续写场景),可把 user 的masked改为False,对应train_on_all策略。
此外要注意SFTDataset对消息序列有强制校验:SFTTransform在 tokenization 前会调用validate_messages(torchtune/data/_messages.py),以下情况会直接抛错:
- 消息少于 2 条(至少需要一轮 user-assistant 往返);
assistant消息出现在任何user/tool/ipython消息之前;- 连续两条
user消息; system消息不在第一条。
编写自定义 transform 时提前保证消息顺序合法,可以避免训练启动后才暴露问题。
三、第二步:用SFTDataset编写数据集 builder
把 transform 包进一个 builder 函数,这是 torchtune 推荐的组件组织方式。原文档示例如下,注意其注释中指明了文件位置约定(data/dataset.py):
# data/dataset.py from torchtune.datasets import SFTDataset from data.message_transform import MyMessageTransform def custom_dataset(tokenizer, **load_dataset_kwargs) -> SFTDataset: return SFTDataset( source="json", data_files="data/my_data.json", split="train", message_transform=MyMessageTransform(), model_transform=tokenizer, **load_dataset_kwargs, )结合 torchtune/datasets/_sft.py 的__init__签名,SFTDataset的参数说明如下:
| 参数 | 说明 |
|---|---|
source | 必填。Hugging Face Hub 上的数据集仓库名;本地文件则写文件类型("json"、"csv"、"text"等),并配合data_files传入路径 |
message_transform | 必填。第二步之外的核心:调用你的自定义 transform,输出存入"messages"键的Message列表 |
model_transform | 必填。模型专属预处理,文本任务直接传 tokenizer(ModelTokenizer实例即可),要求返回至少包含"tokens"与"mask"两个键的字典 |
filter_fn | 可选。在预处理之前过滤数据集,语义同 Hugging Face datasets 的Dataset.filter |
filter_kwargs | 可选。传给filter_fn的额外参数 |
**load_dataset_kwargs | 透传给datasets.load_dataset的任意参数(如data_files、split、streaming等) |
示例中split="train"与data_files属于透传给load_dataset的参数。builder 首参tokenizer之所以存在,是因为默认 recipe 实例化数据集时会把配置中独立定义的 tokenizer 自动注入进来——这一点在下一节详述。
SFTDataset 内部如何生成训练样本
理解SFTTransform.__call__(torchtune/datasets/_sft.py)有助于排查数据问题,它做了三件事:
- 先执行
message_transform,若输出含"messages"则调用validate_messages校验对话合法性; - 再执行
model_transform(即 tokenizer)完成分词,并检查返回值必须包含"tokens"和"mask"两个键,否则抛出带键列表的ValueError; - 由 mask 生成
labels:mask为True的位置填入CROSS_ENTROPY_IGNORE_IDX(忽略项),否则保留 token,整体右移一位并对齐 logits,末尾补一个忽略 token。
这正是masked=True的Message最终不参与 loss 计算的底层机制——你只需要在 transform 里正确设置masked标志,损失屏蔽会由框架自动完成。
相关测试用例位于 tests/torchtune/datasets/test_sft_dataset.py,可作为验证行为是否符合预期的参考。
四、第三步:在 recipe 配置中接入自定义数据集
torchtune 的组件系统通过_component_字段按「相对于tune run启动目录」的导入路径实例化组件。原文档给出的配置写法:
dataset: _component_: data.dataset.custom_dataset即假设你的工程布局为:
my_project/ ├── data/ │ ├── dataset.py # custom_dataset builder │ ├── message_transform.py # MyMessageTransform │ └── my_data.json └── config/ └── custom_config.yaml有两条源码级规则必须遵守,否则会出现导入失败或参数错位的隐蔽错误:
- builder 的第一个位置参数必须是 tokenizer/model transform。从 recipe 源码可以看到,数据集是带 tokenizer 注入的:
lora_finetune_single_device的_setup_data中执行config.instantiate(cfg_dataset, self._tokenizer)(recipes/lora_finetune_single_device.py)。因此不要在配置里再写 tokenizer 字段,配置中单独定义的tokenizer组件会被自动传入 builder 首参。这也是 custom_components.rst 中官方 Note 强调的约定。 - 组件路径相对于启动目录解析。如果组件导入失败或找不到,可参考官方建议调整
PYTHONPATH使项目目录可导入,例如:
PYTHONPATH=${pwd}:PYTHONPATH tune run lora_finetune_single_device --config config/custom_config.yaml一个可直接运行的配置示例
在 custom_components.rst 的完整示例中,自定义数据集 builder 还会暴露自己的可配置参数(如packed: True),并展示与默认 recipe 组合的用法。综合原文档与源码,一份典型的完整 YAML 会是这样(数据集部分照抄原文档,其余为 recipe 常用字段示意):
# config/custom_config.yaml dataset: _component_: data.dataset.custom_dataset # 如需覆盖/追加 load_dataset 参数,可在这里直接写,例如: # max_seq_len 之类与你的数据相关的字段则写进 builder 显式参数 tokenizer: _component_: torchtune.models.llama3_2.tokenizer batch_size: 4 epochs: 1其中**load_dataset_kwargs的设计意义在于:配置中dataset字段下除 builder 显式参数外,还可以透传load_dataset所需参数,无需修改 Python 代码。
五、实用建议与常见坑
- 优先复用内置 transform:如果你的数据只是列名不同(如
prompt/response而非input/output),内置的InputOutputToMessages支持column_map重映射和masking_strategy/new_system_prompt配置(torchtune/data/_messages.py),不必手写 transform。只有内置类都无法配置时才需要自定义。 - 多数据集拼接:默认 recipe 支持
dataset写成 ListConfig,内部会逐个instantiate后包进ConcatDataset(recipes/lora_finetune_single_device.py),因此你的自定义 builder 也可以以列表形式与内置 builder 混用。 - 分词后的键契约:
model_transform必须返回"tokens"与"mask",缺失时SFTTransform会抛错(torchtune/datasets/_sft.py);纯文本场景直接把 tokenizer 作为model_transform传入即可满足该契约。 - 消息校验前置:
validate_messages在每次__getitem__时执行,坏数据会以ValueError形式暴露具体 index,便于定位是哪条样本的消息顺序不合法(torchtune/data/_messages.py)。
六、延伸阅读
- 消息对象构造与更多内置 transform 的详细说明:docs/source/basics/message_transforms.rst
- 自定义组件注册、builder 模式与
tune run/tune cp工作流:docs/source/basics/custom_components.rst - 自定义数据集完整实现源码:torchtune/datasets/_sft.py
Message类定义与消息校验逻辑:torchtune/data/_messages.pySFTDataset测试用例:tests/torchtune/datasets/test_sft_dataset.py
【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考