news 2026/10/10 15:27:43

分布式训练全解析:显存策略、并行方案与工程实践

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
分布式训练全解析:显存策略、并行方案与工程实践

写这篇笔记的起因很实际:前阵子给团队小伙伴讲大模型训练基础,发现大家对“分布式训练”的理解大多停留在“多卡跑起来就是分布式”这个层面。这话不算错,但离真正能上手的距离还很远——显存怎么分、梯度怎么聚合、通信为什么经常是瓶颈、断点续训为什么要保存一堆状态,这些问题如果只看两三个 Demo,根本看不透。

这篇笔记是我自己把分布式训练拆成“显存账本—并行策略—通信模型—工程稳定性”四条线整理的,不追求面面俱到,重点讲清楚每个方案解决什么问题、为什么这么设计、实际跑起来会踩哪些坑。适合刚接触大模型训练、想系统理解分布式原理的同学,也适合已经在用 DDP/DeepSpeed 但没搞懂底层逻辑的工程师对照着看。

1. 为什么必须做分布式:显存与算力的双重瓶颈

很多人以为分布式训练是因为“卡越多越快”,其实真正的起点更朴素:单卡根本装不下,或者单卡算到天荒地老。

1.1 一次训练到底要吃掉多少显存

先算一笔账。一个 7B 参数模型,FP16 权重占 14GB。听起来单张 80GB 的卡绰绰有余,可训练的时候显存里放的不止权重。

正常混合精度训练下,每个参数大致需要 16 字节的“训练状态”:

  • FP16 参数副本:2 字节;
  • FP16 梯度:2 字节;
  • FP32 主权重:4 字节;
  • Adam 优化器的 momentum 和 variance:各 4 字节,合计 8 字节。

加起来是 2+2+4+8 = 16 字节/参数。7B 模型光这部分就是 112GB,单张 A100-80G 直接出局。这还没算前向传播产生的激活值(activation),序列一长,激活值动辄几十 GB,比参数状态还夸张。

这就是为什么分布式训练第一性原理是“拆”:要么拆数据,要么拆模型,要么把参数状态分散开,目标都是让任何一张卡的显存峰值可控。

1.2 是显存不够,还是算力不够

显存只是第一关,算力是第二关。

业界常用一个粗略公式估算训练总计算量:总 FLOPs ≈ 6 × 参数量 × token 数。假设训练一个 7B 模型、1 万亿 token,总计算量大约 6×7×10⁹×10¹² = 4.2×10²² FLOPs。

单张 A100 FP16 算力算 312 TFLOPS,那也得跑 4 年多。这还只是理论满速,真实利用率打三折四折就更离谱。所以不是“想不想分布式”的问题,是根本绕不开。

但要注意,“多卡更快”的前提是每张卡都在做有效计算。如果通信没有跟计算重叠,GPU 大量时间在空等,加的卡越多浪费越严重。后面会专门讲这个。

1.3 看懂并行方式的第一张地图

分布式训练的并行策略可以分四类,先看一张总表:

策略切分对象需要同步的东西典型场景
数据并行训练样本梯度(AllReduce)单卡能放下模型,但数据量太大
张量并行单层内部的矩阵运算每层的中间激活单层超大,放不进单卡
流水线并行网络层层与层之间的激活和梯度整模型按层切分,减通信
ZeRO 系列参数/梯度/优化器状态参数 AllGather、梯度 ReduceScatter显存不够但想保持数据并行语义

一句话总结:数据并行是“每个人一份完整菜谱,各做各的菜,最后交换经验”;张量并行是“一道菜几个人同时切配”;流水线并行是“一条产线上不同工序各有人负责”;ZeRO 则是“同样一份菜谱拆成几页,每人只背一页,用到时再互相借”。

2. 数据并行:大多数人入门的第一个并行方案

2.1 核心思想:复制模型,切分数据

数据并行(Data Parallelism)是最好理解、最容易上手的方案:每个进程持有完整的模型副本,各进程从 dataloader 拿到不同的数据分片,独立完成前向和反向,得到各自的梯度,然后对梯度做一次全局聚合,得到“所有卡梯度的平均值”,再用平均梯度更新每张卡上的模型副本。

