TorchTitan 自定义模型接入实战:基于 TrainSpec 协议从零扩展新模型架构
【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址: https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs
TorchTitan 是 PyTorch 官方的分布式 LLM 预训练框架,通过可组合的 4D 并行(FSDP2、TP、PP、CP)支撑从 8B 到 405B+ 规模的模型训练。本文以本仓库 custom-models.md 为核心指南,完整讲解如何遵循 TorchTitan 既有的协议模式(BaseModelArgs、ModelProtocol、TrainSpec)向框架注册并训练一个全新模型架构,覆盖参数定义、模型实现、并行化编排、注册接入、HuggingFace 权重互转与数值验证的完整闭环。
读完本文你将掌握:如何在不改动训练框架核心代码的前提下,用「单设备语义 + 外部并行化」的方式将任意 Transformer 变体(如自研注意力、MoE 层、新归一化结构)接入 TorchTitan,并跑通从 debug 配置到 8 GPU 真实训练的完整流程。
TorchTitan 的模型接入设计哲学
在动手写代码前,先理解 TorchTitan 的设计约束,这决定了所有实现细节。从本仓库 SKILL.md 可知,TorchTitan 是 PyTorch 原生的分布式预训练平台,核心卖点是「可组合的 4D 并行」(FSDP2、TP、PP、CP)以及 Float8、torch.compile、分布式 checkpoint(DCP)等能力。
为了让这些并行能力对任何模型都生效,TorchTitan 要求遵循四条指导原则(原文档 "Guiding Principles"):
- 可读性优先于灵活性(Readability over flexibility):不要过度抽象,模型代码保持直白;
- 模型改动最小化(Minimal model changes):并行性全部由外部注入,模型本体不感知 TP/FSDP/PP;
- 代码库保持简洁(Clean, minimal codebase):尽可能复用现有组件(优化器、学习率调度器、dataloader、tokenizer、loss);
- 单设备语义(Single-device semantics):模型代码必须在单 GPU 上就能直接跑通。
这套哲学与 fsdp.md 中介绍的 FSDP2 思路一脉相承:FSDP2 用fully_shard直接在模块上施加分片,auto_wrap_policy、FlatParameter等封装都被移除,模型代码与并行代码天然解耦。你的自定义模型只需要实现最朴素的nn.Module前向逻辑,其余交给parallelize_fn统一编排。
目录结构:一个模型的完整形态
原文档给出的标准目录骨架如下,我们逐一说明每个文件的职责:
torchtitan/models/your_model/ ├── model/ │ ├── __init__.py │ ├── args.py # 模型参数(继承 BaseModelArgs) │ ├── model.py # 模型定义(继承 ModelProtocol) │ └── state_dict_adapter.py # HF 权重互转(可选) ├── infra/ │ ├── __init__.py │ ├── parallelize.py # TP、FSDP、compile 的编排入口 │ └── pipeline.py # PP 编排(可选) ├── train_configs/ │ ├── debug_model.toml # 小规模 debug 配置 │ └── your_model_XB.toml # 正式训练配置 ├── __init__.py # TrainSpec 注册 └── README.md其中model/与infra/的分离正是「单设备语义」原则的体现:模型文件只关心计算图本身,并行化逻辑收敛在infra/下,方便与训练主循环train.py对接。
Step 1:定义模型参数(model/args.py)
所有模型参数类必须继承BaseModelArgs(位于torchtitan.protocols.model),它提供了与训练框架交互的统一接口。原文档的完整示例如下:
# model/args.py from torchtitan.protocols.model import BaseModelArgs from dataclasses import dataclass @dataclass class YourModelArgs(BaseModelArgs): dim: int = 4096 n_layers: int = 32 n_heads: int = 32 vocab_size: int = 128256 def get_nparams_and_flops(self, seq_len: int) -> tuple[int, int]: """Return (num_params, flops_per_token) for throughput calculation.""" nparams = self.vocab_size * self.dim + ... # Calculate params flops = 6 * nparams # Approximate: 6 * params for forward+backward return nparams, flops def update_from_config(self, job_config) -> "YourModelArgs": """Update args from training config.""" # Override specific args from job_config if needed return self两个必须实现的方法承担不同职责:
get_nparams_and_flops(seq_len):返回(参数量, 每 token 的 FLOPs),供框架计算吞吐率(TPS)与 MFU(模型浮点利用率)。其中flops = 6 * nparams是业界对「前向 + 反向」的标准近似(一次前向约 2 次乘加、反向约 4 次,合计约 6 倍参数量)。如果你的模型包含 MoE 或稀疏结构,需要按实际计算路径修正此值,否则性能指标会失真。update_from_config(job_config):允许从 TOML 训练配置中覆盖部分超参(例如通过[model] flavor选择 8B 还是 70B 后,把dim、n_layers按 flavor 映射)。默认返回self即可。
值得补充的是:vocab_size = 128256这类取值通常来自实际 tokenizer(如 Llama 3 词表),建议与 tokenizer 保持一致;若模型采用 padding 词表,也应在此处体现。
Step 2:定义模型本体(model/model.py)
模型类必须继承ModelProtocol(同样位于torchtitan.protocols.model)。原文档给出的是一个极简 Transformer 骨架:
# model/model.py import torch.nn as nn from torchtitan.protocols.model import ModelProtocol from .args import YourModelArgs class YourModel(ModelProtocol): def __init__(self, args: YourModelArgs): super().__init__() self.args = args self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim) self.layers = nn.ModuleDict({ str(i): TransformerBlock(args) for i in range(args.n_layers) }) self.norm = RMSNorm(args.dim) self.output = nn.Linear(args.dim, args.vocab_size, bias=False) def forward(self, tokens: torch.Tensor) -> torch.Tensor: h = self.tok_embeddings(tokens) for layer in self.layers.values(): h = layer(h) h = self.norm(h) return self.output(h) def init_weights(self): """Initialize weights recursively.""" for module in self.modules(): if hasattr(module, 'init_weights') and module is not self: module.init_weights() elif isinstance(module, nn.Linear): nn.init.normal_(module.weight, std=0.02)这里有几个关键实现细节(原文档 "Important guidelines"),直接影响后续并行化的正确性:
- 单设备代码:
forward只接收普通的torch.Tensor,不处理任何分布式逻辑。TP 的切分、FSDP 的分片都发生在parallelize_fn中。 - 用
nn.ModuleDict而非nn.ModuleList组织 layers:这是与 Pipeline Parallelism 兼容的关键。PP 需要把层按 stage 切分并删除不属于本 stage 的层,ModuleDict以字符串键索引,删除后其余子模块的 fully qualified name(FQN)保持不变,从而保证 state dict 键名稳定;而ModuleList删除中间元素会导致 FQN 编号漂移。 - 输入/输出层设为可选:
tok_embeddings与output在 PP 场景下只应存在于第一个/最后一个 stage。实现时可通过self.tok_embeddings = ... if not args.ignore_input_output else nn.Identity()之类的开关控制,确保 PP 切分后各 stage 的模块引用完整。 - 递归定义
init_weights():因为 FSDP2 采用 meta-device 初始化(先with torch.device("meta")建模型,再fully_shard分片,最后to_empty(device="cuda")后统一初始化,详见 fsdp.md),权重初始化必须发生在「分片完成之后」。你的模型需要提供一个可递归下发的init_weights(),让nn.Linear、nn.Embedding等子模块各自完成初始化。
Step 3:编写并行化函数(infra/parallelize.py)
这是自定义模型接入的"灵魂":并行化必须以固定顺序施加。原文档给出的顺序是TP → AC → compile → FSDP:
# infra/parallelize.py from torch.distributed._composable.fsdp import fully_shard from torch.distributed.tensor.parallel import parallelize_module def parallelize_your_model( model: YourModel, world_mesh: DeviceMesh, parallel_dims: ParallelDims, job_config: JobConfig, ): # Apply in this order: TP -> AC -> compile -> FSDP # 1. Tensor Parallelism if parallel_dims.tp_enabled: apply_tp(model, world_mesh["tp"], job_config) # 2. Activation Checkpointing if job_config.activation_checkpoint.mode == "full": apply_ac(model, job_config) # 3. torch.compile if job_config.compile.enable: model = torch.compile(model) # 4. FSDP if parallel_dims.dp_enabled: apply_fsdp(model, world_mesh["dp"], job_config) return model各阶段的作用与依据如下:
- TP 先行:先做张量并行切分(如
parallelize_module(model, world_mesh["tp"], {...})将nn.Linear的列/行切到 TP 维度的各 rank),后续 FSDP 的分片粒度基于已切分的参数。 - Activation Checkpointing:
mode可选"full"或"selective"(selective_ac_option = "op"表示按算子粒度选择性重算,见 SKILL.md 的配置示例),在编译之前施加以保证计算图可被完整捕获。 - torch.compile:必须放在 AC 之后、FSDP 之前——FSDP2 对
DTensor的处理与torch.compile有配合要求,且 compile 需要看到的是已注入 AC 的模块。从 SKILL.md 的 Float8 工作流可以看到,--compile.enable同时是 Float8 训练的前提(用于融合 float8 的 scale/cast 内核)。 - FSDP 最后:用
fully_shard对每个 TransformerBlock 分别分片,再对整体模型分片(meta-device 初始化 + 分片 +to_empty的完整流程见 fsdp.md 的 "Meta-Device Initialization" 一节)。
Step 4:创建并注册 TrainSpec(__init__.py)
TrainSpec是 TorchTitan 连接「模型」与「训练框架」的契约对象。原文档的核心代码:
# __init__.py from torchtitan.protocols.train_spec import TrainSpec, register_train_spec from .model.model import YourModel from .model.args import YourModelArgs from .infra.parallelize import parallelize_your_model MODEL_CONFIGS = { "8B": YourModelArgs(dim=4096, n_layers=32, n_heads=32), "70B": YourModelArgs(dim=8192, n_layers=80, n_heads=64), } def get_train_spec(flavor: str) -> TrainSpec: return TrainSpec( model_cls=YourModel, model_args=MODEL_CONFIGS[flavor], parallelize_fn=parallelize_your_model, pipeline_fn=None, # Or your_pipeline_fn for PP build_optimizer_fn=build_optimizer, # Reuse existing build_lr_scheduler_fn=build_lr_scheduler, # Reuse existing build_dataloader_fn=build_dataloader, # Reuse existing build_tokenizer_fn=build_tokenizer, # Reuse existing build_loss_fn=build_loss, # Reuse existing state_dict_adapter=None, # Or YourStateDictAdapter ) # Register so train.py can find it register_train_spec("your_model", get_train_spec)TrainSpec各字段含义:
| 字段 | 作用 |
|---|---|
model_cls | 模型类(继承ModelProtocol),由框架实例化 |
model_args | 按 flavor 选定的参数对象 |
parallelize_fn | 上一步写的并行化入口,签名必须与ParallelizeFn一致 |
pipeline_fn | PP 编排函数(可选);不启用 PP 时传None |
build_optimizer_fn/build_lr_scheduler_fn/build_dataloader_fn/build_tokenizer_fn/build_loss_fn | 尽量复用框架现有实现(这正是"代码库保持简洁"原则的落地) |
state_dict_adapter | 权重互转适配器(可选) |
register_train_spec("your_model", get_train_spec)把模型名注册进全局注册表,train.py即可通过[model] name = "your_model"找到它。
Step 5:State Dict Adapter(可选,HuggingFace 互转)
如果希望与 HuggingFace 生态互换 checkpoint(例如从 HF 加载预训练权重继续训练,或训练完成后导出给 HF 生态推理/微调),需要实现BaseStateDictAdapter:
# model/state_dict_adapter.py from torchtitan.protocols.state_dict_adapter import BaseStateDictAdapter class YourStateDictAdapter(BaseStateDictAdapter): def to_hf(self, state_dict: dict) -> dict: """Convert torchtitan state dict to HF format.""" hf_state_dict = {} for key, value in state_dict.items(): hf_key = self._convert_key_to_hf(key) hf_state_dict[hf_key] = value return hf_state_dict def from_hf(self, state_dict: dict) -> dict: """Convert HF state dict to torchtitan format.""" tt_state_dict = {} for key, value in state_dict.items(): tt_key = self._convert_key_from_hf(key) tt_state_dict[tt_key] = value return tt_state_dict两个方向的键名映射(如tok_embeddings.weight↔model.embed_tokens.weight)分别由_convert_key_to_hf/_convert_key_from_hf实现。把适配器实例传入TrainSpec(state_dict_adapter=YourStateDictAdapter(...))后,训练脚本即可支持 checkpoint 层面的 HF 互操作。更完整的 checkpoint 能力(DCP 分片结构、last_save_in_hf/initial_load_in_hf直接读写、离线转换脚本)可参考同目录的 checkpoint.md。
Step 6:编写训练配置(train_configs/your_model_8b.toml)
配置采用 TOML 格式,通过[job]、[model]、[optimizer]、[training]、[parallelism]等分区组织。原文档的完整示例:
# train_configs/your_model_8b.toml [job] dump_folder = "./outputs" description = "Your Model 8B training" [model] name = "your_model" flavor = "8B" [optimizer] name = "AdamW" lr = 3e-4 [training] local_batch_size = 2 seq_len = 8192 steps = 1000 dataset = "c4" [parallelism] data_parallel_shard_degree = -1 tensor_parallel_degree = 1对照 SKILL.md 的 Llama 3.1 8B 配置,可以补充以下常用分区,使配置更贴近实战:
[lr_scheduler] warmup_steps = 200 [training] max_norm = 1.0 # 梯度裁剪阈值 [activation_checkpoint] mode = "selective" # 或 "full" selective_ac_option = "op" [checkpoint] enable = true folder = "checkpoint" interval = 500关键参数说明:
[model] name = "your_model"必须与register_train_spec的第一个参数一致;flavor = "8B"对应MODEL_CONFIGS中的键;data_parallel_shard_degree = -1表示 FSDP 分片维度自动使用全部可用 GPU(-1即 auto);若想启用 HSDP(跨组复制 + 组内分片),可用data_parallel_replicate_degree,语义详见 fsdp.md 的 HSDP 一节;tensor_parallel_degree = 1表示单节点先不启用 TP,待模型变大(如 70B)再提升到 8(TP 在节点内)。
Step 7:注册到模型注册表
最后一步是把新模型登记进全局注册表torchtitan/models/__init__.py:
from .your_model import get_train_spec as get_your_model_train_spec MODEL_REGISTRY["your_model"] = get_your_model_train_spec完成这 7 步后,即可像内置模型一样启动训练。本仓库 SKILL.md 给出了两种启动方式:
# 方式一:通过 run_train.sh(读取 CONFIG_FILE 环境变量) CONFIG_FILE="./your_model_8b.toml" ./run_train.sh # 方式二:显式使用 torchrun torchrun --nproc_per_node=8 \ -m torchtitan.train \ --job.config_file ./your_model_8b.toml训练日志(TensorBoard)默认输出到dump_folder/tb/,可用tensorboard --logdir ./outputs/tb监控。多节点(SLURM)场景下,用srun torchrun --nnodes=N --nproc_per_node=8 ...并配置--rdzv_backend=c10d --rdzv_endpoint=$MASTER_ADDR:$MASTER_PORT即可扩展。
测试与验证:三关缺一不可
原文档给出了接入新模型必须通过的三类测试,这是保证训练正确性与性能的关键环节:
1. 数值一致性测试(Numerics Test)
将同一份 checkpoint 分别加载进 TorchTitan 实现与 HuggingFace 实现,对比相同输入下的输出:
def test_numerics(): # Load same checkpoint into both implementations tt_model = YourModel(args).load_checkpoint(...) hf_model = HFYourModel.from_pretrained(...) # Compare outputs input_ids = torch.randint(0, vocab_size, (1, 128)) tt_output = tt_model(input_ids) hf_output = hf_model(input_ids).logits torch.testing.assert_close(tt_output, hf_output, atol=1e-4, rtol=1e-4)注意:由于并行切分、编译与数值路径的差异,推荐在单 GPU、关闭并行的条件下做此对比,容差atol/rtol取1e-4级别的经验值。
2. 损失收敛测试(Loss Convergence)
与已验证的基线模型对比损失曲线,确保新架构的收敛行为符合预期,避免出现初始化或前向逻辑的隐性错误。训练若干步后损失应平滑下降,且与同规模基线模型的数量级一致。
3. 性能基准(Performance Benchmark)
在benchmarks/目录下补充基准配置,记录不同并行组合下的 TPS/GPU 与 MFU。可参考本仓库 SKILL.md 中 H100 上的内置模型基线(如 Llama 8B 纯 FSDP 约 5,762 TPS/GPU,加 compile 与 Float8 后约 8,532 TPS/GPU)作为同环境下的对照标尺,评估你的模型并行化编排是否达到预期。
进阶要点与常见坑
(1)并行化顺序不可颠倒。TP → AC → compile → FSDP 的顺序由各机制对计算图与张量布局的依赖决定。若把 FSDP 提前,torch.compile可能无法正确处理已分片的DTensor;若把 AC 放在 compile 之后,重算算子将无法被编译图捕获。
(2)PP 场景务必先造 seed checkpoint。启用 Pipeline Parallelism 前,需要用单卡(所有 parallel degree 均为 1)生成 seed checkpoint,保证各 stage 初始化一致。具体命令模板见 checkpoint.md 的 "Creating Seed Checkpoints" 一节。
(3)Float8 兼容性。若模型层数很多、GEMM 足够大(经验上 K、N 均大于 4096),可叠加 Float8 训练:[model] converters = ["quantize.linear.float8"]并配合--compile.enable,同时用filter_fqns排除收益小的层(如 output 投影),详见 float8.md。
(4)FQN 稳定性。层容器务必使用nn.ModuleDict并用字符串键索引;任何对层列表的增删都必须保持其余键的 FQN 不变,否则 PP 切分与 DCP checkpoint 加载会因键名漂移而失败。
总结
TorchTitan 的自定义模型接入本质上是「实现三个协议(BaseModelArgs/ModelProtocol/TrainSpec)+ 编排一个并行化函数」的过程。得益于「单设备语义 + 外部并行化」的设计,你完全可以把注意力放在模型架构本身,而 FSDP2、TP、PP、CP、Float8、DCP 等分布式能力由框架统一提供。接入后务必依次通过数值一致性、损失收敛与性能基准三关,再逐步把并行维度从纯 FSDP 扩展到 2D/3D/4D。
本仓库中与本文配套的可继续阅读材料:SKILL.md(快速开始与 4D 并行工作流)、fsdp.md(FSDP2 与 meta-device 初始化)、checkpoint.md(DCP 与 HF 互转)、float8.md(Float8 训练配置)。
【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址: https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考