ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

大模型分布式训练必知:TP、DP、PP、CP、EP并行策略全解析

大模型分布式训练必知:TP、DP、PP、CP、EP并行策略全解析 算法同学做 LLM 训练和推理早晚会撞上这样几个问题模型参数往显卡里塞不下了好不容易加了几张卡训练速度反而没上去或者跑一个超长文档单卡直接 OOM。这时候身边 Infra 同学会丢给你一串缩写TP、DP、PP、CP、EP。说实话刚听到这些词容易懵因为网上资料多数是给 Infra 看的充斥着“AllReduce”“bubble”“分片”这些名词算法同学拿起来并不容易吃透。本文不聊源码细节只想把这五种并行策略真正讲明白它们各自在切什么维度为什么需要切切完怎么通信以及实际选型时怎么搭配。适合刚接触分布式训练、自己动手跑过大模型但没系统梳理过并行策略的同学也适合打算把训练和推理框架从单卡迁到多卡的工程师。1. 先理清分布式计算的三个底层问题大多数并行策略的出发点不是“怎么把卡用满”而是“模型塞不下了怎么办”和“塞下之后怎么不白加卡”。想搞懂 TP、DP、PP、CP、EP得先回答三个问题显存被谁吃了、并行到底在切什么、为什么加卡不一定会变快。1.1 显存到底被谁吃掉了一个 Transformer 模型训练时的显存占用可以粗略分成静态状态和动态状态两类。静态状态包括模型参数本身、模型参数对应的梯度以及优化器状态比如 Adam 里保存的动量、二阶动量。动态状态主要是前向过程保存下来的激活值还有分布式训练需要的临时通信缓冲区。以 70B 模型为例用 FP16 存参数仅参数就要占 140GB。参数还需要 FP16 梯度再加 140GB。如果用 Adam 优化器每个参数还要额外存一份 FP32 主参数4字节、一阶动量4字节和二阶动量4字节合计 12字节。所以光静态状态70B 模型大约要 140GB 参数 140GB 梯度 840GB 优化器状态 1120GB。单张 A100 80G 显然装不下。所以第一个底层问题就来了必须把模型“切开”让每张卡只承担一部分。这个“切开”的维度就对应着 TP、PP、CP 这些策略。1.2 并行计算的本质切分和通信所有并行策略本质上都是两件事选择一个维度把计算切碎再设计通信方式把碎片结果拼起来。数据并行切的是“batch 维度”流水线并行切的是“模型层维度”张量并行切的是“单个算子内部的计算维度”。后来出现的上下文并行切的是“序列长度维度”专家并行切的是“MoE 模型里的专家维度”。维度不同通信的模式完全不同。数据并行是一步一次全量梯度同步流水线并行只在相邻层之间传激活值和梯度张量并行每个层内都要做多次聚合通信专家并行则是 token 在不同芯片之间做动态路由。理解了这一点再看各种策略的“优缺点”基本都能自己推出来。1.3 为什么加卡不一定会变快很多算法同学第一次上多卡训练都会发现加速比远低于理想值甚至跑出负优化。原因在于并行计算的加速上限除了受并行度影响还受“串行部分”和“通信开销”两头挤压。假如一个 step 原来单卡要跑 100 秒其中 80 秒是纯计算20 秒是必须串行等待的部分即使计算全并行理论加速上限也只有 5 倍这就是 Amdahl 定律的直觉。实际中还有通信开销。每一点并行都会引入额外的数据搬运、同步等待和气泡空洞。更麻烦的是几种并行叠加后通信和气泡会相互作用。所以做并行策略不是“越高越好”而是要在显存、计算、通信三者之间取平衡。2. DP 数据并行最基础也最容易理解数据并行Data ParallelismDP是很多人接触到的第一种并行策略也是五个缩写里唯一“不切模型”的策略。2.1 DP 的核心流程DP 的做法很直接每张卡都保存一份完整的模型副本然后把一个 global batch 切分成多个小 batch分别发给不同的卡做前向和后向。每张卡算出来的梯度不一样所以每个 step 结束时需要做一次梯度 AllReduce把各卡的梯度取平均再用平均梯度去更新所有卡上的模型。由于每卡都有完整模型DP 的显存占用其实没有降低。但它能让多卡一起处理更多数据提高吞吐量。理论上如果每卡一个 step 处理 8 条样本4 卡就相当于每 step 处理 32 条样本训练速度接近线性增长前提是通信开销足够小。DP 的通信量跟模型大小强相关不跟 batch 大小强相关。每次梯度同步需要传输的是模型所有参数的梯度所以模型越大DP 的通信开销越大。对于 7B 模型FP32 梯度大约 28GB做一次 AllReduce每卡实际产生几十 GB 的通信流量这也是大模型不单独用 DP 的原因。2.2 从 DP 到 ZeRO再到 FSDP既然 DP 每张卡都存完整模型状态那还是太占显存。ZeRO零冗余优化器就是解决这个问题的把模型状态切开分到不同卡上需要时再组装。ZeRO 分成几个 stageStage 1 只分片优化器状态Stage 2 再分片梯度Stage 3 把参数也分片。Stage 3 下每张卡只拥有模型参数的一小部分计算前需要通过 Gather 拿到当前层完整参数计算后还可以丢掉。PyTorch 里的 FSDPFully Sharded Data Parallel就是基于 ZeRO 思想实现的分片数据并行工具。方案参数梯度优化器状态通信量DP每卡完整每卡完整每卡完整1次全梯度同步ZeRO-1每卡完整每卡完整分片小幅增加ZeRO-2每卡完整分片分片增加梯度通信ZeRO-3 / FSDP分片分片分片抬高需省显存2.3 DP 的适用场景和关键参数DP 最适合“单卡能装下完整模型但想提高吞吐量”的场景。如果单卡已经放不下模型那就得靠 ZeRO/FSDP 先把参数分掉或者配合后面要讲的 TP/PP。实践中一个小模型比如 7B 以下做指令微调FSDP 是首选方案因为它比纯 DP 省显存使用难度也不算高。还有一个问题是全局 batch 怎么确定。并行卡数多了以后global batch per_gpu_batch × grad_accumulation × data_parallel_size。改并行卡数时如果要保持 global batch 不变就得调整梯度累积步数。学习率缩放也要留意batch 翻倍时很多场景学习率可以近似线性放大但放大有上限建议用 cosine schedule 和 warmup 配合别一上来就暴力放大。3. TP 张量并行把每个算子都拆开张量并行Tensor ParallelismTP是把一个算子的计算矩阵按行或按列切成多份分给不同卡最后再把结果合并。它是解决“单个 Transformer 层太大放不进一张卡”的核心手段。3.1 从矩阵乘法说起一个最简单的线性层计算是 Y XW其中 X 是输入W 是权重。如果想把 W 切开有两种方式。第一种按列切把 W 切成左右两块 W1、W2分别放在两张卡上。X 每张卡都有一份各自算出 XW1 和 XW2最后把输出横向拼接得到完整的 Y。这种切法不需要立刻做 AllReduce代价是每张卡都要持有完整输入。第二种按行切把 W 切成上下两块每张卡只需要一部分输入特征算出部分输出最后要把两边的输出加在一起这时必须做一次 AllReduce。Megatron-LM 中的标准做法是第一个线性层用列切第二个线性层用行切这样在 MLP 中间只需要同步一次激活值减少通信次数。Transformer 的注意力模块也可以用同样思路切分。QKV 投影做列切注意力头天然可以分配到不同卡上最后输出再接一个行切线性层做聚合。整体上每个 Transformer 的前向过程会包含几次关键 AllReduce这正是 TP 的主要成本。3.2 TP 的通信成本为什么比 DP 高DP 的通信是一个 step 一次通信次数少但每一次数据包巨大。TP 则恰好相反它在每个 Transformer 层内部都有多次 AllReduce通信次数非常频繁。TP 的通信数据量跟 hidden_size 和 batch 相关单次 AllReduce 的数据量其实不算大但架不住次数多所以对 GPU 之间的通信带宽和延迟很敏感。这也是为什么 TP 通常只在同一个节点内部使用节点内的 NVLink 带宽几百 GB/s可以承受频繁通信而跨节点的网络带宽远不如 NVLink强行把 TP 拉到 16、32 会立刻被通信拖垮。实践中一个节点的 GPU 数量常是 8 或 16所以 TP 一般取 4 或 8。3.3 TP 的组合策略和显存估算TP 能把参数、梯度和优化器状态大致按 TP 大小均分所以能极大降低单卡显存压力。比如 70B 模型TP8 时每卡参数降低为 17.5GB梯度 17.5GB。但 Adam 优化器状态依然很重还需要继续用 ZeRO 或 PP 切分才能把显存压到合理范围。另外需要注意的是TP 下每张卡依然要计算完整的 LayerNorm、Dropout 等元素级操作这些部分不做张量切分。如果开启了 Sequence Parallelism可以把 LayerNorm 和 Dropout 也按序列维度切分进一步省显存这也是很多框架里 “Sequence Parallel” 选项的含义。4. PP 流水线并行把模型一层层切开流水线并行Pipeline ParallelismPP的思路更直观Transformer 有几十甚至上百层把不同层的计算分给不同的卡卡与卡之间只传递层与层之间的激活值和梯度。4.1 为什么需要按层切朴素层并行的问题最朴素的按层切分是把第 1 到第 N/2 层放在 GPU0第 N/21 到第 N 层放在 GPU1。跑前向的时候GPU0 必须等 GPU1 完成才能继续等跑反向时GPU1 要等 GPU0 传梯度。如果同一时刻只有一个 GPU 在工作其他 GPU 都在等待加速比会很低。于是出现了 micro-batch 流水化把一个 batch 进一步切成多个 micro-batch第一个 micro-batch 进入 GPU0 后GPU0 继续处理第二个 micro-batch不必等第一个跑完整条流水线。这样 GPU0 和 GPU1 能在不同 micro-batch 上同时工作重叠度提高气泡bubble变小。4.2 气泡率估算流水线并行有一个逃不开的气泡率。假设阶段数为 pmicro-batch 数为 m理想状态下气泡率可以近似为 (p-1)/(mp-1)。举个例子p2、m4 时气泡率约 20%如果 p8、m4气泡率约 58%大部分算力都在空转。这也解释了为什么 PP 的 stage 数量不能随意增大。常见实践里 PP 一般取 2 到 8再往上就要靠增加 micro-batch 数量来填补气泡但增加 micro-batch 会增大激活显存也会增加通信压力。4.3 PP 的边界通信与负载均衡PP 的通信量比 TP 小很多因为只需要在每一段边界传递激活和梯度传输的数据形状是 [batch, seq_len, hidden_size]。但 PP 依然有三个问题要注意一是气泡二是各 stage 计算量要尽量均衡三是反向传播顺序会影响显存峰值。你在框架里看到的一些调度名词比如 GPipe、PipeDream、1F1B、Interleaved其实都是在处理“什么时候做前向、什么时候做反向、怎么减少气泡、怎么压低显存”。对于算法同学来说不需要死记所有调度名只要记住PP 适合大模型 大批次场景且 micro-batch 数要足够否则性能很差。5. CP 上下文并行与 EP 专家并行TP、DP、PP 是经典三板斧但是当模型开始往超长序列和 MoE 方向演进之后又出现了两个更“新潮”的并行维度CP 和 EP。它们不是替代前三者而是补充。5.1 CP当序列长度成为瓶颈上下文并行Context ParallelismCP切的是序列长度维度。为什么需要切序列长度因为一个样本的序列长度一旦到了 128k、1M token注意力矩阵的大小会随序列长度平方增长哪怕有 FlashAttention单卡也可能存不下一个样本的中间结果。CP 的做法是把长序列分成多段每张卡持有一部分 token 对应的 Query、Key、Value 和激活。每张卡只需要计算自己这一段的本地注意力但要得到完整注意力每张卡还要知道其他卡上的 Key 和 Value所以需要在卡之间循环传递 KV 片段。Ring Attention 就是这个思路把 KV 在卡之间像接力棒一样传递一圈把显存压力从“单卡一次全量”变成“随时间序贯传递”。5.2 CP 的适用场景和与 TP 的配合CP 主要面向超长文本比如长文档理解、长上下文推理、多轮对话场景。如果序列长度只有 4k、8k直接用 FlashAttention 就够了没必要上 CP但如果序列长度超过 64kCP 会非常有用。CP 和 TP 可以一起用TP 切注意力头CP 切序列长度两者叠加能进一步降低每张卡的 KV 显存。实际操作中如果上层框架不直接暴露 CP 参数它往往会以 Sequence Parallelism 或 Ring Attention 的形式藏在长序列训练方案里。看到seq_len维度的并行逻辑把它理解成 CP 就好。5.3 EP把专家分散到不同卡专家并行Expert ParallelismEP专门针对 MoE 模型设计。MoE 模型里的 FFN 被替换成多个“专家”每个 token 只会被路由到其中 top-k 个专家。如果不做 EP每一张卡都要保存所有专家权重参数量会非常大。EP 的做法是把不同 expert 分布在不同卡上token 通过 router 被送到目标卡去计算计算完再传回来。这个通信模式叫 All-to-All特征是在一批 token 中每个 token 可能被发往任意卡所以通信模式不像 AllReduce 那么规整更容易成为瓶颈。EP 最大的意义是让模型总参数可以远超单卡显存。比如 Mixtral、DeepSeek-V2 这类模型模型总参数量很大但每个专家单独看都很小用 EP 可以让专家分散存放和计算同时保留较高的单卡计算效率。EP 通常还要和 DP/TP 组合在一个“每组几个卡”的并行布局里交错设置。6. 组合拳混合并行选型与实操估算真实的大模型训练很少只用一种并行方式。70B 或更大模型的训练方案几乎都是“DP TP PP activation checkpointing”甚至再加 CP、EP 的混合体。这一节直接给选型思路和估算方法方便你拿到一个模型后快速定方案。6.1 选型路线图先把话放前面没有最优配置只有“够用”的配置。下面这个表是一个通用出发点实际还要根据显存、节点数、互联方式调整。模型规模典型场景推荐并行组合1B~7B单卡能跑需要吞吐DP 或 FSDP13B~30B单卡放不下TP DP或 FSDP70B 以上单卡绝对放不下TP PP DP加 ZeRO开 activation checkpointing超长序列文档、代理、多轮长上下文上述方案加 CP / Sequence ParallelismMoE 大模型专家路由EP DP必要时加 TP/PP6.2 一张卡到底怎么算要多少张以 70B 模型训练为例前面算过静态状态大约需要 1120GB。如果使用 TP4、PP4、DP8总卡数4×4×8128 张那么每个模型的“模型维度”被切成了 16 份每卡承担的静态状态大约是 1120/1670GB。再加激活值、通信缓冲和显存碎片已经逼近 80GB 上限还得继续开 activation checkpointing 或 ZeRO 去压。所以一个简单的手算公式是单卡显存需求 ≈ (参数量 × 每个参数需要的训练状态字节数) ÷ (TP × PP × DP分片因子) 激活值。其中“DP分片因子”取决于是否用 ZeRO纯 DP 分片因子是 1ZeRO-3 下优化器、梯度、参数都会被进一步分片。核算时还要记得留 20%~30% 余量给通信和临时张量别顶着显存上限设计。6.3 在框架里怎么配置不同框架的配置入口不一样但核心参数基本一致。以 Megatron-LM 为代表的训练框架常见两行参数就是--tensor-model-parallel-size 4 --pipeline-model-parallel-size 4总卡数 world_size 除以 TP×PP剩下的就是数据并行度。比如 128 卡TP4、PP4则 DP128/(4×4)8。另一个框架 DeepSpeed 会用 ZeRO 和流水线引擎需要额外指定--num_stages或者 pipeline stage 数。PyTorch 用户用 FSDP 时设置sharding_strategyShardingStrategy.FULL_SHARD即可开启类似 ZeRO-3 的分片。推理侧也一样vLLM 常用--tensor-parallel-size来配置 TP 大小。如果做长上下文推理一些框架还支持上下文并行参数或者用--max-model-len限制长度后用 FlashAttn 和分块 prefill 来缓解显存。6.4 实测调优顺序我个人的调优顺序是这样的先在小规模配置上跑通正确性比如 2 卡、TP2、DP1。逐步增加 TP 和 PP看吞吐tokens per second和显存变化。用 profiler 看每个 step 中通信时间占比。如果通信超过 30%大概率是并行维度拆得太碎或 batch 太小。保持 global batch 不变调整 per_gpu_batch 和梯度累积观察端到端耗时。最后再开 activation checkpointing、混合精度、异步通信这些优化开关。不要一上来就朝“尽量多切”的方向堆并行度。先满足显存底线再在速度、通信、稳定性之间找平衡这样不容易出大问题。7. 常见问题与排查技巧实录这一部分是我实际在算法同学迁多卡过程中见过最多的问题整理成几个典型场景直接说现象和做法。7.1 加了卡之后训练反而变慢最常见的原因是并行方式不适合场景比如模型不大却强行开 TP8或者 PP8 但 micro-batch 很少。另一个常见原因是每张卡的 batch 太小计算时间短通信时间占比反而高。排查时先看 GPU 利用率如果利用率不高且训练耗时长跑一次 profiler 看通信和 kernel 耗时占比。通信高就减少 TP/PP 程度或提高 micro-batch / 梯度累积计算低就看看是不是数据加载成了新瓶颈。我见过不少模型从 4 卡加到 8 卡反而变慢最后发现是每卡 batch 没变导致整体 batch 翻倍学习率没调对收敛不了还浪费算力。7.2 OOM但预算明明够显存超限不一定是模型太大常见原因是激活值堆积、通信缓冲区预留、显存碎片化以及 ZeRO 分片状态下临时 Gather 的峰值。解法通常有三个开 activation checkpointing把前向激活值丢弃、反向时重新计算减小 per_gpu_batch 或调低 micro-batch开启 sequence parallelism 或 CP 来切分激活值。activation checkpointing 会增加约 20%~30% 计算时间但显存下降常常非常明显。如果开完还 OOM就继续降 batch或者检查是不是开了过大的 extra padding / padding。7.3 MoE 卡在 All-to-AllMoE 模型做 EP 时如果 token 路由不均匀有些卡会特别忙有些卡闲着GPU 利用率波动剧烈。All-to-All 通信本身也会有大量小数据包可能把带宽打满。解决方向包括增大 batch、保证路由 token 数量统计更平稳把专家并行组大小调得和卡间拓扑匹配确保 All-to-All 尽量在节点内部发生通信和计算重叠比如在等待远程 expert 结果时先把本地非专家部分的计算做完必要时用 TP 切分单个专家权重降低专家维度的单卡显存。这些手段都值得在当前框架里逐个试。7.4 长序列推理 OOM长序列推理时KV Cache 会非常大。如果发现超过一定长度就 OOM先看max-model-len是否设置合理再看有没有开 PagedAttention 或 FlashAttention。如果序列确实超长就需要上 CP 或序列并行把 KV 分散到多卡。另外推理不一定要把输入截断可以分 chunk 处理前面的 token再拼接输出。很多推理框架的 “chunked prefill” 就是这个目的。7.5 各类问题速查表现象排查方向常用解法加速比低、GPU利用率低profiler 看通信/计算占比调小 TP/PP、增大 micro-batch、检查数据加载显存 OOM看激活和临时缓冲activation checkpointing、减小 batch、开 CPPP 性能差气泡大增加 micro-batch、降低 PP stage 数EP 通信慢All-to-All 阻塞增大 batch、调整 EP 组、通信计算重叠超长序列 OOMKV Cache 过大开 CP、FlashAttention、限制 max-model-len最后分享一点个人经验这五个缩写其实对应着五个不同的切分维度DP 切数据TP 切矩阵PP 切层CP 切序列EP 切专家。实际工作中真正困难的不在于记住每个缩写而是选好组合并定位通信瓶颈。我测试过几十个卡数配置最快速度不是靠猜出来的而是先用一个小规模 profiler 跑一遍找出通信时间占比再决定调 TP 还是 PP。建议你手上有一个已经能跑的模型之后专门拿一个晚上的时间把所有并行策略分别在 2 卡/4 卡上跑一遍你会对开销有个非常直观的感觉后边做方案会从容很多。
返回列表