这里的逻辑要点是:多卡一起更新等于用了更大的 batch。单卡 batch=2、4 卡并行,全局 batch 就是 8,梯度是 8 个样本梯度的平均。所以数据并行并不是“各学各的”,而是“合起来学同一批数据”。

2.2 AllReduce 与 Ring AllReduce:梯度是怎么同步的

梯度同步是数据并行的核心动作。最简单粗暴的方式是把所有梯度汇总到一台参数服务器,算完平均再广播回去,但这种方式在大规模集群上有单点瓶颈。

实际用到的是 AllReduce,它的语义是“所有进程把梯度数据汇总后,每个进程都拿到最终平均值”。实现方式很多,业界主流是 Ring AllReduce:把多个进程排成一个环,每个进程只和相邻进程通信,把一次大梯度的全局归约拆成 reduce-scatter 和 all-gather 两个阶段。

通信量估算有个常用结论:消息大小为 M、卡数为 P 时,单卡通信量约 2×(P−1)/P×M,卡数多的时候逼近 2M。举个例子,7B 模型 FP16 梯度约 14GB,32 卡做一次梯度 AllReduce,每卡要搬运约 27GB 数据。这个数字很吓人,所以现代框架都在挖空心思让通信和计算重叠。

2.3 为什么现在的框架默认用 DDP 而不是 DP

PyTorch 早期有 DataParallel,后来官方基本推荐 DistributedDataParallel(DDP)。两者都叫数据并行,差别在工程实现。

DataParallel 是单进程多线程,一张卡当主卡,前向后把所有梯度收集到主卡再广播。主卡既是计算节点又是通信中心,负载极不均衡,而且 Python GIL 还会造成瓶颈。

DDP 是多进程,每张卡一个独立进程,梯度同步用的是 AllReduce,且实现里会把梯度按 bucket 打包,在反向传播过程中就异步开始通信,从而把通信时间“藏”在计算里。实测中 DDP 的扩展效率通常远高于 DP,代码上也只多了两三行初始化逻辑。

2.4 数据并行在微调场景下的三个注意点

微调大模型时用 DDP 很常见,但有几个点新手容易忽略。

第一,模型里有 BatchNorm 时要注意。大模型基本以 LayerNorm 为主,BatchNorm 只在 CV 模型里常见。如果模型里真有 BatchNorm,数据并行下各卡算的 running_mean 不一致,要么同步 BN 统计量,要么干脆换成 LayerNorm/GroupNorm。

第二,梯度累积和 AllReduce 的频率。梯度累积的正确姿势是:累积 N 步的梯度,再做一次 optimizer.step(),但每一步都做反向,梯度会不断叠加。DDP 的梯度归约默认每个 bucket 都会触发,所以如果你设置了梯度累积,最好等累积完成后再触发一次全局归约,否则通信次数白白翻 N 倍。

第三,数据分片。用 DistributedSampler 或类似机制保证每张卡拿到的样本不重叠,同时注意 shuffle 的随机种子。我踩过一次坑:微调时每个 rank 用同一随机种子加载数据,结果所有卡读了同一批样本,本质等于单卡 batch,分布式完全失效,loss 曲线还异常平滑。

3. 模型并行与流水线并行:显存放不下的下一站

当单卡放不下完整模型,数据并行就失效了,这时需要把模型本身切开。

3.1 张量并行:把一层切成几块,Megatron 的思路

张量并行(Tensor Parallelism)针对的是“某一层的单个算子太大,放不进一张卡”的情况。把矩阵乘法按维度切开,比如词嵌入矩阵或 FFN 的权重矩阵,切成多个分片分布在不同 GPU 上,每卡算一部分,最后汇总。

Megatron 的典型做法是:对一个线性层 Y = XW,把权重 W 按列切分到多张卡,每卡算出部分和,再通过一次 AllReduce 把部分和合并成完整输出。下一层如果紧跟另一个线性层,则可以把第二层的 W 按行切分,让两次切分错开,减少一次 AllReduce。

张量并行的特点是通信极其频繁,每个 Transformer 层的前向后向都要做多次 AllReduce。所以它必须用在卡间互联极高的环境,比如单机内走 NVLink,机器之间很少直接做张量并行。

