
最近一直有人在后台问我大模型训练动不动就要几十张甚至上千张显卡显存到底是怎么省下来的那么多卡又是怎么协同工作的我之前在系列前几篇里聊了Transformer架构和反向传播的底层逻辑这一篇终于要进入最硬核的实战环节了混合精度训练和分布式训练。严格来说这两块内容分开讲都是一门课但实际训练大模型时它们又是深度绑定的。做混合精度是为了省显存、加快计算做分布式是为了把省下来的显存和算力池化到更大规模。如果你打算从“能跑通小模型”迈向“能训练大模型”这两个技术点是绕不开的坎。这篇文章我尽量用大白话把原理和实操串起来讲明白希望能帮你少踩几个坑。1. 为什么要动精度和分布式的念头一个7B模型的显存账本先算一笔账你就知道问题出在哪了。一个70亿7B参数的大模型按FP32精度来存每个参数占4字节光模型权重就是 7 × 10^9 × 4 28GB。可训练过程中不能只存权重吧还得存优化器状态AdamW的话每个参数要存一阶动量 m 和二阶动量 v又是两个FP32以及反向传播时要用的梯度又一个FP32。再加上前向和反向过程的中间激活值显存需求轻松超过100GB。一张目前主流的专业显卡比如NVIDIA A100 80GB或者H100 80GB单独一张卡连一套完整的7B训练流程都跑不起来更别提动辄几百B的大模型了。这就是为什么行业里常说“训练大模型不是算力问题而是显存问题”。显存摆不平算力再高也是空的。解决办法有两条路一是把每个字节“抠”着用这就是混合精度训练要干的事二是把数据、模型、优化器状态分散到多张卡上这就是分布式训练要干的事。两条路不是二选一而是同时用这也是目前大模型训练框架的默认配置。2. 混合精度训练为什么不是“纯低精度”混合精度训练的核心直觉很简单模型里的某些东西可以用更少的位数来存只要不让精度崩掉就行。但这里有个关键点名字叫“混合”不是“全换”。很多人一开始以为把权重从FP32换成FP16就完事了这是最容易翻车的地方。2.1 FP16、BF16和FP32之间的差异先过一下基础知识。FP32是单精度浮点数1位符号位、8位指数位、23位尾数位它能表示的数值范围大约是 ±3.4×10^38精度能达到小数点后约7位有效数字。FP16是半精度1位符号位、5位指数位、10位尾数位范围只有 ±65504有效数字约3位。问题就出在这个“3位有效数字”上。大模型训练时梯度值经常非常小比如1e-7这个量级。在FP16下这个数根本表示不了直接就会变成0这就是所谓的“下溢”underflow。梯度要是变成0了参数就再也不更新了模型直接废掉。为了解决这个问题NVIDIA和各大框架搞出了BF16BFloat161位符号位、8位指数位、7位尾数位。它的指数位和FP32一样多所以能表示的数值范围和FP32几乎一致但尾数少了精度低一些。这意味着BF16不会像FP16那样出现严重下溢小梯度仍然能表示只是不够精确。对大模型训练来说梯度“丢精度”的影响远小于“直接清零”因此BF16几乎成了大模型训练的默认选择。2.2 混合精度训练的完整闭环光是换数据类型还不够混合精度训练真正难的是整个训练流程里每种数据用多少位。标准流程是这样主权重master weights保存在FP32这是模型参数的“权威版本”每次更新都对着它来前向传播把FP32权重临时转成FP16/BF16用低精度计算目的是加速矩阵乘法和省显存反向传播得到的梯度也是低精度的但存到梯度缓冲区时会转回FP32优化器状态Adam里的m和v始终是FP32更新时用FP32梯度去更新FP32主权重。你可能想问为什么不一开始就把权重存成FP16我试过这种“激进”做法结果训练没几步loss就开始失控。原因在于每次更新都是一个很小的修正量比如1e-6量级如果主权重是FP16更新的步长小于FP16的精度分辨率时这次更新就白做了。FP16的尾数只有10位一个数越大能表示的最小区间越大更新量很容易被“吃掉”。FP32主权重相当于给整个训练过程一个可靠的参考系。2.3 损失缩放Loss Scaling的实际操作虽然BF16缓解了下溢问题但FP16仍然怎么绕都绕不开。这就是损失缩放登场的原因。在混合精度训练的早期版本里尤其是FP16时代训练loss经过反向传播得到梯度后很多梯度的绝对值已经小于FP16能表示的最小正常数约6×10^-8基本等同于0。为了不让它们消失做法是在前向传播时先把loss放大比如乘以1024反向传播时梯度整体也被放大了等梯度存回FP32再加回来。实操中一般推荐动态损失缩放Dynamic Loss Scaling框架会自动监测梯度的溢出情况。如果一段时间内没有溢出就逐步把缩放因子调大比如从1024调到2048如果检测到有inf或nan就回退并调小。PyTorch的torch.cuda.amp.GradScaler就是这个逻辑autocast负责把算子自动切换到低精度GradScaler负责缩放梯度。我用过一个比较稳妥的组合混合精度策略bfloat16如果是Ampere及以上架构损失缩放BF16其实不需要动态缩放因为它的指数范围够大但FP16必须配如果你在V100上只能支持FP16别偷懒一定要上动态缩放。实测中不缩放的FP16训练十次有八次会因为梯度下溢而loss不动。3. 分布式训练把一张卡装不下的模型“拆开”分布式训练的本质是把模型、数据、梯度拆到多张卡上并行计算再用通信把结果汇聚起来。这里最大的误解在于分布式训练不只是“多个GPU跑同一个脚本”而是有不同层级的并行策略选择哪种策略直接决定了训练效率。3.1 数据并行最直观的并行方式数据并行Data Parallelism是最直觉的思路。每张卡都复制一份完整模型喂进去的数据是不同的batch。前向计算各自算反向计算出梯度后要经过一次全reduce操作把不同卡上的梯度加起来求平均然后每张卡再用这个平均梯度更新自己的模型副本。数据平行的通信开销是梯度同步每轮每张卡都要收发全部模型参数对应的梯度。模型越大通信量越大。当模型大到单卡装不下时纯数据并行就失效了。另外还有个细节数据并行通常要求每张卡上的batch size一致总batch size 单卡 batch size × 卡数。如果总batch size过大收敛效果会有变化一般会配合学习率调整。3.2 模型并行与流水线并行把模型切开当模型超出单卡显存数据并行就不够了这时要把模型本身切开。这里有三类常说的并行张量并行如Megatron-LM的做法、流水线并行如GPipe、PipeDream、以及深度学习的另一种分布式范式ZeRO。张量并行是把一个Layer里的权重矩阵按行或列切成多块每张卡负责计算一部分计算过程中每步都要做all-reduce来同步中间结果。通信频率极高适合在服务器内部用NVLink组网的高带宽环境。流水线并行则是按Transformer层来切第1-8层放卡A第9-16层放卡B数据像流水线一样依次通过各卡。它的通信频率低很多但存在流水线气泡问题各卡会有空闲等待时间。你会发现这些策略本质都在做相同的一件事用通信换显存。选择什么并行策略取决于你的硬件拓扑和模型大小。3.3 ZeRO把冗余状态也算清楚ZeROZero Redundancy Optimizer是微软DeepSpeed提出的方案它洞察到一个数据并行中的关键问题数据并行时每张卡都存着一份完整的模型权重、梯度和优化器状态这些是多卡间完全冗余的。ZeRO把优化器状态、梯度、模型参数分段切开分到不同的卡上。三阶段是ZeRO-1切分优化器状态显存节省约4倍ZeRO-2进一步切分梯度显存节省约8倍ZeRO-3连模型参数也切分训练时按需做all-gather取回。ZeRO-1和ZeRO-2几乎不影响通信效率因为它们本来同步梯度时就要通信。但ZeRO-3每个layer前向/反向都要实时收集该层权重通信量比数据并行高了不少。所以ZeRO-3通常会配合梯度检查点activation checkpointing和更精细的通信调度来用。4. 实战配置与踩坑记录从跑通到跑快的心路历程原理听再多不上手永远是纸上谈兵。这一节我记录一次实际调参过程用的是一个70亿参数模型在8张A10080GB上的训练配置。机器环境是CUDA 12.1、PyTorch 2.1、DeepSpeed 0.12。4.1 第一步显存估算与模型策略选择70亿参数按FP32主权重算权重28GB、梯度28GB、Adam优化器状态56GB合计112GB。这就是为什么单张80GB卡完全装不下的原因。如果开了BF16训练时每张卡上的进程仍会保留一份FP32主权重和优化器状态所以112GB这个数值并不是按卡数均分的“总额”而是实实在在每张卡都要维护的一份。用上ZeRO-1后优化器状态被切成8份每张卡只存一份 56GB/8 7GB显存压力大幅降低。如果只是做ZeRO-1每张卡仍然需要存完整的28GB权重和28GB梯度加上中间激活感觉还是比较紧。稳妥起见我直接选了ZeRO-2并开了activation checkpointing。激活值这一块非常占显存尤其是序列长度长的时候开启activation checkpointing后以约30%的计算开销换来大量显存回退这买卖划算。最终显存占用大约是FP32权重28GB 梯度切分后 28/83.5GB 优化器状态(56/87GB) 激活约10GB。总计约50GB在80GB卡上余量充足。这个余量给了batch size和sequence length调整空间训练起来踏实得多。4.2 第二步DeepSpeed配置解析DeepSpeed的配置文件是训练流程里最关键的一份文件直接决定显存怎么分、通信怎么优化。我当时用的核心配置长这样{ train_batch_size: 64, train_micro_batch_size_per_gpu: 8, gradient_accumulation_steps: 1, fp16: { enabled: false }, bf16: { enabled: true }, zero_optimization: { stage: 2, offload_optimizer: { device: none }, contiguous_gradients: true, overlap_comm: true }, activation_checkpointing: { partition_activations: true, cpu_checkpointing: false }, communication_data_type: fp16 }几个关键点值得解释一下。train_batch_size是全局的总batch size它等于gradient_accumulation_steps × train_micro_batch_size_per_gpu × 卡数。用64 1 × 8 × 8正好对上。如果你改了这个数值但没同步改另外两个训练可能会直接报错或者出现batch对不上的问题。bf16.enabled和fp16.enabled只能二选一。在高阶卡上我强烈建议开BF16。我当时在A100上对比过FP16需要额外维护GradScaler而且因为动态损失缩放的存在偶尔会看到loss出现小幅跳变BF16则不需要训练过程平滑了很多。offload_optimizer默认是none即优化器状态保持在显存里。如果你显存实在吃紧可以把optimizer状态offload到CPU内存速度会慢一些但能把单卡显存占用压到极低。我当时80GB够用就没开offload。contiguous_gradients和overlap_comm是DeepSpeed特别有用的两个开关。前者把梯度整理成连续内存块后者让梯度计算与通信重叠。开启后整体吞吐能提升10%-15%属于白捡的优化。4.3 第三步启动训练的命令与日志解读DeepSpeed启动命令通常长这样deepspeed --num_gpus8 train.py \ --deepspeed ds_config.json \ --model_name_or_path path/to/model \ --per_device_train_batch_size 8 \ --gradient_accumulation_steps 1 \ --learning_rate 1e-5 \ --bf16跑起来以后我盯了几个关键指标cpu_mem和gpu_mem确认显存占用与你预估相符。如果显存占用持续增长先查是不是缓存泄漏loss在混合精度场景下loss出现nan或inf九成是梯度溢出或学习率过大throughput每秒处理的样本数如果吞吐偏低优先检查通信瓶颈其次是数据加载。我第一次把8卡跑通的时候发现有个卡的利用率只有60%多。排查了半天发现是数据加载用了默认的dataloadernum_workers太小每轮到卡取数据都要等。把num_workers调到16并开启prefetch_factor之后GPU利用率拉到了90%以上。这个坑很基础但很多人都会踩。4.4 第四步通信瓶颈的识别与解决分布式训练最大的隐形杀手是通信。有时候你看到8张卡都在跑但吞吐就是上不去此时八成是卡在梯度同步上。判断方法很简单在DeepSpeed配置里把overlap_comm关闭如果训练时间反而变短了说明你的通信和计算其实是串行的并且通信时间占比太高。这时可以优先做三件事检查是否用的NCCL通信库并确认网络接口选择正确调大train_micro_batch_size_per_gpu增大单卡计算粒度减少通信频率在ZeRO-2下把梯度切分打开用多张卡的通信冗余换取更小的单次通信量。我调参时把micro batch从4调到8吞吐直接提升了近30%因为通信被分摊到了更大的计算量上。这也是为什么很多框架教程反复强调分布式训练优先调大单卡batch size再去想其他花活。5. 常见问题与排查技巧实录训练大模型的过程本质就是不断遇见新bug、不断排查的过程。以下是我踩过的一些坑整理成表格方便你对照。故障现象可能原因解决方案loss为nan训练直接崩梯度溢出学习率过大开BF16或动态损失缩放降低学习率重启loss长期不变化FP16梯度下溢换BF16开FP16的GradScaler显存OOMbatch size过大激活值过多减小micro batch开启activation checkpointingZeRO升级到Stage 2/3多卡吞吐远低于预期数据加载瓶颈通信瓶颈调大num_workers检查NCCL增大单卡batch size不同卡显存占用严重不均流水线切分不平衡改用ZeRO或调整层切分策略BF16的loss比FP32略高正常现象无需处理收敛趋势正常即可还有一个很有意思的坑是我在第一次尝试ZeRO-3时碰到的。模型权重被切分后每张卡按需去别的卡上取参数训练速度比ZeRO-2慢了不少。当时我以为配置有问题后来才发现ZeRO-3必须配合gradient_checkpointing和精心调好的communication_data_type否则通信量会大到拖垮训练。实测下来70B以上的模型ZeRO-3才真正划算7B-13B这个量级ZeRO-2性价比通常更好。另外提醒一下千万不要在混合精度训练中用手动乘一个缩放系数来代替框架的GradScaler/autocast。这种看似“控制力更强”的做法实际很容易遗漏某些算子导致精度不一致而且排查起来极难定位。框架提供的工具是经过大规模验证的直接信任它。6. 训练完成后的模型保存与继续训练技巧模型训练到一定阶段要保存checkpoint这里面也有讲究。混合精度训练下pytorch的save默认保存的是FP32主权重这没问题。但如果你用的是DeepSpeedcheckpoint的保存和加载最好也走DeepSpeed的接口否则权重切分状态会和优化器状态对不上暖启动从checkpoint继续训练时会出现各种怪问题。我习惯每隔500步保存一个checkpoint如果一个checkpoint保存失败整个训练就白跑了。保存频率不是越高越好因为保存checkpoint会打断训练流程、产生额外IO开销。500步对我来说是一个在恢复时间和性能损失之间比较平衡的值。暖启动时的学习率也要注意。如果是从头训练前面一两千步一般用warmup把学习率缓慢升高避免早期梯度震荡如果是加载checkpoint继续训练可以把学习率降为原先的50%-70%。我一般会在训练中断后给一个更小的学习率让loss先从“恢复期”平稳过渡回正常下降轨道。还有一个容易被忽略的点多卡训练时checkpoint只保存主卡的数据。加载时如果用单卡直接加载需要先把分布式环境初始化好否则torch.load时会因为缺少分布式上下文而报错。这些细节往往不会写在框架文档里但实际跑训练时几乎都会遇到。