模型并行切分粒度的选择没有绝对最优解,核心在于匹配你的集群拓扑、显存容量与计算效率;粗粒度省通信但显存浪费严重,细粒度显存均衡但通信开销指数级上升,实操中往往需要在两者之间找到性价比拐点。
这几年做大模型训练,并行策略几乎是绕不开的坎,你可能会搜到各种关于模型并行训练切分粒度怎么选的文章,但真正落地时,很多人还是在张量并行、流水线并行和序列并行之间反复试错,今天这篇不聊空泛的理论,直接切入切分粒度的真实权衡场景,以及我们在实践中总结的取舍经验。
模型并行切分粒度怎么选:先看清通信与显存的天平
选择切分粒度,本质上是在做一道资源配置题,你的GPU卡间通信带宽、每张卡的显存容量、以及训练样本的序列长度,共同决定了最优解的位置。
Tensor Parallelism与Pipeline Parallelism的粒度差异
-
张量并行(TP):将单个Transformer层的权重矩阵按行或列切开,分别放在不同GPU上,粒度细,但每个Transformer层的前向和反向计算都需要all-reduce通信,通信频率极高,比如在8卡A100节点内,NVIDIA NVLink的带宽能到600GB/s,TP8是可行的;一旦跨节点,网络带宽骤降为100Gb/s级别,同一切分粒度会导致通信等待时间成倍放大。
-
流水线并行(PP):按层切分,将不同层放在不同设备上,粒度相对较粗,通信只发生在相邻切分点之间,每轮micro-batch只需传输一次激活值,通信频率远低于TP,但粗粒度带来的问题也很直接如果某几个层的显存占用极不均匀,卡的闲置概率会明显增加,出现“前卡忙死、后卡闲等”的负载失衡现象。
行业共识认为,TP适合单机多卡的高带宽环境,PP更适合跨节点的长管道传输场景,你问模型并行训练切分粒度怎么选,第一原则就是看你的通信硬件边界在哪。
切分粒度粗与细的实际开销对比
| 对比项 | 细粒度(TP为主) | 粗粒度(PP为主) |
|---|---|---|
| 通信频率 | 每个算子后都触发 | 仅切分边界触发 |
| 显存碎片 | 少,但单卡冗余量小 | 较多,需预留较大buffer |
| 负载均衡 | 强依赖均匀切分 | 层间大小不均时波动大 |
| 扩卡上限 | 受限于层内并行度 | 受限于总层数 |
| 典型场景 | 单节点8卡A100/H800 | 多节点大规模集群 |
这里有一个经常被忽视的细节:切分粒度越细,计算通信比越差,业内专家指出,当TP并行度超过单节点内的物理卡数时,通信时间可能占到总训练时间的40%以上,这个比例在千亿参数模型上会进一步恶化。
大模型并行切分方案对比:不同场景下的粒度偏好
不同规模的模型和不同硬件环境,对粒度的偏好差异很明显。
单节点场景下的细粒度实践
如果你只有一台8卡A100服务器,想训练一个13B左右的模型,此时TP=8往往是较为自然的选择,NVLink+NVSwitch的全互联拓扑能提供足够的通信带宽,细粒度切分带来的显存节省能让你塞进更大的batch size,实际操作上,配合序列并行(Sequence Parallelism),把LayerNorm和Dropout的激活值也切分开,显存占用能再降一截。
关键参数参考:
- 微批次大小(micro-batch)建议设为1或2,避免过多中间激活累计
- 梯度累积步数与PP的micro-batch数相乘等于全局batch size时,收敛稳定性更好
- 开启AMP混合精度后,注意TP通信的梯度压缩策略,避免精度损失叠加
多节点场景下粗粒度的妥协
跨节点训练时,网络延迟成为主要瓶颈,此时如果将TP粒度延续到跨节点,你会发现每个step都很慢网络往返延迟直接加在关键路径上,更常见的做法是:节点内使用TP细粒度切分,节点间使用PP粗粒度衔接。
一个典型的12节点96卡训练配置可能是这样的:
- 每个节点内TP=8,充分利用NVLink
- 节点间PP=12,将96层网络按层分布到12个节点
- 数据并行(DP)维度根据全局batch size调节
这种混合并行方案的存在本身,就说明了一个问题:真实的切分方案往往是多粒度共存的,单独谈论粒度的粗细没有多大意义,组合策略才是日常训练中的常态。
切分粒度影响训练效率的关键机制
要理解为什么粒度调整会产生这么多连锁反应,得稍微深入一点看计算图层面的事件序列。
气泡与激活值:两个最大的隐性成本
-
气泡:PP的粗粒度必然引入流水线气泡,即某些GPU在等待前一阶段的输出时处于空闲状态,切分越粗,阶段数越少,单个气泡的时间越长;但切分越细,阶段数越多,气泡出现的频率也越高,实践中,前向与反向的计算量比例约1:2,通过调整micro-batch数量来填充气泡,是一个需要反复试错的过程,我们一个具体的案例里,把1F1B调度下的micro-batch从4升到8,吞吐提升了约18%,但显存压力也随之增加。
-
激活值显存:细粒度TP切分能分摊激活值的存储压力,但代价是每一层都需要在通信算子处同步,模型并行训练时的显存权衡是个长期存在的痛点,尤其在序列长度超过4096的场景中,激活值的显存占用甚至可能超过权重本身,这时就需要考虑激活值重计算(activation recomputation),用额外的计算换显存通常是保留每一层的输入,反向时重算输出,显存能减少一半以上,但重算带来的计算开销约在20%-30%之间。
2D/2.5D并行与切分粒度的新变化
现在英伟达的Megatron-LM、华为的MindSpore,以及PyTorch的FSDP方案,都在往更灵活的混合并行方向演进,FSDP本身是把参数、梯度、优化器状态全部切分,属于比较粗粒度的策略,但它内部也支持逐层注册钩子来控制切分边界。
实操中,我们用PyTorch FSDP时,通常使用sharding_strategy=ShardingStrategy.SHARD_GRAD_OP,配合auto_wrap_policy按TransformerBlock维度切分,具体路径:torch.distributed.fsdp里的FullShard表示全切分,ShardGradOp表示只切梯度算子,不同配置在显存占用和通信量上的差异,多试几次就会有直观感受。
模型并行训练显存权衡与性价比曲线
聊完机制,回归到最实际的决策维度成本。
显存节省幅度与通信延迟的性价比拐点
当切分粒度从粗到细逐步变化时,显存占用确实是单调下降的,比如一个7B模型在BF16精度下,权重就占约14GB,如果再加上Adam优化器状态,单卡训练至少需要56GB左右显存,切分到2卡,每卡压力减半;切分到4卡,进一步减少。
但通信的开销却是阶梯式上升的:
- 从单卡到2卡,通信量增加一截
- 从2卡到4卡,通信占比明显提升
- 从4卡到8卡,如果跨节点,通信延迟可能突然恶化数倍
这个“显存收益衰减、通信成本陡增”的交点,就是我们说的性价比拐点,找到这个拐点的方法其实不复杂:固定你的global batch size,记录不同切分粒度下的有效吞吐(MFU),画出曲线,取峰值左侧的点作为生产配置。
多机多卡环境下的过犹不及
多机场景下,粗粒度能有效减少跨机的通信频率,但过粗也会带来另一个问题:单机故障影响面变大,如果你的PP深度等同于节点数量,一台机器挂了,整个训练任务基本就停了,所以实务中,很多团队会选择PP深度为节点数量的1.5倍到2倍,留出故障转移的余地。
如果你用的是公有云GPU实例,比如简米云的灵骏集群或AWS的ParallelCluster,网卡带宽和延迟在不同机型上的差异很大,通用的做法是在选型阶段用NCCL的all_reduce_perf
测试工具跑一遍,实际测出不同通信量下的时间开销,再代入你的切分方案进行计算。
模型并行训练调参实战:从理论到收敛
调参这件事,理论说再多都不如跑一次来得实在,但有一些经验可以作为起点,帮你少走弯路。
通信算子与切分策略的选择建议
我们团队在最终落地时,通常遵循这样一套判断流程:
- 先确认单卡显存能否装下权重+优化器状态+激活值,如果能,优先考虑纯数据并行
- 单卡放不下,尝试TP,但TP度不超过单节点GPU数
- TP无法覆盖时,叠加PP,PP深度根据机器数量决定
- 最后用ZeRO/FSDP兜底,处理剩余显存缺口
这套流程被很多内部项目反复验证,虽然不是唯一的路径,但至少是一个不太容易踩坑的起点。
从显存不足到训练加速的完整过程
举个例子,我们曾经在一个4节点、每节点8张A800的集群上训练一个70B模型,初始方案使用PP=4、TP=8,结果发现每个micro-batch的耗时达到了将近3秒,大部分时间都花在了节点间的激活传输上。
简单的调整是,把PP提高到8,TP降为4,虽然每层的计算并行度降低了,但单次通信的数据量也变小了,再配合将micro-batch从2调整到4,管道填充效果提升,最终训练吞吐提升了约25%,显存峰值也控制在安全线以内。
这中间的测试过程涉及一套可以复用的NCCL调试路径:
- 用
nsys profile抓取一个step的详细时间线,定位通信算子的耗时占比 - 使用
NCCL_DEBUG=INFO查看每次集合通信的消息大小和耗时 - 调整
NCCL_MAX_NCHANNELS或NCCL_NTHREADS参数,观察带宽利用率变化 - 结合
torch.profiler确认是通信瓶颈还是计算瓶颈
这些路径看起来繁琐,但在线上调优时,效果往往比拍脑袋猜测要好得多。
Q&A:关于模型并行切分粒度,你还需要知道的
Q:对于不太熟练的算法团队,一开始选粗粒度还是细粒度更稳妥?
A:粗粒度起步会更稳妥,PP的实现逻辑相对直观,调试成本更低,尤其适合刚接触分布式训练的小规模团队,细粒度TP每次改动切分边界都要重新验证正确性,且通信问题排查难度较大,如果不是有明确的显存压力,不建议一上来就挑战。
Q:模型并行训练时显存不够,主要有哪些可行的改进方向?
A:顺序依次是:开启激活值重计算、切换到梯度/优化器状态分片(ZeRO/FSDP)、调整Transformer层的并行切分方式、降低batch size,前两者通常能释放出相当可观的显存空间,且对训练吞吐的影响在可控范围内,如果仍不足,再考虑序列并行或上下文并行(Context Parallelism)等进阶方案。
首发原创文章,作者:王坚,如若转载,请注明出处:https://idctop.com/article/624977.html