3.2 流水线并行:按层切分与气泡问题

流水线并行(Pipeline Parallelism)的切法更直观:把网络按层分组,卡 0 负责第 1 层到第 10 层,卡 1 负责第 11 层到第 20 层。数据像流水线一样流过各卡,只有层与层之间的激活和梯度需要跨卡传输。

它的通信量比张量并行小得多,因为不需要每层内部做多次集体通信,只要把边界激活传过去。代价是“流水线气泡”:前一段算得再快,后一段还没轮到,中间总会有一段 GPU 空闲。

最简单的顺序执行是一段一段地跑,气泡极大。业界主流用 1F1B(one forward one backward)之类的调度策略,把大 batch 拆成 micro-batch,让前向和反向交叉执行,流水线时刻保持尽量多的卡在计算,气泡明显减小。

3.3 两种并行方式的取舍与典型组合套路

维度张量并行流水线并行
切分粒度单层内部层之间
通信量高,每层多次 AllReduce低,只有边界激活和梯度
负载均衡层内相对均衡前后段计算量可能不均
对硬件要求极高带宽,建议机内 NVLink普通高速网络即可
扩展瓶颈通信带宽气泡和调度复杂度

实际大规模训练几乎都是混合并行,不会拿一种方案单打独斗。比较成熟的套路是:机器内部用张量并行,因为 NVLink 带宽高;机器之间用流水线并行和纯数据并行,减少跨机通信量。

3.4 从“并行策略”到“并行维度”:3D 并行到底在说什么

聊到 3D 并行(DP + PP + TP),很多人一开始就懵。我的理解方式是把三种并行当成三个正交的维度:数据并行管“数据怎么分”,张量并行管“单层怎么分”,流水线并行管“层怎么分”。

三者的作用不是重复的,而是各解决一个方向的瓶颈。训练超大模型时,三个维度会同时作用在同一批卡上。比如把 64 卡分成 8 组,组内 4 卡做张量并行,组间 8 个组做流水线并行,再叠加数据并行把 batch 分到所有卡。

理解 3D 并行的最佳方式是先拆开看每一维,不要指望一次性消化。实际工程里也很少有人拍脑袋直接配 3D,通常是从 DDP 跑到显存不够,再加 ZeRO;ZeRO 还不够,再上模型并行。

4. ZeRO 系列:把数据并行做深,而不是换一种并行

ZeRO 是个很巧妙的思路:不改变数据并行的整体结构,只把每张卡上重复存储的参数、梯度、优化器状态分散到所有卡上,谁需要谁临时聚合。

4.1 ZeRO 的三个阶段到底切掉了什么

ZeRO 分三个阶段递进:

  • ZeRO-1:只切优化器状态。每卡只保存自己负责的优化器状态分片,参数和梯度仍然每卡完整。每参数字节数大约从 16B 变成 4 + 12/P。
  • ZeRO-2:梯度也切。反向传播时用 ReduceScatter 把梯度分散到对应卡,不再做完整 AllReduce。每参数大约 2 + 14/P。
  • ZeRO-3:参数也切。前向和反向需要用到某层参数时,临时 AllGather 拉回该层参数,用完释放。稳态存储接近 16/P,但执行瞬间仍然需要临时存放完整参数副本。

关键要理解:ZeRO 不是模型并行,因为计算图本身没有切分,每个 forward/backward 的算子逻辑还是完整模型,只是输入权重需要动态聚合。

4.2 ZeRO-Offload:把参数与优化器状态挪到 CPU/NVMe

显存不够时,除了分片,还有一个思路是借 CPU 内存甚至 NVMe 硬盘。DeepSpeed 的 ZeRO-Offload 会把优化器状态、参数甚至是部分梯度放到 CPU 内存,GPU 只负责真正的前向反向计算。

这个方案最适合“单机多卡 + CPU 内存充足”的微调场景。比如 4 卡 A100-80G 想跑一个 13B 模型,显存紧巴巴,但 CPU 内存有 256GB,完全可以把优化器状态和主权重放在内存里,GPU 上只保留当前需要的 FP16 参数。

