TorchTitan 数据流水线实战指南:基于 Grain 的源、数据集、打包与分布式加载
【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan
本文系统讲解 TorchTitan 中统一基于 Grain,配套源码在 torchtitan/components/data/ 目录。
统一数据流水线的心智模型
TorchTitan 用一个基于 Grain 的流水线同时服务文本预训练、SFT 与图像训练,共分五层。先建立整体直觉:
1. Define a source (e.g. jsonl): class: SourceConfig output: RandomAccessDataSource | IterDataset 2. Define a dataset (filter/process applied to the source): class: SingleDatasetConfig does: pre-filter -> process -> post-filter output: MapDataset | IterDataset 3. Compose datasets (optional): class: e.g. FirstFitPackingConfig(dataset=DatasetMixConfig(...)) input: one or more child DatasetConfig values does: mix, concatenate, and/or pack output: MapDataset | IterDataset 4. Dataloader: runtime: GrainDataLoader config: GrainDataLoader.Config input: MapDataset | IterDataset does: convert to iterable if needed -> batch -> collate -> prefetch output: TrainerBatch 5. Trainer: input: TrainerBatch does: model forward and backward这五层分别对应 torchtitan/components/data/ 下的sources.py(源)、dataset.py(数据集与组合)、packing.py(打包)、loader.py(加载器)和collators.py(批处理),公共类型定义在 types.py。
值得注意的细节:SourceConfig与DatasetConfig在源码中都是协议(Protocol)而非基类,只要实现约定的build()方法即可参与组装;而SingleDatasetConfig、DatasetMixConfig等则是冻结数据类(frozen dataclass),以配置值(Config value)的形式被声明,由 GrainDataLoader 在初始化时统一构建为 Grain 数据集图。从源码结构看,这套设计的核心是“声明式配置 + 延迟构建”:你在配置中描述流水线形状,真正的构建发生在 loader.py 的GrainDataLoader.__init__中,它把 tokenizer、max_context_length、num_tokens_per_batch等运行时值打包进DatasetBuildContext,把seed/shuffle/repeat/dp_rank/dp_world_size打包进DatasetIterationPolicy,再自上而下构建整张数据集图。
文本预训练:从 JSONL 到打包好的 Token 批
本地 JSONL 源
JSONL 文件要求每个非空行恰好是一个 JSON 对象:
{"title": "First", "body": "The first document."} {"title": "Second", "body": "The second document."}最简的文本预训练配置:
from torchtitan.components.data import ( ConcatThenSplitPackingConfig, GrainDataLoader, IndexedJsonlSource, SingleDatasetConfig, ) from torchtitan.hf_datasets.text_datasets import TextProcessor def article_text(row): return row["title"] + "\n\n" + row["body"] books_ds = SingleDatasetConfig( source=IndexedJsonlSource.Config( patterns=( "/datasets/books/*.jsonl", ), ), processor=TextProcessor.Config(text_fn=article_text), post_filters=(lambda sample: sample is not None,), ) books_packed_ds = ConcatThenSplitPackingConfig(dataset=books_ds) config.dataloader = GrainDataLoader.Config( dataset=books_packed_ds, )这里IndexedJsonlSource(实现见 sources.py)提供的是基于字节偏移的随机访问:初始化时扫描每个匹配的 JSONL 文件,记录每个非空行的(path_id, byte_offset);__getitem__(index)直接seek到偏移处读取该行并json.loads,因此既不需要把整个语料加载进内存,也能支持dataset[index]式随机访问。有几点实现约束值得注意:
- 每个 glob pattern 必须至少匹配一个文件,否则抛出
FileNotFoundError; - 多个 pattern 解析出的路径若出现重复,会抛出
ValueError拒绝,防止同一文件被索引两次(sources.py); - 源码中的 TODO 也说明了当前限制:每个 rank 和 worker 启动时都会重新扫描全部 JSONL 文件,未来计划构建一个可被所有进程 mmap 的共享偏移索引。
TextProcessor(text_datasets.py)把text_fn返回的字符串编码为带 BOS/EOS 的 token 序列,然后前移一位切分为 next-token 对齐的TextSequence(input_ids, labels)——注意这个移位发生在数据集处理器里,而不是 trainer 中,TextSequence的注释明确强调“trainer does not do it”。短于 2 个 token 的样本返回None,由post_filters过滤掉。
ConcatThenSplitPackingConfig实现于 packing.py:先把文档拼接成连续 token 流,再切分成固定长度(num_tokens_per_batch)的行,并丢弃未填满的行(_packing_output_is_full)。当设置了max_num_documents(每行最多包含的文档片段数)时,会切换到文档感知(document-aware)实现,在行内保留文档边界、记录 remainder 状态以保证续训精确恢复(见 packing.py)。
Hugging Face 源:随机访问还是流式?
需要物化(materialize)数据集时用随机访问源:
from torchtitan.components.data import HuggingFaceRandomAccessSource source = HuggingFaceRandomAccessSource.Config( path="openai/gsm8k", name="main", split="train", )语料不应物化时用流式源:
from torchtitan.components.data import HuggingFaceStreamingSource source = HuggingFaceStreamingSource.Config( path="allenai/c4", name="en", split="train", )两类源都接受path、split、name、revision以及透传的load_dataset_kwargs,并在__post_init__中拒绝把split/name/revision/streaming重复放进 kwargs(sources.py)。它们的底层实现差异很大:
HuggingFaceRandomAccessSource以streaming=False加载一个datasets.Dataset,非Dataset类型会直接报错,随后只做len/__getitem__包装;HuggingFaceStreamingSource本身继承grain.IterDataset,以streaming=True加载IterableDataset,并用datasets.distributed.split_dataset_by_node在源级按 DP 坐标切分(dp_rank/dp_world_size来自DatasetIterationPolicy)。它还要求 HF 数据集实现state_dict()/load_state_dict(),否则拒绝使用——这是精确续训的前提。内部的_HuggingFaceCursorIterator把流式游标暴露给 Grain 的检查点递归:记录epoch与 HF 内部状态,repeat=True时在每个 epoch 递增epoch、必要时set_epoch重设 shuffle 后再从头迭代(sources.py)。
添加自定义源:预 token 化数据示例
当需要把“每条索引对应一篇已 token 化的文档”暴露为随机访问源时,可以自定义源,并复用既有的处理与打包配置。完整的可运行示例见 torchtitan/components/data/README.md,核心是把 memmap 的 token 数组与文档偏移数组包装成RandomAccessDataSource:
from dataclasses import dataclass import numpy as np from torchtitan.components.data import ( ConcatThenSplitPackingConfig, DatasetBuildContext, DatasetIterationPolicy, RandomAccessDataSource, SampleProcessor, SingleDatasetConfig, TextSequence, ) from torchtitan.config import Configurable class PretokenizedMemmapSource(Configurable, RandomAccessDataSource): @dataclass(kw_only=True, slots=True) class Config(Configurable.Config): tokens_path: str document_offsets_path: str def __init__( self, config: Config, *, dataset_iteration_policy: DatasetIterationPolicy, ): del dataset_iteration_policy self.tokens = np.memmap(config.tokens_path, dtype=np.uint32, mode="r") self.offsets = np.load(config.document_offsets_path) def __len__(self): return len(self.offsets) - 1 def __getitem__(self, index): start, end = self.offsets[index : index + 2] return np.asarray(self.tokens[start:end], dtype=np.int64) class TokensToTextSequence(SampleProcessor): @dataclass(kw_only=True, slots=True) class Config(SampleProcessor.Config): pass def __init__(self, config: Config, *, context: DatasetBuildContext): del config, context def __call__(self, token_ids, rng): del rng if len(token_ids) < 2: return None return TextSequence( input_ids=token_ids[:-1], labels=token_ids[1:], ) token_documents_ds = SingleDatasetConfig( source=PretokenizedMemmapSource.Config( tokens_path="tokens.bin", document_offsets_path="document_offsets.npy", ), processor=TokensToTextSequence.Config(), post_filters=(lambda sample: sample is not None,), ) packed_tokens_ds = ConcatThenSplitPackingConfig( dataset=token_documents_ds, )这个例子展示了三层复用:自定义源只负责“按索引返回原始行”;SampleProcessor(dataset.py)负责把原始行转成TextSequence,其中rng参数是 Grain 提供的确定性随机数生成器,可安全用于数据增强;打包与加载则完全复用通用组件。SingleDatasetConfig在构建时会执行pre_filters -> processor(random_map) -> post_filters -> (shuffle) -> DP shard -> (repeat)的固定顺序(dataset.py),这解释了为什么post_filters里常见的lambda sample: sample is not None能过滤掉处理器返回的None。
SFT:换处理器与打包策略,不换加载器
SFT 与预训练的差异只体现在处理器和打包策略上,加载器保持不变:
from torchtitan.components.data import ( FirstFitPackingConfig, GrainDataLoader, HuggingFaceRandomAccessSource, SingleDatasetConfig, ) from torchtitan.hf_datasets.text_datasets import ChatProcessor def gsm8k_messages(row): return [ {"role": "user", "content": row["question"]}, {"role": "assistant", "content": row["answer"]}, ] gsm8k_ds = SingleDatasetConfig( source=HuggingFaceRandomAccessSource.Config( path="openai/gsm8k", name="main", split="train", ), processor=ChatProcessor.Config(messages_fn=gsm8k_messages), post_filters=(lambda sample: sample is not None,), ) gsm8k_packed_ds = FirstFitPackingConfig(dataset=gsm8k_ds) config.dataloader = GrainDataLoader.Config( dataset=gsm8k_packed_ds, )无 renderer 时的单轮对话处理
不配置 renderer 时,ChatProcessor走 tokenizer 的 chat template 处理单轮[user, assistant]对话:先校验消息严格为两轮且 role 分别为user、assistant(否则ValueError),再用apply_chat_template渲染完整对话并追加 EOS,生成 next-token 的 input/label 对,最后把 prompt 部分的 label 置为IGNORE_INDEX(即 -100,见 torchtitan/components/loss.py 的IGNORE_INDEX),只对 assistant 回答计算损失。
prompt/response 边界的定位方式是:单独用add_generation_prompt=True渲染 prompt,并要求它的 token 序列恰好是完整渲染 token 序列的前缀,否则抛出ValueError(text_datasets.py)。这里选择“报错”而非“丢弃样本”是有意的:前缀不匹配意味着标签边界无法确定,且这类问题通常是模板系统性的,静默丢弃会让模型只在小部分数据上训练。超出max_context_length的样本则被整条丢弃(因为超长是个别样本问题)。ChatProcessor也会在首个样本时打印完整渲染文本便于人工核对。
使用 renderer 处理多轮对话
多轮对话需要显式选择模型对应的 renderer:
from renderers import Qwen3RendererConfig from torchtitan.components.renderer import RenderersLibraryConfig processor = ChatProcessor.Config( messages_fn=lambda row: row["messages"], renderer=RenderersLibraryConfig(renderers_config=Qwen3RendererConfig()), )RenderersLibraryConfig(renderer.py)基于renderers库,通过RendererTokenizerWrapper把 TorchTitan 已加载的 tokenizer 适配成 renderer 需要的接口(字符偏移、token 到 id 的查找、raw 编码),不会二次加载 tokenizer。它明确禁止auto与default类型的 renderer:前者依赖name_or_path精确匹配且可能落到不受支持的 DefaultRenderer,后者依赖 HF 的apply_chat_template,而 TorchTitan 的模板渲染缺少其特殊 token 变量,会产生静默不同的 token。
使用 renderer 时,一次渲染即返回 token 与 loss mask:mask 只监督模型生成的 token(含回合终止符),排除 prompt token 与模板脚手架。约束包括:
- 对话必须以 assistant 消息结尾,否则报错;
- 每个对话是一个样本,打包时在对话之间重置位置,而不是回合之间;
- 超过
max_context_length的样本整条丢弃; - 格式化与思维链(reasoning)保留策略由所选 renderer 决定,例如 Qwen3 会在最后一个 user 提问之前省略 assistant 的 reasoning,被省略的 token 不计损失;若想对每个回合的 reasoning 都训练,需要在源数据集里准备独立的对话前缀;
thinking_retention控制的是 renderer 的 rollout 桥接,不作用于此处使用的完整渲染。
混合与拼接数据集
加权混合:DatasetMixConfig 与 WeightedDataset
把权重放在每个数据集旁边:
from torchtitan.components.data import DatasetMixConfig, WeightedDataset # `books_ds` and `code_ds` are SingleDatasetConfig values. pretraining_mix_ds = DatasetMixConfig( datasets=( WeightedDataset(dataset=books_ds, weight=0.75), WeightedDataset(dataset=code_ds, weight=0.25), ), )示例中概率之和为 1.0,但实现上接受任意正相对权重并在内部归一化——weight=0.75配weight=0.25表示第一个数据集被抽中的频率是第二个的 3 倍。DatasetMixConfig(dataset.py)在构建时校验所有权重有限且为正,并给每个子数据集分配seed + index的偏移(插入或重排子数据集会重新播种其后的所有子数据集,从而改变数据顺序——断点续训要求代码与配置完全不变的原因之一)。
混合的语义取决于子数据集的粒度,这由 README 中的两个小节给出:
先混合再打包:权重按“文档”计数
packed_pretraining_ds = ConcatThenSplitPackingConfig( dataset=pretraining_mix_ds, )先打包再混合:权重按“定长行”计数
books_packed_ds = ConcatThenSplitPackingConfig(dataset=books_ds) code_packed_ds = ConcatThenSplitPackingConfig(dataset=code_ds) token_ratio_mix_ds = DatasetMixConfig( datasets=( WeightedDataset(dataset=books_packed_ds, weight=0.67), WeightedDataset(dataset=code_packed_ds, weight=0.33), ), )源码注释(dataset.py)补充了微妙差异:全 map 路径下MapDataset.filter会把被拒索引留作None,因此权重作用于尝试取样的索引而非被接受的样本;而一旦混入 iterable 子数据集,权重则作用于各子数据集实际发射的元素——混合TextSequence子项按文档计数,混合打包好的定长子项则按物理 token 计数。repeat=True时混合是无限的;repeat=False时混合在第一个耗尽的子数据集处停止,较大的子数据集不会被完整覆盖。从源码结构看,按观察到的文档数或监督 token 数自动调节权重的功能尚未实现(源码中有 TODO 注释),文档也明确说明:自定义混合可维护各数据集的移动平均并自行再平衡权重。
拼接:把有限数据集当作一个语料
拼接用于把多个有限数据集视为一个语料:每行出现一次,占比由各数据集大小决定;需要显式权重或流式数据集时改用混合:
from torchtitan.components.data import DatasetConcatConfig pretraining_corpus_ds = DatasetConcatConfig( datasets=(books_ds, code_ds, math_ds), )若希望某个有限数据集的每一行在一个 epoch 内出现多次,可以在拼接中重复该子项:
pretraining_corpus_ds = DatasetConcatConfig( datasets=(books_ds,) * 3 + (code_ds, math_ds), )此时每个books_ds行在合并后的有限索引空间中出现 3 次;当shuffle=True时,TorchTitan 会在 DP 分片之前对合并后的索引空间做全局 shuffle。实现上(dataset.py)DatasetConcatConfig要求全部子项为 map 风格,构建时先把每个子数据集以shuffle=False, repeat=False, dp_rank=0, dp_world_size=1的迭代策略构建,再MapDataset.concatenate合并,最后统一全局 shuffle、DP 分片、repeat。若需求是“相对采样频率”而非“精确的有限重复”,应改用混合:
pretraining_mix_ds = DatasetMixConfig( datasets=( WeightedDataset(dataset=books_ds, weight=0.6), WeightedDataset(dataset=code_ds, weight=0.2), WeightedDataset(dataset=math_ds, weight=0.2), ), )最后注意:GrainDataLoader.Config(repeat=True)重复的是整个拼接/混合后的数据集,并不会改变某个子数据集的相对贡献。
图像与多模态数据
图像训练复用同一套 source / dataset / loader / sharding / checkpoint 契约:处理器保留模态相关的样本字典,collator 生成模型专属的 batch。Qwen 多模态示例:
from torchtitan.components.data import ( GrainDataLoader, HuggingFaceStreamingSource, SingleDatasetConfig, ) from torchtitan.hf_datasets.multimodal.mm_collator import MultiModalCollator from torchtitan.hf_datasets.multimodal.mm_datasets import ( MMSamplePackingConfig, MultiModalProcessor, _process_cc12_wd_sample, ) mm_processor = MultiModalProcessor.Config( sample_processor=_process_cc12_wd_sample, ) mm_ds = SingleDatasetConfig( source=HuggingFaceStreamingSource.Config( path="pixparse/cc12m-wds", split="train", ), processor=mm_processor, post_filters=(lambda sample: sample is not None,), ) packed_mm_ds = MMSamplePackingConfig( dataset=mm_ds, num_packing_bins=8, ) config.dataloader = GrainDataLoader.Config( dataset=packed_mm_ds, collator=MultiModalCollator.Config( build_mrope_positions=True, patch_size=mm_processor.patch_size, temporal_patch_size=mm_processor.temporal_patch_size, spatial_merge_size=mm_processor.spatial_merge_size, ), streaming_shuffle_buffer_size=128, )几个关键语义(均有源码对应):
MultiModalProcessor(mm_datasets.py)适配 Grain 的 map 契约,内部调用sample_processor(如_process_cc12_wd_sample、_process_obelics_sample)。处理流程为:解码图像/视频字节 -> 缩放到patch_size * spatial_merge_size的倍数并归一化 -> 在文本中插入<|vision_start|><|image_pad|>...<|vision_end|>占位 token -> 编码文本;vision 占位 token 在 labels 中被置为IGNORE_INDEX(mm_datasets.py)。默认配置包括patch_size=16、temporal_patch_size=2、spatial_merge_size=2、min_pixels=65_536、max_pixels=16_777_216、max_patches=4096等。处理后的样本若超过max_context_length会被整条跳过。MMSamplePackingConfig把整个多模态文档装进定长行:先过滤超长样本,再用FirstFitPackIterDataset打包(input_ids/labels/positions为定长结构,pixel_values作为 meta features 保留),num_packing_bins是保持打开的候选打包行数量,而不是输入样本缓冲区——更大的值可以减少 padding 但会滞留更多媒体数据(mm_datasets.py)。MultiModalCollator(mm_collator.py)负责两件事:collate_images把图像/视频张量切块(patch)并 pad 到统一 patch 数,产出pixel_values与grid_thw(形状(num_images, 3),乘积给出每个条目的 patch 序列长度);collate_text拼接整条样本并只 pad token 批的尾部。当build_mrope_positions=True时,还会在 CPU 数据侧构建三维(时间/高/宽)的 MRoPE 位置 ID(_build_mrope_positions,返回(num_tokens, 3)),此时patch_order必须为"block"(raster 顺序会使 MRoPE 与 patch 序列失同步)。max_images_per_batch(默认 128)限制每个批的视觉条目数。- 自定义图像增强应放进
SampleProcessor。
Loader 策略:一次配置,全局生效
GrainDataLoader 完整配置
运行级行为在GrainDataLoader.Config中一次配好(loader.py):
import grain.python as grain config.dataloader = GrainDataLoader.Config( dataset=packed_pretraining_ds, seed=42, shuffle=True, repeat=True, streaming_shuffle_buffer_size=1_000, read_options=grain.ReadOptions( num_threads=16, prefetch_buffer_size=500, ), num_prefetch_batches=2, )各字段的默认值与作用(源码注释):
| 字段 | 默认值 | 说明 |
|---|---|---|
dataset | 必填 | 任意DatasetConfig(叶子、混合、拼接或打包后的数据集) |
collator | TextCollator.Config | 行到批的转换,多模态用MultiModalCollator.Config |
seed | 42 | 全局随机种子,贯穿 shuffle、random_map、打包 |
shuffle | True | 是否全局/流式 shuffle |
repeat | True | 是否无限重复;见下文分布式小节中的限制 |
streaming_shuffle_buffer_size | 1_000 | 每个 rank 保留的流式原始行数,用于近似 shuffle |
read_options | grain.ReadOptions() | MapDataset转IterDataset时的并发读参数 |
num_prefetch_batches | 2 | 预留给 trainer 的完整 collate 批数量 |
GrainDataLoader.__init__里还有一道防线:当dp_world_size > 1且repeat=False时直接抛ValueError,提示必须用repeat=True配合 trainer 控制的步数(原因见“分布式与检查点行为”一节)。加载器构建完数据集图后执行batch(collator.num_rows_per_batch(), drop_remainder=config.repeat, batch_fn=collator),再用ThreadPrefetchIterDataset预取num_prefetch_batches个完整批。TextCollator(collators.py)把若干TextSequence行合并进**预分配、页锁定(pin_memory)**的定长张量,padding 位置 label 为IGNORE_INDEX、padding_mask=True,并在 batch 字典中预先计算num_valid_tokens(参与损失的 token 数),避免 trainer 在关键路径上重复扫描。
读取带索引的数据集:MapDataset 与 IterDataset 的转换规则
随机访问源构成MapDataset(支持dataset[index]);当后续某个阶段需要顺序消费样本时,它会转成IterDataset。转换时机遵循以下规则(README 原文):
all children are MapDataset: DatasetMixConfig remains a MapDataset any child is IterDataset: DatasetMixConfig converts each MapDataset child packing: converts its child if needed and returns IterDataset GrainDataLoader: converts a MapDataset if no earlier stage didread_options控制每一次MapDataset到IterDataset的转换:
grain.ReadOptions( num_threads=16, # indexed samples read concurrently prefetch_buffer_size=500, # samples waiting for the consumer )每次转换都有自己独立的线程和缓冲区:全 map 的混合只转换一次;混有流的混合会为每个 map 子数据集各自转换一次。loader.py 中的 TODO 也提到当前只使用多线程而非多进程做 CPU 密集处理,未来计划在更早的边界引入共享 worker 池。
流式 shuffle 与就绪批
streaming_shuffle_buffer_size:保留用于近似 shuffle的原始行数(每个 rank)。缓冲区越大,混合越均匀,但内存占用越高。流式源在SingleDatasetConfig._build_iter_dataset里通过grain.experimental.WindowShuffleIterDataset实现(dataset.py),窗口大小即此值。num_prefetch_batches:允许等待 trainer 的完整 collate 批数量:
trainer computes batch 10 background thread prepares batches 11 and 12即 trainer 计算第 10 批时,后台线程已在准备第 11、12 批。
分布式与检查点行为
有效数据并行(effective DP)决定数据归属
只有有效 DP 坐标负责选取数据:
effective DP = data_parallel_replicate_degree * data_parallel_shard_degree different effective-DP ranks -> disjoint source rows TP/PP/CP peers -> same rows for their effective-DP coordinate也就是说,TP/PP/CP 的 peer rank 与它们的有效 DP 坐标共享相同数据行,而不同的有效 DP rank 之间数据行互不相交。这是保证前向/反向集合通信语义一致的关键。
数据所有权的判定时机
数据归属在批处理之前决定,不同源/组合的归属方式不同:
random-access source -> global shuffle -> contiguous balanced DP shard Hugging Face stream -> source-level DP shard DatasetMixConfig -> combines children already owned by this DP rank DatasetConcatConfig -> concatenates, globally shuffles, then DP-shards packing -> packs samples locally on each DP rank GrainDataLoader -> batches and collates that rank's samples对随机访问训练而言,每个 rank 拿到的是全局 shuffle 后索引空间的连续切片,而不是原始语料的连续区域。SingleDatasetConfig与DatasetConcatConfig中divmod分片逻辑(dataset.py)保证了dp_world_size不整除时余数行被均衡分配。当len(dataset) < dp_world_size时会直接报错,避免分片后出现空集。
有限数据与 hang 的陷阱
当有效 DP 大于 1 时,repeat=False会被拒绝(loader.py):因为各 rank 可能在不同的 step耗尽数据,导致训练集合通信挂起。正确做法是repeat=True,让 trainer 的步数控制训练何时停止。加载器还定义了DataloaderExhaustedError(loader.py),刻意继承Exception而非StopIteration,以避免 PEP 479 把生成器内部的StopIteration包装成RuntimeError导致程序崩溃。
state_dict 与断点续训
GrainDataLoader.state_dict()(loader.py)记录以下内容:
- 源游标与 shuffle/repeat 进度(含 HF 流式源的 epoch 与内部状态);
- mix 子数据集与打包缓冲区的状态(如 document-aware 打包的 remainder 与偏移);
- batching 与 prefetch 状态;
- 有效 DP 维度(
dp_world_size)。
恢复(load_state_dict)的硬性要求:代码、配置、源内容、tokenizer 与有效 DP 维度都必须与保存时一致。源码会显式校验state_dict["version"] == 1以及dp_world_size是否变化,并确保 checkpoint 中包含当前dp_rank_{i}的条目,缺失或版本不符都会抛错。这也解释了上文混合时“插入或重排子数据集会重新播种其后所有子数据集”为何影响续训一致性。配套的单元测试覆盖了这些行为,例如test_indexed_jsonl_random_access、test_hf_shuffled_repeat_advances_epoch、test_hf_resume_mid_second_epoch、test_loader_requires_repeat_with_data_parallelism、test_weighted_map_mix_keeps_weight_with_dataset、test_concat_shards_after_one_global_index_space、test_document_aware_concat_packing_restores_exactly等,见 tests/unit_tests/cpu/components/data/test_grain_data.py,可作为理解各阶段精确语义的补充证据。
小结
TorchTitan 的数据流水线把“源、处理、组合、加载”四层职责分离,全部以声明式 Config 表达,最终由GrainDataLoader统一构建并接入 trainer。文本预训练、SFT 与多模态训练共用同一套骨架,差异仅在SampleProcessor与打包/ collator 的选择上;分布式场景下“有效 DP 决定数据归属 + 全局 shuffle 后分片 + repeat=True”的约定,加上可递归恢复的 Grain 状态,为大规模多机训练提供了可复现与可断点续训的坚实基础。动手实践时,建议从本文的 JSONL 最小示例出发,逐步替换源、处理器与组合方式,并用仓库中的单元测试验证每一步的语义。
【免费下载链接】torchtitanA PyTorch native platform for training generative AI models项目地址: https://gitcode.com/GitHub_Trending/to/torchtitan
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考