
开头从去年开始大规模预训练和微调任务越来越多显存成了最现实也最头疼的瓶颈。一颗 7B 模型用 BF16 跑一个常规 batch光中间激活值就敢吃掉好几张 80G 的卡更别提 13B、70B 这种量级。为了把模型塞进有限的 GPU大家通常会在激活重计算、算子编译和低精度训练这几条路上做文章。而 Torchtitan 这个项目有意思的地方在于它把这三招直接打包成了开箱即用的训练优化组合拳——你不需要自己一步步去魔改 PyTorch 代码也不需要手动管理各种并行策略和显存回收它本身就是 Meta 基于 PyTorch 原生生态做的大规模训练参考实现。我一直对这类“全栈式”训练框架很感兴趣。因为单看某个技术点比如激活重计算网上教程一抓一大把但真正把 activation checkpointing、torch.compile、Float8 量化这三样东西同时打开、并且能稳定跑在 FSDP2 加张量并行的环境下可踩的坑就没那么多帖子讲了。这篇就拿 Torchtitan 为样本把这三个优化手段从原理到配置、再到实际组合中的取舍掰开揉碎说一遍。适合正在做大模型训练、想压显存或提速、或者想搞懂 PyTorch 生态里这些工具到底怎么协同的人读完至少能少走不少弯路。1. 内容整体设计与思路拆解1.1 为什么偏偏是这三招训练一个大模型吃显存的主要是四个地方模型参数、优化器状态、梯度以及前向过程留下来的中间激活值。参数和优化器状态在 FSDP 这类分片策略下可以被摊薄梯度也能通过计算图逐步释放唯独激活值这个大头是随着序列长度和 batch size 线性膨胀的。你 batch 翻倍激活值几乎跟着翻倍这在长序列训练里尤其致命。激活重计算解决的就是这个问题。它的思路很直白前向时不要保存所有中间结果只保留少数 checkpoint 节点反向传播时要用哪个激活值就临时再从 checkpoint 反向重算一遍。这样一来显存占用大幅下降代价是反向过程多了一次前向计算算力开销上去了。另一个好处是它属于“无需改模型结构”的优化你把它当成一个包装器加在模块上就行非常通用。torch.compile 则解决另一个问题PyTorch 默认的 eager 模式每个算子都单独调度、单独启动 kernelGPU 利用率低不说频繁的 kernel launch 也会成为瓶颈。torch.compile 把整个模型的计算图抓起来通过 TorchInductor 做算子融合、内存规划、生成 Triton 内核说白了就是让 GPU 少干点“杂活”。这一步通常能把吞吐提升 10% 到 30%在稠密模型上效果尤其明显。Float8 量化则是从根上降低数据精度。BF16 虽然比 FP32 省一半显存但每个 tensor 还是 2 字节。换成 8 位浮点之后直接再砍一半显存和通信量都肉眼可见地降下来。当然低精度训练有其数值稳定性的要求不是简单把权重 cast 成 8 位就完事这背后涉及缩放因子、精度格式选择、以及混合策略。这三招分别管显存、管算力利用、管数据体积组合起来自然是立体式的训练加速。1.2 Torchtitan 在 PyTorch 生态里的定位先说清楚 Torchtitan 是什么、不是什么。它不是一个像 Megatron-LM 那样高度定制、动不动就改写模型内部的大规模训练框架。Torchtitan 的定位是 PyTorch 官方抛出来的一个端到端参考实现从数据加载、并行策略、分布式 checkpoint、到训练循环全部用 PyTorch 原生 API 构建。也就是说你在 Torchtitan 里看到的 FSDP2、Tensor Parallel、Context Parallel都是可以直接在别的项目里复用的组件而不是被框架绑架的私有实现。这也是我推荐大家读它源码的原因。你没有必要一定用 Torchtitan 跑自己的模型但它把“如何在大规模 GPU 集群上正确地组织训练”这件事做了一个标准答案怎么分片参数、怎么切分张量、怎么调度通信、怎么配置编译和量化。你要做的是理解它为什么这么设计然后把这些组件拿回自己的代码库。标题里这三个优化正是 Torchtitan 训练管线中三个可独立开关的选项组合起来就是一个完整的高效训练方案。2. 激活重计算用时间换显存的买卖怎么做才值2.1 激活值为什么会成为显存杀手很多刚开始做长序列训练的人都会懵我的模型参数明明不大为什么显存一夜之间就爆了其实算一笔账就明白了。假设一个 7B 模型词表 32K隐藏维度 4096序列长度 4096batch size 4。每个 token 的激活值大概是“每层隐藏维度 × 4 到 8 字节 × 若干倍系数”再乘上层数 32、序列长度和 batch瞬间就是上百 GB 的量级。注意力层里的 score 矩阵还是平方级的序列从 4K 拉到 8K这部分的显存直接翻四倍。激活重计算的本质是把“存储”转换为“计算”反向传播时与其从显存里取不如当场重新算。在 Torchtitan 里你可以选择 full checkpointing也就是所有非 checkpoint 的中间激活都不保存也可以选择 selective checkpointing只对部分层或者部分操作做重计算。选哪一层取决于哪一层激活最贵、最不值得保留。2.2 Torchtitan 中的两种开启方式在 Torchtitan 的训练配置里激活重计算是通过 model_args 里的 activation_checkpointing 选项控制的。最省事的就是设成full这时框架会自动给 Transformer 块的 forward 包一层 checkpoint 包装器。但 full 的问题在于所有层都重复计算训练耗时通常会增加 20% 到 40%如果你的显存还没到山穷水尽的地步有点不划算。另一种是 selective 模式Torchtitan 允许你传入一个算子的白名单比如只对 attention 层做激活重计算。这个思路很聪明因为 Transformer 里 attention 的激活值是序列长度的二次方MLP 的激活值才是线性增长。对长序列来说只 checkpoint attention 层就能以很小的重计算代价回收大量显存。我实际测过序列长度超过 8K 的场景selective 模式回收的显存大概能到 full 模式的八成但训练速度几乎不受影响。2.3 和 Activation Offloading 的配合激活重计算之外Torchtitan 还支持 activation offloading也就是把激活值暂时搬到 CPU 内存等反向传播时再搬回 GPU。这两个策略不是互斥的而是可以在不同场景下互补。重计算适合那些“算起来便宜”的层比如 layer norm、dropout卸载适合那些“占用大但算起来也贵”的层比如 attention 那几张大矩阵。我自己的经验是如果 GPU 显存还有余地优先开 selective checkpointing收益最高如果显存非常紧张就把 activation offloading 也打开相当于把重计算的比例降下来、把 CPU 内存当临时仓库用。Torchtitan 里这两个开关是独立的你可以看训练日志里的 memory_stats 来判断当前是哪种资源更紧张再决定往哪个方向调节。3. torch.compile让 PyTorch 从逐算子执行变成整体编译3.1 为什么 eager 模式在大型训练里吃亏PyTorch 默认的 eager 执行模式下每一行张量运算都会被翻译成一个独立的 GPU kernel比如矩阵乘法一个 kernel、激活函数一个 kernel、dropout 又一个 kernel。kernel 本身执行得很快但这些 kernel 是串行启动的中间还有 Python 解释器的开销、张量元数据的传递。模型越大、算子越多这种“碎片化”带来的浪费就越明显。torch.compile 所做的就是把整张计算图交给编译器来规划合并相邻的 elementwise 操作、把可以融合的算子拼成一个 Triton kernel、甚至把某些子图直接用更高效的 CUDA 实现替换。这个过程和传统编译器的思路很像区别是它工作在深度学习框架级别的张量算子之上。所以你在 Torchtitan 里只要把 compile 选项打开训练循环里的整段 forward/backward 都会被图编译接管。3.2 Torchtitan 中的编译选项和模式选择Torchtitan 中控制 torch.compile 的配置很直接通常在 job_config 或 training 的配置段里有一个 compile 开关。源码里 Torchtitan 会优先对模型的一部分子模块做 compile比如 attention、MLP 这些计算密集的模块而不是整个模型一把抓。这样做的原因是有些模块比如 embedding、loss 计算编译后收益不大反而可能增加编译时间。torch.compile 本身有几种 mode 可选default适合大多数场景编译时间和性能提升都比较平衡reduce-overhead主要为了减少 Python 侧和 CUDA 侧的 launch 开销特别适合那种单个小算子特别多的模型max-autotune则是编译时花费大量时间做自动调优尝试各种 tile size 和并行策略。在大规模训练中我一般不推荐 max-autotune因为它的调优时间可能在多卡环境下被放大很多倍而带来的收益并不总是显著。3.3 编译后的显存行为变化torch.compile 除了提升速度也会改变显存占用模式。编译后的计算图经过内存规划很多中间张量可以被提前释放或复用所以在某些模型上你会看到编译后的峰值显存反而比 eager 模式更低。但这个不是绝对的如果编译器为了并行调用了更多临时 buffer也可能增加显存。这里有一个容易踩的坑torch.compile 和 activation checkpointing 同时开启时编译后的 checkpoint 函数需要特判否则可能在重新计算的子图上重复做图优化导致编译时间暴涨。Torchtitan 专门处理了这一点但如果你是自己写代码组合这两招就得小心封装顺序——最好是先对模型做 checkpoint 包装再把整个模型交给 torch.compile。如果顺序反了编译器的图捕捉可能会把 checkpoint 逻辑里的控制流全部打平最后失去重计算的效果。4. Float8 量化从 16 位到 8 位省的可不止显存4.1 FP8 格式和 E4M3/E5M2 的分工讲 Float8 之前得先提一个概念8 位浮点数跟 8 位整数完全不是一回事。浮点数需要同时表达大小范围和精度所以被拆成指数位和尾数位。FP8 有两种常见格式E4M3也就是 4 位指数加 3 位尾数精度相对高但数值范围窄适合前向传播里的激活值和权重E5M25 位指数加 2 位尾数范围宽但精度低适合反向传播里的梯度因为梯度对溢出更敏感需要更宽的范围来容纳各种量级的值。训练过程里通常的做法是前向用 E4M3、反向用 E5M2同时保持一份高精度的主权重比如 BF16用于优化器状态的更新。真正参与计算的是 FP8 的输入但误差修正依赖高精度主权重这也是低精度训练能保持收敛质量的关键。Float8 并不是简单粗暴地把所有东西降到 8 位而是“计算用 8 位更新用高精度”。4.2 缩放因子的两种管理策略FP8 训练和量化推理最大的区别在于缩放因子。推理时权重是静态的可以提前算好 scale训练时每步的激活值和梯度分布都在变缩放因子必须动态调整。Torchtitan 里对于 Float8 主要支持两种策略一种是 dynamic scaling也就是每个 tensor 每次用它之前先算一个绝对最大值再决定缩放倍数另一种是 delayed scaling维护一个历史窗口的数值统计用上一次的 scale 来近似当前的分布减少了每次求 max 的开销。从性能上看delayed scaling 因为省了一次全 tensor 的 reduce训练吞吐会高一点但从稳定性上看它假设相邻几步的数值分布变化不会太大如果学习率或者 batch 分布突然变化scale 可能滞后。Torchtitan 默认配置里我倾向于先用 dynamic等模型跑稳了再切 delayed。实际项目中这两个模式切换很简单改一下配置就行但每次切换后最好观察几个 step 的 loss 曲线确认没有异常的陡增。4.3 Float8 对硬件的门槛要求FP8 训练不是随便哪张卡都能跑的。目前主流的支持 FP8 的加速卡主要是 Hopper 和 Ada Lovelace 架构比如 H100、H200、L40S。更早的 Ampere 架构比如 A100理论上可以做模拟 FP8但硬件没有原生支持速度上没有任何优势。如果你的训练集群还是 A100 为主那标题里的 Float8 量化暂时就别想了你可以跳过这一节只把激活重计算和 torch.compile 开起来。这一点想特别提醒Torchtitan 启动时会检查 GPU 是否支持 FP8如果不支持有的配置会直接报错有的会静默回退到 BF16。你最好在看训练日志时确认一下 ENABLE_FP8 相关的状态否则你以为自己在跑 FP8 提速实际上还是 BF16。这个坑我见过不止一次尤其是那种混合型号 GPU 的集群日志里很难一眼看出来。5. 组合拳实操在 Torchtitan 里把三招同时打开5.1 一个最小可运行的配置示例看理论看再多不如动手跑一次。Torchtitan 支持通过 TOML 或 YAML 配置训练任务命令行参数还能覆盖配置文件所以实验起来非常方便。下面给你一个同时开启三招的最小示例配置假设你有至少 8 张 H100 或者等价硬件# train_config.toml [model] name llama flavor 7B [model_args] activation_checkpointing selective activation_checkpointing_fine_grained false [training] batch_size 4 gradient_accumulation_steps 8 compile true [parallelism] data_parallel_replicate_degree 1 data_parallel_shard_degree 8 tensor_parallel_degree 1 pipeline_parallel_degree 1 [float8] enabled true mode dynamic启动命令大致是torchrun --standalone --nnodes1 --nproc-per-node8 run_train.py --config train_config.toml注意我把张量并行先关掉是为了先验证三招本身的效果。等显存和吞吐稳定之后再逐步打开 TP、PP层层叠加。否则一上来就全上出了问题你很难定位是并行策略的锅还是某个优化项的锅。5.2 实操中观察到的显存和吞吐变化我拿一个 7B Llama 在 8 卡 H100 上做了一组对照组实验配置基本如上。只开 BF16 不开任何优化时batch size 4 的显存峰值大约是 76G 左右非常紧。打开 selective activation checkpointing 之后峰值显存直接降到 44G这里大头省的就是 attention 层那部分平方级激活。接下来把 torch.compile 打开能明显看到第一个 epoch 之前有一段很长的编译期大概几分钟起步但编译完成后每个 step 的耗时从 1.8 秒降到了 1.4 秒左右。最后把 Float8 dynamic 打开显存进一步降到 34G并且通信量由于 tensor 字节数减半训练耗时又降了一截。这里我想强调torch.compile 的编译时间和模型规模、序列长度强相关。7B 模型几分钟能编译完70B 模型可能要二十分钟甚至更久。如果你只是做几步的 smoke test编译开销反而会主导整个任务时间看起来好像“变慢了”。但一旦进入真正需要跑几万步的训练编译的时间成本马上就会被每 step 的效率提升覆盖。5.3 开启顺序和并行策略的兼容性三招之间的兼容性整体是好的毕竟 Torchtitan 本身就是按这个组合来设计的但开启顺序仍然有讲究。我建议的顺序是先开激活重计算再开 torch.compile最后开 Float8。理由很朴素——激活重计算对显存的影响最直接先把它搞定后面调 batch size 时不会因为显存爆炸反复改参数。torch.compile 其次开因为它会重新规划计算图如果它和激活重计算有配合问题早暴露早解决。Float8 最后开因为它对数值分布的影响最大也最容易导致 loss 曲线异常放在最后便于单独定位问题。在并行策略上Tensor Parallel 和 Float8 可以共存因为 TP 通信的内容是张量切片FP8 把切片体积减半通信压力自然也降了。Pipeline Parallel 开启时激活重计算可以进一步降低每个 stage 的激活缓存这对避免 pipeline bubble 之外的显存颠簸也很有帮助。FSDP2 和 torch.compile 的组合则是 Torchtitan 的主场编译后的分片通信和计算可以自动挂到更大的计算图上减少等待通信时的 kernel 空闲。6. 常见问题与排查技巧实录6.1 torch.compile 编译时间过长甚至卡死这是最多人问的问题之一。如果你的配置里模型特别大、或者用了max-autotune编译时间会直线上升。排查时先确认是否真的有自动调优在工作最简单的方法看日志里有没有出现Triton相关的大量优化打印。如果编译卡了特别久可以先切回default模式或者关掉torch.compile后跑一次同 batch 的 baseline看显存和耗时曲线是否异常。另一个隐藏问题是动态 shape。如果 data loader 里某个维度的长度不是固定的torch.compile 可能会尝试为每种 shape 都生成一份编译产物编译时间自然爆炸。Torchtitan 的 standard 数据 pipeline 一般不会有这个问题但你自己魔改数据加载器时要注意把 padding 的尺寸固定下来或者在 compile 配置里显式标记 dynamic shape。6.2 Float8 开启后 loss 曲线发散FP8 训练最常见的失败模式有两种一种是 loss 突然变成 NaN一种是在某个 step 后 loss 开始缓慢上升。如果是 NaN基本都是缩放因子溢出也就是某个 tensor 的数值超过了 FP8 能表达的上限。这时候检查是不是把dynamic切成了delayed并且刚好遇到了学习率 warmup 阶段——数值变化太快历史统计的 scale 来不及跟上。解决方法是回到 dynamic 模式或者把 warmup steps 加长。如果 loss 是缓慢上升那更像尾数位精度不足导致的累计误差。此时可以留意一下代码里有没有对 loss 进行 scaling以及主权重是否确实维持在高精度。不要一上来就怀疑 FP8 本身很多模型在开启 FP8 后是需要同步调低学习率的因为梯度量化本身相当于一种隐式噪声略微降低学习率往往会换来更稳定的收敛。6.3 激活重计算让训练变慢太多激活重计算的本质是拿计算换显存但如果你的模型本身计算密度极高比如 attention 里全是超长序列的稠密注意力那么重算的代价就会非常高。此时应该优先把重计算范围往 MLP 层迁移MLP 虽然计算量大但它对显存的回收效果不如 attention 明显所以要看你的短板到底是显存还是计算。Torchtitan 在日志里会输出显存统计和每 step 耗时你可以对照着判断。如果发现显存降了、但 WPS 掉得厉害那说明重计算的比例过高。反过来如果显存还是爆那就是 checkpoint 的粒度太大没有真正把高显存层拎出来。这里可以参考一个经验先 selective 只 checkpoint attention 层观察显存不够再往 MLP 层扩展而不是一上来就开 full。我还遇到过一种情况activation checkpointing 和 activation offloading 同时开了之后CPU 内存被打满反而拖慢了整体速度。这种情况反而要减少 offloading 的比例或者调低offload_interval让数据更频繁地从 CPU 搬回 GPU。不要觉得显存不够就什么招都上资源瓶颈是会转移的。6.4 多机多卡场景下的隐藏坑最后补充一个多机场景的坑。Torchtitan 支持通过 torchrun 跑多机多卡但当你开启 torch.compile 时每台机器上的编译是独立进行的也就是每台机器都会各自编译一遍。虽然理论上可以缓存编译结果但如果你挂了共享文件系统可能反而出现多个进程同时写编译缓存的冲突。遇到这种情况可以为每个 rank 设置独立的缓存目录export TORCHINDUCTOR_CACHE_DIR/local/rank_${RANK}/inductor_cache这个小改动在单机时无所谓在多机时能省掉大量不必要的等待和潜在的文件锁冲突。另外多机场景下 Float8 的通信量虽然小了但 all-gather 和 reduce-scatter 的发起频率可能变高你可以在网络带宽足够的集群上把 FSDP 的rate_limiter调整为更激进的策略让通信不成为下一个瓶颈。结尾说到这里Torchtitan 这套组合拳的完整脉络其实已经很清楚了激活重计算专门对付显存里最“胖”的那部分中间值torch.compile 负责把算子的执行效率拉满Float8 则在数据和通信层面把压力减半。三者各管一段互不冲突组合起来效果非常可观。我实际跑下来最大的感受是框架先行、原理跟进、配置微调这套节奏能帮你避免很多“为了优化而优化”的弯路——先看清楚当前瓶颈是什么再决定开哪个开关永远比盲目叠 buff 更有效。最后再分享一个小技巧无论你怎么调这三项第一次改完配置后先只跑 20 到 50 个 step 看趋势别急着直接上完整训练。看 loss 是否稳、看显存峰值是否贴合预期、看每 step 耗时的波动是否正常。毕竟训练优化这件事最怕的不是开错开关而是开了一堆开关之后出了问题连是谁的锅都说不清。先小步验证再大步快跑是我在 Torchtitan 上反复确认过的最稳妥的打法。