要注意 offload 的代价是 PCIe 和内存带宽瓶颈。实测中如果 CPU 内存带宽不够,训练速度会明显下滑。它解决的只是“跑不跑得动”,不是“跑得快”。

4.3 显存优化四件套:梯度累积、激活检查点、混合精度与梯度压缩

真正在工程里帮上大忙的,常常不是某种并行策略,而是四样基础显存优化手段。

梯度累积(Gradient Accumulation):把多个小 batch 的梯度累加后再更新一次参数,等效于增大 batch size,同时不增加单卡显存峰值。配合数据并行时,全局 batch = 单卡 batch × 梯度累积步数 × 卡数。

激活检查点(Activation Checkpointing):前向时不保存所有激活,反向需要时重新计算一段前向。显存大幅下降,代价是计算量上升 30% 到 40% 左右,在大模型训练里几乎是必开选项。

混合精度(AMP):FP16 计算 + FP32 主权重,既能省显存又保证训练稳定。现代 GPU 对 FP16/FP32 混合的吞吐远高于纯 FP32,不用的基本等于白扔算力。

梯度压缩:理论上能降低通信量,但容易引入精度损失,在常规同步训练里用得少,更多出现在去中心化或异步训练研究里。我的建议是常规场景不要碰它。

4.4 一张表看懂:选 ZeRO 还是选模型并行

场景推荐方案理由
单机 8 卡,模型 7B/13B,微调DDP + ZeRO-2/3调参简单,保持 DP 语义
单层超大或超长序列TP解决的是“层放不下”
多机大规模预训练PP + TP + DP通信量被分层控制
显存极度紧张ZeRO-Offload用 CPU 内存换 GPU 显存
只想低成本做 LoRA 微调DDP + LoRALoRA 可训练参数极少,普通多卡足够

一个判断依据:如果只是显存不够,优先考虑 ZeRO,因为它不改变数据并行的语义,调试成本最低。如果 ZeRO 之后通信开销太大、扩展效率上不去,再考虑引入 TP/PP。

5. 通信模型与硬件常识:训练时间的另一半

分布式训练里,计算时间往往容易估算,被低估的是通信时间。很多时候卡数翻倍但吞吐没翻倍,问题就出在通信上。

5.1 同步训练中的通信节奏

大模型训练主流的做法是同步训练:每步更新前,所有卡的梯度必须保持一致。所以一个 step 的时间大致是:前向计算 + 反向计算 + 梯度通信 + 参数更新。通信时间如果没法跟反向计算重叠,就会直接裸露出来。

DDP 之所以快,核心就在它能做通信与计算重叠:梯度的归约被拆成 bucket,反向算到某一层就立刻对这一层的梯度发起通信,GPU 算下一层时,网络在传上一层的数据。这样通信被“藏”在计算里。

但重叠不是万能的。当梯度数据量太大、网络带宽跟不上时,通信还是会浮出水面。典型表现是:扩大 batch 后单卡计算时间变长,通信量不变,扩展效率反而提升;但模型变大、卡数变多时,通信时间随模型大小线性增长,压力陡增。

5.2 带宽、拓扑与等待:NVLink、PCIe 和集群网络

通信快慢取决于硬件拓扑。一台服务器内的多张卡通过 NVLink/PCIe 互联,带宽远高于跨服务器的以太网/InfiniBand。在分配并行策略时,一个基本法则是:通信频率越高,越要放在卡间互联好的范围内。

机内 NVLink 适合张量并行这种“每层都要同步好几次”的通信模式;机间网络适合流水线并行这种“每步只传一次边界数据”的模式;数据并行的 AllReduce 规模与模型大小成正比,对网络的要求介于两者之间。

实际集群规模一大,网络拥塞就变成最常见的性能杀手。我见过不少 case:单机 8 卡跑到 90% 多 GPU 利用率,跨两台机器跑 16 卡立刻掉到 60%,代码没有任何变化,纯粹是网卡队列和路由问题。

5.3 用 FLOPs 估算训练时间:理论算力与现实损耗

估算训练时间用的是“FLOPs 账本 + 算力利用率(MFU)”。模型总计算量用 6ND 近似,然后除以集群总理论算力和 MFU 的乘积。

