尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
TorchTitan 自定义模型接入实战:基于 TrainSpec 协议从零扩展新模型架构
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-SKILLsTorchTitan 是 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、分布式 checkpointDCP等能力。为了让这些并行能力对任何模型都生效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 forwardbackward 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, biasFalse) 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, std0.02)这里有几个关键实现细节原文档 Important guidelines直接影响后续并行化的正确性单设备代码forward只接收普通的torch.Tensor不处理任何分布式逻辑。TP 的切分、FSDP 的分片都发生在parallelize_fn中。用nn.ModuleDict而非nn.ModuleList组织 layers这是与 Pipeline Parallelism 兼容的关键。PP 需要把层按 stage 切分并删除不属于本 stage 的层ModuleDict以字符串键索引删除后其余子模块的 fully qualified nameFQN保持不变从而保证 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(devicecuda)后统一初始化详见 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 Checkpointingmode可选full或selectiveselective_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__.pyTrainSpec是 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(dim4096, n_layers32, n_heads32), 70B: YourModelArgs(dim8192, n_layers80, n_heads64), } def get_train_spec(flavor: str) - TrainSpec: return TrainSpec( model_clsYourModel, model_argsMODEL_CONFIGS[flavor], parallelize_fnparallelize_your_model, pipeline_fnNone, # Or your_pipeline_fn for PP build_optimizer_fnbuild_optimizer, # Reuse existing build_lr_scheduler_fnbuild_lr_scheduler, # Reuse existing build_dataloader_fnbuild_dataloader, # Reuse existing build_tokenizer_fnbuild_tokenizer, # Reuse existing build_loss_fnbuild_loss, # Reuse existing state_dict_adapterNone, # 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_fnPP 编排函数可选不启用 PP 时传Nonebuild_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 5State 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_adapterYourStateDictAdapter(...))后训练脚本即可支持 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再提升到 8TP 在节点内。Step 7注册到模型注册表最后一步是把新模型登记进全局注册表torchtitan/models/__init__.pyfrom .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_node8 \ -m torchtitan.train \ --job.config_file ./your_model_8b.toml训练日志TensorBoard默认输出到dump_folder/tb/可用tensorboard --logdir ./outputs/tb监控。多节点SLURM场景下用srun torchrun --nnodesN --nproc_per_node8 ...并配置--rdzv_backendc10d --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, atol1e-4, rtol1e-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 之后重算算子将无法被编译图捕获。2PP 场景务必先造 seed checkpoint。启用 Pipeline Parallelism 前需要用单卡所有 parallel degree 均为 1生成 seed checkpoint保证各 stage 初始化一致。具体命令模板见 checkpoint.md 的 Creating Seed Checkpoints 一节。3Float8 兼容性。若模型层数很多、GEMM 足够大经验上 K、N 均大于 4096可叠加 Float8 训练[model] converters [quantize.linear.float8]并配合--compile.enable同时用filter_fqns排除收益小的层如 output 投影详见 float8.md。4FQN 稳定性。层容器务必使用nn.ModuleDict并用字符串键索引任何对层列表的增删都必须保持其余键的 FQN 不变否则 PP 切分与 DCP checkpoint 加载会因键名漂移而失败。总结TorchTitan 的自定义模型接入本质上是「实现三个协议BaseModelArgs/ModelProtocol/TrainSpec 编排一个并行化函数」的过程。得益于「单设备语义 外部并行化」的设计你完全可以把注意力放在模型架构本身而 FSDP2、TP、PP、CP、Float8、DCP 等分布式能力由框架统一提供。接入后务必依次通过数值一致性、损失收敛与性能基准三关再逐步把并行维度从纯 FSDP 扩展到 2D/3D/4D。本仓库中与本文配套的可继续阅读材料SKILL.md快速开始与 4D 并行工作流、fsdp.mdFSDP2 与 meta-device 初始化、checkpoint.mdDCP 与 HF 互转、float8.mdFloat8 训练配置。【免费下载链接】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),仅供参考
RELATED

相关推荐

3个维度拆解地图导航地图实战项目选型

3个维度拆解地图导航地图实战项目选型

3个维度拆解地图导航地图实战项目选型 官方文档动辄几百页,翻到第三章就头晕,根本抓不住重点。很多团队做 实战项目 时,卡在技术选型的泥潭里,到底是用WebGL还是原生Canvas,是用Leaflet还是Mapbox,往往决定了一个导航应用是…

📅 2026/9/23 2:26:33
立项书工程化指南:从模糊共识到可执行技术契约

立项书工程化指南:从模糊共识到可执行技术契约

简介:本资源是一份标准化、可直接套用的软件项目立项书模板文档,面向软件开发项目经理、需求分析师、项目申报人员及高校计算机相关专业学生,用于规范启动阶段的项目规划与干系人对齐。文档严格遵循立项书核心结构,涵盖项目名称与…

📅 2026/9/23 2:26:33
基于Java+MySQL的微信小程序网上花店毕设源码拆解与运行指南

基于Java+MySQL的微信小程序网上花店毕设源码拆解与运行指南

简介:基于微信小程序的网上花店系统源码,是一套面向计算机专业毕业设计、课程设计场景的完整前后端项目,也适合新手学习小程序与Java后端的数据交互。项目采用Java开发、搭配MySQL 5.7数据库,前端包含微信小程序与后台管理端。压缩…

📅 2026/9/23 2:21:33
MORE NEWS

更多资讯

📰

BP神经网络+Adaboost:时间序列预测的集成提升实践

做时间序列预测的人,多数都会被同一个问题反复缠住:单模型的精度上不去,怎么调都差那么一点。这个基于BP神经网络的Adaboost算法的时间序列预测项目,本质是把"一个BP网络"升级成"一堆BP网络投票决策"&#xf…

📰

HTML列表表格表单实战:语义化与移动端适配

这节内容我从实际开发的角度聊聊HTML里最容易忽略、但也最见功力的三个组件:列表、表格、表单。很多人学HTML时觉得这些标签简单——无非就是ul里放li、table里放tr、form里放input——但真到了做项目的时候,导航菜单怎么搭才语义清晰,课程表…

📰

用WebGPU在浏览器跑DeepSeek-R1:端侧推理实战指南

直接放结论:DeepSeek-R1 是能跑进浏览器的,而且不是玩具级演示。我用 WebGPU 后端 Transformers.js 把量化后的 R1 蒸馏模型装进了 Chrome,完全端侧推理,数据不出本地,生成速度在我的 M 系列芯片上能到每秒 30~60 tok…

📰

Excel参数表分块秒传方案:前端解析、批量提交与增量比对实战

1. 车间里那张20MB的参数表,为什么每次上传都要点好几遍重试机械制造行业的MES、工艺管理、ERP这些系统,我接触过不少,几乎每个项目里都会遇到同一个尴尬场景:工艺员手里有一张Excel工艺参数表,十几兆甚至几十兆&#…

📰

java获取项目路径的5种姿势与面试避坑指南

java获取项目路径的5种姿势与面试避坑指南 Java 8 升级到 Java 17 后, ClassLoader.getResource 的行为突变,导致大量 实战项目 在打包成 Jar…

📰

Yii2 缓存机制全面解析:数据缓存、查询缓存、片段缓存、页面缓存与 HTTP 缓存实战指南

后端Web框架 【免费下载链接】yii2 Yii 2: The Fast, Secure and Professional PHP Framework 项目地址: https://gitcode.com/gh_mirrors/yi/yii2 点击查看 免费下载 缓存是 Web 应用提升性能的一种廉价而有效的手段:把相对静态的数据存入缓存&#xf…

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

读完文章,想聊聊您的网站?

告诉我们您的行业与需求,资深顾问一对一梳理方案与报价,全程免费。

📞 💬