Torchtitan训练优化组合拳:激活重计算、torch.compile与Float8实战

发布时间:2026/9/9 2:47:54
Torchtitan训练优化组合拳:激活重计算、torch.compile与Float8实战 做大规模预训练的人可能都有同一个感受显存永远不够用。模型权重、梯度、优化器状态、中间激活值每一项都在跟GPU抢显存而模型的规模还在越卷越大。PyTorch团队开源的Torchtitan这段时间我拿它当主力工具跑了几轮对比实验最大的感受是它把激活重计算、torch.compile 和 Float8量化这三套优化方法组合在一起形成了一套非常实用的“训练优化组合拳”。这篇文章就是记录我实际打开这三层优化时的配置、数据以及踩过的坑。先说结论再讲细节如果你要训练一个7B以上的稠密模型尤其是序列长度超过4096的场景那么激活重计算基本是必需品torch.compile 是性价比最高的加速手段而Float8量化则适合已经暴露通信瓶颈或显存瓶颈时的下一级选择。三者叠加往往能让你在同样的八卡机器上从“勉强跑起来”变成“有余量做更大的batch和更长的序列”。这篇内容适合正在用PyTorch做大模型训练、想找到一个可以直接照抄的优化配置或者单纯想搞明白这些技术背后逻辑的人。我会把原理、实操、数据、坑位一次讲完。1. Torchtitan是什么为什么值得专门聊一聊1.1 它不是又一个训练框架而是一套“原生PyTorch的规模化配方”提到大规模训练很多人第一反应是Megatron-LM、DeepSpeed这类框架。Torchtitan的定位不太一样它是PyTorch团队自己维护的参考实现目标不是造一个新的轮子而是用最纯粹的PyTorch API搭出一套能训练Llama 3.2 90B甚至更大规模模型的完整方案。这里面的分布式能力比如FSDP2、DTensor、张量并行、上下文并行全部基于PyTorch原生的分布式API没有魔改。我最初接触它也是带着怀疑的毕竟做训练优化的人早就习惯了“框架帮你把所有脏活做完”的体验。但实际跑下来你会发现正是因为Torchtitan没有藏东西所有的并行策略、优化开关、通信策略都是显式的配置项它非常适合做训练优化方案的“实验床”。你想看激活重计算带来了多少收益开和关跑两个实验就行。你想验证Float8量化对吞吐的影响也是一个flag的事。这在其他框架里往往要翻很多源码。1.2 为什么三种优化技术要放在一起谈如果你单独问哪一项优化最有效其实很难回答。因为大模型训练的资源瓶颈通常不是单点显存吃紧的时候可能是激活值占得太多吞吐上不去的时候可能是kernel启动和显存读写太频繁多卡通信卡住的时候往往是带宽不够。这三样问题恰好对应这三项技术激活重计算专门对付显存里的“中间激活值”torch.compile 用算子融合和代码生成把计算效率往上顶Float8量化则把线性层的计算和通信量降下来。关键在于这三项技术在资源上互补而不是简单叠加。激活重计算节省显存代价是增加了约三分之一的计算量这个多出来的计算量正好能被torch.compile优化掉一部分Float8量化只动精度、不动模型结构和另外两项没有冲突。在Torchtitan里它们都是配置文件里的开关你可以像搭积木一样组合它们然后通过MFU和显存占用两个指标来评估效果。这也是为什么聊训练优化时这三项总是一起出现。2. 三种优化的原理以及它们省下的资源到底从哪来2.1 激活重计算用“重新计算”换显存先解释清楚激活值为什么这么吃显存。Transformer每一层的前向传播都会产生若干中间张量attention的Q、K、V、注意力矩阵、MLP的中间激活等。这些张量在反向传播时还要用所以训练过程中默认会全部保存在显存里。一个粗略的估算量级是当seq_len4096、hidden_size4096、层数为32时仅激活值就可能达到数十GB比模型权重本身还大。序列越长、batch越大激活值增长得越夸张。激活重计算的思路非常直接前向传播时不再保存这些中间激活而是等到反向传播需要的时候再临时重新跑一遍前向计算把它“算回来”。你可以把它类比成做饭时不再把每道菜切好的半成品都放进冰箱冷藏而是记住菜谱需要下一步时重新处理食材。代价是前向计算量增加了大约30%到40%收益是激活值占用的显存可以下降一个甚至几个数量级。实操中它有两种粒度一种是全量重计算整个Transformer层的前向都重算显存省得最狠但计算开销最大另一种是选择性重计算比如只丢弃注意力和MLP块内的一部分中间结果显存省得少一些但计算增量也更小。在模型规模不大、序列长度中等时选择性重计算的性价比更高到了70B以上或者超长序列往往就得上全量了。2.2 torch.compile算子融合和代码生成带来的计算加速PyTorch默认的eager模式其实很“奢侈”。每执行一个算子都需要GPU启动一次kernel算完的结果写回显存下一个算子再读出来。当一个模型有几百上千个小算子时大量时间都耗在了kernel launch和数据搬运上而不是真正的计算上。torch.compile 的解决方案是把计算图交给Triton编译器分析把能合并的算子融合成一个大kernel减少kernel启动次数同时优化显存读写。它在Torchtitan里几乎是“一键开启”的。开启后首次训练会有一段明显的编译过程日志里能看到“Optimizing module”“Compiling graph”之类的信息。编译完成后后面的迭代就会用上优化后的kernel。我实测下来对kernel比较多的模型收益尤其明显原因很简单小算子越多融合的空间越大。如果你的模型里线性层、归一化层很多torch.compile基本是白给收益。但它也有代价编译本身要吃掉不少时间第一次跑一个7B模型编译可能要花几分钟到十几分钟编译过程对CPU和内存有压力如果模型代码里有编译器不支持的算子会退回到eager模式你需要从日志里识别出来。另外强烈建议在训练时固定输入shape避免动态shape触发反复重编译否则收益会被编译时间吞噬掉。2.3 Float8量化8位浮点如何在不伤精度的情况下降低显存和算力开销Float8训练不是简单地把模型换成FP16/FP8就完事。FP8这个格式很有意思它有两种主流细分E4M34位指数、3位尾数和E5M25位指数、2位尾数。前者精度更高适合前向计算后者动态范围更大适合反向传播和梯度累计这类对范围要求更高的场景。在实际训练里主权重仍然保留在BF16或FP32只有线性层的计算部分用FP8执行这样既享受低精度的速度又不会因为精度不够导致模型不收敛。用生活里的事情打比方就好比记账时你始终用精确的小本子记总账但给合作方看过程数据时用的是更精简的报表——总量不会错但传输更快、更省空间。具体到显存FP8把线性层的前向计算量减半、激活值和梯度可以按FP8格式传输直接降低了通信带宽的压力。在Torchtitan里开启浮点8以后还会涉及一个挺关键的选项是否使用8x8缩放。默认的缩放是每张量一个缩放系数8x8则按区块独立缩放精度更好但计算略有增加通常用于对数值稳定性要求高的场景。需要特别提醒的是FP8需要硬件支持目前主流是NVIDIA H100、H200以及AMD MI300系列。A100和更早的卡是跑不了FP8的别开了以后发现kernel直接报错或退化成模拟模式。此外FP8对loss scaling的依赖比较高数值范围如果下溢会损失梯度信息所以在Torchtitan里开启float8后要关注loss曲线是否出现尖峰。2.4 三种技术叠加时它们各自扮演什么角色收拢一下这三种技术的分工激活重计算解决的是“显存装不下”的问题torch.compile 解决的是“GPU算得慢”的问题Float8量化解决的是“通信和计算饱和”的问题。大模型训练通常不是被某一个瓶颈卡死而是几个瓶颈同时存在所以组合的收益往往会超过单项收益的简单相加。我刚开始做这一套组合的时候最容易忽略的是激活重计算和Float8之间的配合。激活重计算在反向时多出来的那部分前向计算恰恰可以用FP8的kernel来加速而FP8降低了通信量后重计算带来的额外迭代时间又被部分抵消了。torch.compile则像是给所有kernel都加了一层加速buff。简单说这不是三个孤立的开关而是三个互相喂饭的齿轮。理解这一点你在调参时就不会只看单项指标的变化而是能判断一个指标下降会不会被另一个指标上涨补回来。3. 在Torchtitan里把这些优化逐级打开3.1 基线配置先让一个模型稳定跑起来先别急着开优化第一步永远是跑通一个基线。我用的是单机8卡NVIDIA H100 80GB环境PyTorch 2.4以上版本CUDA 12.4拉取最新Torchtitan代码后安装依赖。Torchtitan的入口是一个run.py脚本加一个TOML配置文件里面可以指定模型规模、优化器、并行策略、检查点等所有内容。以Llama 3 8B为例官方直接给了训练配置python run.py --config training/llama3_8b.toml --override model.flavor8B --checkpoint.enable_checkpoint这里我故意先关掉了所有优化让模型在默认eager模式下跑起来。基线跑通以后记录两个最关键的指标峰值显存可以通过nvidia-smi或Torchtitan日志里的显存统计看和模型吞吐量MFU模型浮点利用率。说白了后续每开一个开关你都要拿这两个数跟基线对比才能知道收益到底花在了哪里。注意别忽略一个小事数据集和参数要一致。我在做对照实验时固定使用同一份小数据子集batch size和序列长度全部保持一致否则后面看到的显存和速度差异会被数据规模波动污染掉。3.2 第一拳开启激活重计算在Torchtitan里激活重计算的打开方式非常直接。命令行加一个flag即可python run.py --config training/llama3_8b.toml --activation-checkpoint --activation-checkpoint-modefine_grainedfine_grained对应选择性重计算默认会丢掉部分中间激活如果你已经走了全量重计算的极端路线就把模式改成full。跑起来之后别急着看速度先看显存在8B、seq_len8192的场景下峰值显存通常能从原来的接近OOM降到剩余几十GB的安全区间。这个变化非常直观属于“一开就知道有效”的优化。有一个值得注意的细节激活重计算本身不改变浮点运算的数值它只是重新计算同样的前向所以理论上不该影响loss收敛曲线。但如果某个重计算的实现有bug比如有时重算的输入是更新后的权重而不是该时间步的权重loss就会异常。如果你发现开启重计算后loss曲线和基线对不上先确认Torchtitan版本再考虑是不是混合了别的优化器或梯度累积设置。3.3 第二拳接入torch.compile激活重计算开完之后模型很可能已经能塞进显存了但速度可能下降了计算量增加了。这时候就可以把torch.compile加上把多出来的计算量“找补”回来python run.py --config training/llama3_8b.toml --activation-checkpoint --activation-checkpoint-modefine_grained --compile首次运行会有一段明显的编译期我见过8B模型编译花十几分钟的情况这是正常的。Torchtitan的compile底层走的是Inductor也就是PyTorch的默认编译后端。如果你愿意牺牲一点编译时间来追求更好的kernel可以往更激进的后端调但不建议一上来就这么干先把默认compile跑通后续再看要不要优化。编译完成后训练吞吐会有一个明显提升。我印象最深的是同样一个7B模型在短序列场景下compile开与不开MFU可以差出10到15个百分点。这个提升对任何训练任务都太香了。需要提醒的是如果启用了compile之后日志里出现“fallback to eager”之类的字样说明模型里有些算子编译器不支持这部分会退回原生执行但整体影响一般不大只是收益会打折。3.4 第三拳Float8量化在确认激活重计算和torch.compile都能稳定运行之后最后加Float8python run.py --config training/llama3_8b.toml --activation-checkpoint --activation-checkpoint-modefine_grained --compile --float8-linearfloat8-linear是Torchtitan里对线性层启用FP8训练的开关。如果你的基础模型是Llama结构那么attention里的qkv投影、o投影以及MLP里的所有Linear层都会被替换成FP8计算。主权重仍是BF16Adam优化器状态仍是FP32所以模型质量不会因为量化而崩掉。同时你还可以考虑开启更细粒度的缩放方式比如--float8-linear-configfloat8_linear_8x8不过这个开关会稍微增加计算开销我个人建议在小模型上先不开8x8直接跑到大模型、长序列时再打开因为8x8的数值优势在序列很长、梯度很小的场景里才更明显。Float8叠加到compile之上以后你可能会看到loss曲线出现极轻微的变化这通常是正常的但如果loss出现明显的spike或者发散优先检查是不是loss scaling没有配好或者硬件不支持FP8导致走了模拟路径。3.5 完整组合配置参考上面给的命令是手动一个flag一个flag加实际使用中我通常把配置固化到TOML里方便团队里别人直接复现。下面是一个我在8卡H100上跑Llama 3 8B、seq_len 8192时用的参考配置片段[training] batch_size 1 sequence_length 8192 max_steps 5000 [parallelism] data_parallel_replicate_dim -1 tensor_parallel_degree 1 context_parallel_degree 1 [optimizer] name AdamW lr 3.0e-4 [optimizer.fsdp] mode full_bf16 [activation_checkpoint] mode fine_grained [compile] enable true [float8] linear_enabled true不同版本的Torchtitan里配置段的名称可能稍有出入以你拉取的版本为准。但整体结构差不多。实际跑的时候如果用命令行覆盖配置Torchtitan会优先读命令行参数所以你可以用一套base配置配合命令行微调这很方便做对照实验。最后说明一点——我建议你先按3.2到3.4的顺序一项一项加每加一项跑50到100个step记录显存和MFU再决定要不要继续而不是一把梭把三个开关全打开。原因很简单如果全打开之后出了问题你很难定位是哪一个环节导致的。4. 实测数据与调参经验4.1 显存和吞吐对比我以Llama 3 8B、seq_len 8192、micro batch size 1、8卡H100 80GB的环境为基准把三轮优化前后的数据整理成了下表。需要声明的是不同Torchtitan版本、不同PyTorch版本、不同GPU集群下数据会有浮动这个表的意义在于看趋势而不是当绝对标准。方案峰值显存/卡约说明MFU约基线bf16 eagerFSDP2接近OOM约70GB激活值占大头35%左右开启激活重计算fine_grained约35GB显存压力大幅下降约30%计算量上升拉低效率重计算 torch.compile约33GB显存再降一点计算效率明显提升45%左右重计算 compile Float8约28GB通信和计算同时下降50%左右从这张表里能读出几个信息第一激活重计算对显存的贡献最直接但也把MFU打下去几个点torch.compile恰好把这个损失吃回来了甚至还有富余Float8是压轴项它让显存进一步下降也让MFU再上一档。三者全都打开的时候你的显存余量会非常大这时候就可以考虑增加batch size、加长序列或者干脆把训练目标换成更大的模型。这也是“组合拳”最吸引人的地方它不是让你固定一个小的batch死磕而是给你空间去探索更大的训练规模。4.2 先开哪个后开哪个顺序很重要我在跑实验时踩过一个很典型的坑一上来就把三个开关全打开结果编译阶段报了显存不足。原因在于torch.compile的编译过程和FP8的kernel launch都依赖一定的显存空间如果你用激活重计算把显存压得太狠编译时反而可能触发OOM。所以更稳妥的顺序是先把激活重计算打开等训练稳定下来再叠加torch.compile确认编译通过且速度提升后最后再开Float8。每加一项跑一小段确认没有OOM和loss异常再动下一项。另一个经验是激活重计算的粒度要跟着模型规模走。8B规模我用fine_grained就够了显存降幅够大额外计算量也可控。70B以上我见过直接上full的配置虽然计算变多了但只有这样才能在有限显存里塞下更大的batch。所以你也可以把“重计算粒度”当成一个可调的旋钮而不是非黑即白的开关显存不够、但计算有余就往full方向调计算吃紧、显存还够就退回fine_grained。4.3 模型规模不同效果差异很大不要以为这套组合在任何规模下都一样。小模型1B以下本身显存占用不大激活重计算带来的收益就不明显反而增加了计算量这时不如只开torch.compile。7B到13B是收益最均衡的区间显存压力有但不至于非全量重计算不可torch.compile的加速又非常可观Float8也能明显减少通信。到了70B以上我基本建议把三项全开并且还要配上张量并行和上下文并行否则单靠FSDP2可能很难塞进单卡显存。还有一个此前容易被忽略的点序列长度。如果你的业务以短文本为主激活重计算的收益会被稀释但如果你在做长文档、多轮对话序列长度动辄8K、16K激活重计算省下的显存就非常可观。换句话讲组合拳的最佳使用场景是大模型 长序列 多卡训练如果只是小模型推理加速这些都是杀鸡用牛刀。5. 常见问题与排查技巧实录5.1 编译时间过长或编译期内存不足torch.compile最让人劝退的就是首次编译时间。7B模型编译十几分钟是常见的更大的模型可能更久。解决办法第一升级PyTorch到较新版本编译性能提升明显第二如果编译期CPU或内存爆炸检查是不是同时开了很多并行任务可以把训练脚本里会触发编译的并行度调低第三在容器或本地机器上给训练进程预留足够内存我见过编译阶段因为内存不足直接退出导致误以为显存问题的案例排查时看日志比猜更有效。另外不要小看编译期的显存占用。在部分环境里编译生成的临时GPU kernel也会占用显存所以刚开compile之后出现短时间显存飙升是正常的。如果确实因为这个OOM可以考虑先关掉其他占显存的优化等编译完成后再恢复或者给compile换一个更省内存的inductor配置。5.2 loss不收敛或出现尖峰开启Float8之后loss曲线偶尔会比bf16基线抖一点这是预期内的。但如果loss持续不降或者出现明显尖峰优先排查三件事第一是否开启了loss scaling相关的配置FP8梯度一定要有合适的缩放因子第二是否关闭了某些不该动的精度选项例如对比一下full_bf16和full_fp32混合精度策略的差异第三看看是不是多卡通信时梯度被截断成FP8造成了大值损失这种情况可以考虑改用8x8缩放或加大缩放范围。这里有个排查小技巧先把Float8关掉只开激活重计算和compile跑同一段数据看loss是否恢复到基线水平。如果恢复了那问题就锁定在Float8如果没恢复问题可能出在compile或建模代码上。用“二分法”定位问题比盯着loss曲线瞎猜高效得多。5.3 显存还是不够怎么办把三项优化全开之后仍然OOM说明规模确实超过当前硬件规格了。这时候我会按顺序再做三件事第一开启张量并行tensor_parallel_degree调整把模型的层参数切到多卡上这是最直接的显存卸载手段第二降低micro batch size牺牲一点吞吐让训练能跑起来第三如果序列特别长开启上下文并行把序列维度也切分到多卡。Torchtitan对这几项并行策略的配置都很透明改配置文件就行。还有一个容易被忽略的变量是checkpoint和日志的存储路径。有的环境会把训练日志、profile文件、checkpoint写到显存盘上挤占可用显存或IO带宽。实际排查时可以先关掉checkpoint落地、关掉profile看看显存是否立刻宽裕如果明显变好就把存储路径挪到容量足够的普通磁盘上。5.4 其他容易踩的坑记录几个零散的坑。一是版本问题Torchtitan迭代很快同一个flag可能在相隔几个月就有完全不同的名字网上搜到的旧教程直接复制可能报错建议以仓库的文档和--help输出为准。二是多机训练时的NCCL版本Float8的梯度通信依赖NCCL对FP8的支持如果NCCL版本太旧可能不报错但你也不知道它悄悄丢了精度。三是tp和cp的粒度当张量并行加大以后激活重计算、Float8的收益会被切分组合效果需要重新评估。四是监控训练时建议打开Torchtitan自带的内存统计和吞吐统计别只看nvidia-smi因为nvidia-smi显示的是进程预留的显存而不是实际张量占用两者差距可能很大。我实际用下来的体会是这三项优化里最容易被低估的是torch.compile最容易被高估的是Float8。很多人一听到量化就觉得一定省很多但在小模型上Float8的收益并不惊艳反而是torch.compile几乎每次都能带来稳定的吞吐提升。而激活重计算则是那个“闷声干大事”的角色它不直接提速但给整个组合创造了显存余量让你有底气去开更大的batch、更长的序列。最后再分享一个小技巧做这类对比实验时我会把每个配置跑完后的显存和MFU记录到一张表里连同loss曲线一起存档。下次要扩规模或换卡时这张表就是最可靠的调参依据——毕竟优化组合拳打得好不好最终看得还是数据和经验而不是感觉。