举个例子:7B 模型、1 万亿 token、256 张 A100。A100 FP16 算力按 312 TFLOPS 算,理论总算力 256×312T≈8×10¹⁶ FLOP/s,MFU 按 40% 算,有效算力约 3.2×10¹⁶。总计算量 4.2×10²² 除以它,约 1.3×10⁶ 秒,也就是 15 天左右。

要注意,这个数字只是纯算力下的理想值,没有算通信、数据加载、checkpoint 写盘。MFU 达到 40% 在真实分布式训练里已经算不错了,很多配置连 30% 都不到。

5.4 张量并行与流水线并行在通信上的差异

三种并行策略的通信特征差异极大:

  • 数据并行:通信量与模型大小成正比,全局 AllReduce,频率一个 step 一次。
  • 张量并行:通信量正比于每层激活大小和层数,频率高,必须走 NVLink。
  • 流水线并行:通信量正比于模型深度方向的激活边界,频率低,对网络要求低。

理解了这一点,就知道为什么大规模训练喜欢组合“机内 TP + 机间 PP/DP”。机内做高频大流量通信,机间只传低频小流量,把宝贵的网络资源用在最需要的地方。

6. 实操层:配置、启动、稳定性与断点续训

理论说完了,落实到跑起来这步,很多坑才会真正浮出来。这一节按我自己的实操顺序写。

6.1 用 PyTorch DDP 与 DeepSpeed 把训练跑起来

如果你只是微调模型,建议从 PyTorch DDP 开始。代码只比单卡多几行:

import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP dist.init_process_group(backend="nccl") local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) model = MyModel().to(local_rank) model = DDP(model, device_ids=[local_rank])

启动命令用 torchrun:

torchrun --nproc_per_node=4 train.py

第一次跑 DDP 时,最常踩的坑是没设LOCAL_RANK、使用 CPU 集群却配了 nccl 后端、或者忘记调用dist.destroy_process_group()。这些在单卡上都不会暴露,多卡一跑就报错。

如果显存紧张,直接上 DeepSpeed 也很快。一个 Stage-2 的配置大概长这样:

{ "train_batch_size": 32, "gradient_accumulation_steps": 4, "zero_optimization": { "stage": 2 }, "fp16": { "enabled": true }, "optimizer": { "type": "AdamW", "params": { "lr": 3e-5, "weight_decay": 0.01 } }, "scheduler": { "type": "WarmupDecayLR", "params": { "warmup_min_lr": 0, "warmup_max_lr": 3e-5, "warmup_num_steps": 500 } } }

建议从 Stage-2 起步,跑通后再考虑 Stage-3 和 offload。一步到位上 Stage-3,经常把“显存问题”变成“通信问题”,排查起来更痛苦。

6.2 全局 Batch Size、学习率与 warmup 该怎么调

数据并行后,全局 batch size 是决定学习率的关键变量。调大 batch 时学习率通常按比例上调:新学习率 ≈ 原学习率 ×(新全局 batch / 原全局 batch)。比如原 batch 32、lr 3e-5,改成 batch 64,lr 可试探到 6e-5。

但这套线性缩放不是无穷尽的。batch 大到一定程度,收益递减,学习率继续放大反而会导致训练不稳定。所以几乎所有的预训练和微调都会安排 warmup,让学习率在前几百步从 0 线性升到目标值,等优化器状态稳定后再进入常态。

warmup 步数一般取总步数的 1% 到 5%。微调任务数据量小,可以适当放宽;预训练任务数据量大,warmup 太长会浪费算力。

6.3 断点续训:不只保存模型权重

单卡训练的 checkpoint 习惯是“保存模型就行”,分布式训练里这套做法不够。

断开续训需要恢复的状态至少包含:

  • 模型参数、优化器状态、学习率调度器状态;
  • 各 rank 的 dataloader 位置和随机种子;
  • 分布式通信所需的 RNG 状态;
  • 如果是混合精度,还有 loss scaler 的状态。

只有模型权重没有优化器状态的 checkpoint,恢复后 optimizer 的 momentum 全部丢失,实际等于重新开始训练,之前那一段算力基本白费。

保存时要注意各 rank 协同。简单方案是只在 rank 0 保存,但加载时要把模型和优化器状态从那个文件均匀分发到所有卡。DeepSpeed 自带 checkpoint API 会处理好分片逻辑,建议直接用它的save_checkpoint和load_checkpoint,顺手把调度器状态也存了。

6.4 常见分布式训练故障与排查思路

列一张我平常用得最多的排查表:

现象大概率原因排查动作
训练刚开始卡住不动NCCL 初始化失败,各 rank 联系不上检查 IP、端口、防火墙,确认所有 rank 都能连通
某张卡 OOM单卡 batch 太大、激活检查点没开、ZeRO 分片不均开激活检查点、减 batch、看显存监控确认负载是否均衡
GPU 利用率波动大网络拥塞、dataloader 瓶颈、通信没重叠看通信时间占比,把数据加载换成异步加载
loss 突然暴涨学习率过高、梯度更新顺序不一致、数据 shuffle 有误检查学习率、确认各 rank 数据不重叠
中途掉卡/训练中断硬件不稳定、驱动崩溃、显存过热查 dmesg、跑一遍显存测试,先排除硬件
扩展效率极低TP/PP 配置不当、机间通信成了主瓶颈重新审视机内/机间并行维度分配

一个挺反直觉的经验:很多“慢”的问题不是卡的问题,而是数据加载线程把 CPU 占满了,GPU 每步都在等数据。用 profiler 先看时间分布,再决定调什么,别一上来就换模型并行方案。

最后说几句实在话。分布式训练的坑,十个里有八个不是“并行策略选错”,而是通信没藏好、状态没存全、数据没分对。我自己最推荐的入门路径是:先在单卡上把模型跑通,再上 DDP 感受通信开销,然后开 ZeRO 解决显存,最后才考虑 TP/PP。每一步都只解决一个新问题,排查范围就小得多。这篇笔记里的公式和配置,都是帮你在踩坑之前先看清账本——账算明白了,方案自然就出来了。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/10 15:27:41

Token到底是什么?一文讲透编程、认证与大模型三种身份

最近总有人用Loongwise这个ID来找我讨论一个问题:Token和Token到底有什么区别。乍一看像个绕口令,但真较起真来,问到了很多人的盲区。写代码的人天天接触词法Token,做大模型应用的人张口闭口Token计费,搞安全的人又在讲…

作者头像 李华
网站建设 2026/10/10 15:27:36

编译原理课程设计:用LL(1)和四元式实现IF-ELSE翻译

简介:面向编译原理课程设计与实验的IF-ELSE条件语句翻译程序实现包,采用LL(1)预测分析并生成四元式中间代码,适合计算机专业学生、编译器入门开发者用来对照词法/语法/语义分析流程,完成或改进同类翻译任务。压缩包共17个文件&…

作者头像 李华
网站建设 2026/10/10 15:27:11

Java Web环境搭建全攻略:从JDK、Maven、Tomcat到第一个Servlet项目

1. 项目整体思路与关键选择1.1 这一篇到底要解决什么问题Java Web这个词,很多新手第一反应是“要学一堆框架”,但实际走一遍会发现,最劝退的往往不是语法,而是“环境怎么搭起来”。我见过不少朋友在B站看教程,视频里三…

作者头像 李华
网站建设 2026/10/10 15:24:14

ConvNeXt轻量食物图像识别实战:边缘部署11类水果分类

简介:本资源是一套基于ConvNeXt架构的11类水果与食物图像识别完整实践方案,面向深度学习初学者与计算机视觉项目开发者,解决自定义图像分类任务中模型选型、数据准备、训练调优与结果可视化等核心问题。压缩包共2000个文件,主体为…

作者头像 李华
网站建设 2026/10/10 15:23:57

DeepSeek Harness 0.2.1 Web部署与插件热更实战指南

1. 项目概述:这不是一次普通更新,而是一次部署范式的切换DeepSeek Harness 0.2.1 这个版本号看起来平平无奇,但如果你真把它当成“小修小补”来对待,接下来的三天你大概率会卡在--public-url配置上反复重启服务,或者发…

作者头像 李华