ARTICLE DETAIL

资讯详情

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

Sequence Parallel与Context Parallel:长序列训练推理的显存优化实战

Sequence Parallel与Context Parallel:长序列训练推理的显存优化实战 今年年初我在调一条超长上下文的推理链路序列长度从32K往百万token上拉的时候单卡显存先崩了。崩的不是模型参数是激活和KV cache。折腾了一个多月把Sequence Parallel和Context Parallel这两类并行策略翻了个底朝天才搞清楚它们到底在解决什么问题、怎么配合、论文脉络是什么。这篇就把这些积累梳理成一篇能直接用的经验贴给同样在做长序列训练或推理的兄弟们一个参照。先说结论Sequence Parallel和Context Parallel本质上都盯上了同一个维度——序列长度。传统的并行无外乎数据并行、张量并行、流水线并行分别切batch、切hidden、切层。但模型越来越长后真正卡住你的是“长度”本身这一维度。把序列维度也切成多份、分发到不同设备这就是SP和CP的立身之本。下面我会把它们的原理、差别、选型和一堆吞吐与显存的坑讲透。1. 先把并行维度的地图盘一遍为什么序列会被单独拆出来1.1 数据并行、张量并行、流水线并行各自管哪一段聊SP和CP之前得先把原有的三维并行说清楚不然很多概念挂不上。数据并行是最直观的每张卡放一份完整模型batch切成多份每卡算一个子batch算完把梯度做一次AllReduce。它的问题也很明显——模型太大时单卡放不下数据并行无能为力。张量并行是模型并行的一种把每个Transformer层内的矩阵按hidden维切开比如一个8卡TP的组里每卡只算4096维hidden中的512维最后靠AllReduce或AllGather把结果拼回来。它效果好但对卡间通信带宽要求极高通常需要NVLink或NVSwitch级别的互联跨节点用TP就会被打爆。流水线并行则是按层切把第1-8层放卡09-16层放卡1模型像工厂流水线一样一卡接一卡。它的显存收益很直接但会有流水线气泡需要靠微批次调度把气泡压缩。这三种并行都成熟很久了可是它们都没解决一个问题单条序列太长时激活和KV cache在单卡上占用的显存照样爆炸。数据并行不解决单序列长度TP虽然能把参数摊开但attention的中间激活和每层的LayerNorm结果仍然会随序列长度线性上涨KV cache更是直接和序列长度成正比。1.2 序列维度为何成了真正的瓶颈来算一笔账感受一下。以7B模型为例hidden size取4096、32层、bf16输入序列长度L32768。单个token的激活大小大致是几个KB到几十个KB量级32768个token的激活总量轻松破几十GB。KV cache更吓人每层每个token要存key和value各4096维bf16下每token每层就是16KB32层就是512KB乘以32768个token单序列的KV cache高达16GB以上。这还没算query端的临时张量。所以序列一长显存瓶颈就从“模型参数装不装得下”变成了“激活和KV cache装不装得下”。这时候你不可能把整条序列塞进一张卡只能把序列本身切开分给多卡让每张卡只负责其中一段。这就是Sequence Parallel和Context Parallel共同的基本盘。这里先给一个简单的对比帮助建立直觉并行类型切分对象每卡保留的内容核心通信算子数据并行batch完整模型AllReduce梯度张量并行hidden部分参数行/列AllReduce、AllGather流水线并行layer部分层点对点激活传递Sequence Parallel序列长度序列片段 模型分片AllGather、ReduceScatter、All-to-AllContext Parallel上下文/序列长度序列片段 模型或KV分片AllReduce、P2P Send/Recv2. Sequence Parallel把非注意力计算也摊到序列维上2.1 用Megatron-LM论文的路子理解SP通信量不变激活却减半很多人知道SP是从NVIDIA那篇《Reducing Activation Recomputation in Large Transformer Models》开始的。这篇论文解决的核心痛点是在Tensor Parallelism里两个线性层之间的LayerNorm和Dropout其实是被“重复计算”的。为什么TP模式下MLP或Attention输出会先做一次AllReduce得到完整的hidden维数据然后每个设备分别拿着完整的序列和hidden各自算LayerNorm、Dropout、残差。这意味着同一份数据在每张卡上都被算了一遍白占显存和算力。SP的做法很巧妙把LayerNorm和Dropout的操作按序列维度切分。每张卡不需要拿到完整的序列只需要拿到自己负责的那一段序列在这一段上做LayerNorm做完之后再AllGather拼回完整序列去喂给下一层。这样操作下来LayerNorm和Dropout的激活显存都只剩原来的1/PP是并行度而通信量几乎没有增加——TP原本的AllReduce在数学上等价于ReduceScatter加AllGatherSP只是把这个组合拆开用顺手把中间激活省了。一句话记住这个套路SP不是一种新算法它是对既有TP通信算子的重新编排把冗余计算去掉让显存均摊到序列维上。我之前在自己集群上复现这个思路时把TP8、序列长度16K的激活优化跑了一遍LayerNorm和Dropout相关的临时张量从单卡8GB左右降到了1GB出头。当时就觉得这套操作真是“白捡的显存”。2.2 从局部走向全体DeepSpeed Ulysses把注意力内部也切开了Megatron-SP解决的是“LayerNorm、Dropout、残差”这类非注意力算子的序列化但注意力计算本身还是完整的——每张卡依旧对整条序列做注意力。想进一步压显存就必须把注意力内部也按序列维度切开这就到了DeepSpeed Ulysses的领地。Ulysses的核心思路是序列块与注意力块之间的置换。输入序列被切成P份每张卡先各算自己那一段的QKV投影然后通过All-to-All算子把“序列维分布”转换成“注意力头分布”。转换完之后每张卡手里是完整序列、但只有1/P的注意力头注意力计算就在头维度上并行不再需要持有完整的序列和QKV张量。这个设计有一个特别好的性质它没有引入任何近似计算。不管切多少卡注意力结果和单卡完整计算完全一致数值上无损失。同时它对通信拓扑的要求是“每个设备能和其他所有设备交换数据”因此All-to-All在NVSwitch这类全互联架构上效率很高。我在实际跑Ulysses的时候最大的感受是它把“序列并行”这个概念从非注意力部分扩展到了注意力内部才算真正解决了长序列训练中激活矩阵O(L²)的爆显存问题。如果用Megatron-SP来处理一条极长序列注意力分数矩阵依旧会在单卡上膨胀而Ulysses切完头之后每个设备只需要算L/P × L/P的局部注意力块显存压力是天壤之别。2.3 我理解的SP本质一类“以序列长度为切分对象”的实现思想看了几篇论文再到代码里翻了一圈我发现一个有意思的事SP其实并没有唯一的标准实现不同论文之间的共性是“按序列维度切分”差异在于切完之后每个设备负责什么、用什么通信算子聚合。Megatron-SP切的是非注意力算子的输入Ulysses切的是注意力中的QKV块和头分布Ring Attention切的是KV块并在线更新softmax统计量。三者的定位和适用的集群都不一样但它们都被归到Sequence Parallel这个大帽子下面。所以读论文的时候千万别抱着“SP就是某一套代码”的心态。更好的理解是SP是一类分布式策略它的核心矛盾是长序列的激活与中间张量如何在多设备间均匀摊开。把这一点吃透再去看任何一篇具体论文都能快速定位它到底在哪个环节切、用什么算子通信、牺牲了什么换来了什么。3. Context Parallel为超长上下文而生的专注选手3.1 CP与SP的分工差异同样的切法不同的战场Context Parallel这个词在长上下文训练和推理场景里出现得更频繁。它的切分对象和SP一样是序列长度但侧重点明显不同SP更多是在训练阶段解决激活显存和计算并行CP则更像一个“为超长上下文优化”的完整系统——它要管好KV cache怎么切、注意力怎么跨设备算、softmax怎么合并、decode阶段怎么低延迟地把片区结果拼回完整结果。拿推理场景举例。用户输入一条长文档长度可能到百万token级。此时模型参数可能是几十GB每张卡放得下但整条文档的KV cache叠加起来就是几百GB甚至上TB。CP的做法非常直接把序列切成长度相等的几段每张卡只持有并维护其中一段的KV cache。整个注意力计算时query端是完整的或者每卡持有一份但要跟所有卡上的KV段分别做注意力最后把结果合并起来。这就引申出一个SP不太突出、但CP必须解决的关键技术点跨设备的softmax合并。3.2 核心难点分布式softmax怎么合并才不出错先看为什么不能简单地把各段注意力输出相加。注意力公式里有一个softmax归一化而softmax的分母是整条序列的exp之和。你把序列切成4段每一段只能算出局部最大值和局部exp和直接加权拼接是错的必须先找出全局最大值再用它去修正每一段的计算结果。工程上最常用的办法是维护两个统计量每个设备局部logits的行最大值m以及exp累加和l。每处理一段KV就把局部的m和l与当前累计的m和l做合并用rescale的方式更新输出。这个就是flash attention里online softmax的分布式版本。核心公式大致是# 每个device持有部分KV计算局部注意力分数 # 合并时: m_local local_row_max # 局部最大值 l_local row_exp_sum # 局部exp和 # 跨device求全局最大值 m_global all_reduce(m_local, opMAX) # 用全局最大值rescale局部概率 p_local exp(s_local - m_global) # 跨device求和得到全局归一化项 l_global all_reduce(p_local.sum(-1), opSUM) # 最终输出 sum(p_local v_local) / l_global实际Ring Attention的代码里会用分块递归的方式更新m、l、o三个状态大致是这样的逻辑m_new maximum(m_old, row_max(s)) l_new l_old * exp(m_old - m_new) row_sum(exp(s - m_new)) o o * exp(m_old - m_new) exp(s - m_new) v这里有个被无数人踩过的坑数值稳定性。如果你图省事先对局部softmax做归一化再跨卡平均输出会吃掉很小但正确的误差长序列叠加后会看到loss异常震荡或生成质量莫名其妙劣化。必须严格按全局max的流程走不能贪便宜。3.3 Ring Attention和它的工程变体通信拓扑决定上限CP最有代表性的实现之一是Ring Attention。它的名字很形象设备连成一个环每个设备持有序列的一段和对应的KV块模型参数在每个设备上完整保留或者再配合TP做参数切分。计算时每个设备按序把自己当前的KV块发给下一环同时接收上一环的KV块每收到一块就更新一次输出统计量。这种环形P2P通信的巨大优势是它不要求全互联拓扑。每台设备只需要跟相邻设备通信在普通的以太网、甚至PCIe互联的机器上也能跑不像TP和Ulysses那样对NVSwitch有硬性要求。代价是延迟一次完整的KV遍历需要环上转一圈通信步数是P步而不是一步到位。Ring Attention的变体里Striped Attention值得单独拿出来说。它解决的是一个很隐蔽的负载不均衡问题当因果mask出现时不同的序列块内部有效query数量差别很大靠近开头的块看到的历史很少靠近末尾的块要处理几乎全部历史。如果KV块按连续区间分配会出现明显的“一头忙死、一头闲死”。Striped Attention把KV块按条纹状交错分配让每个设备手里既有靠前的块也有靠后的块从而在每一轮通信中都保持大致均匀的计算量。这个思想我后来在长上下文训练中直接套用到CP的数据切分上吞吐提升了15%以上。4. 训练和推理场景下怎么选型与组合4.1 一张表看懂SP/CP/TP/DP怎么组合回到实务。你面对一个具体场景到底该开SP、开CP还是两个都开我给一张我常用的决策表场景主要瓶颈推荐组合理由长序列训练32K-128K激活显存 O(L²)注意力TP SP ZeRO/DPTP管参数SP管激活DP管batch超长序列训练128K以上激活 KV cache都爆TP SP CPCP负责把注意力内部和KV也摊开超大batch短序列梯度通信和吞吐DP PP通信按梯度走流水线加大吞吐长上下文推理单请求超长KV cache显存TP CPCP把KV分片TP压缩参数显存多用户长上下文并发动态KV cache调度DP按请求 CP请求级DP天然隔离CP按需扩展单请求上下文特别说明一下TP和SP几乎是黄金搭档。原因很简单SP的那套ReduceScatter/AllGather编排本来就是TP通信流程的拆解两者共用同一套通信组组合起来不会引入额外通信步骤。Ulysses这一类SP则可以直接替代TP中的attention部分但需要谨慎处理与原有TP的交互。总的原则是先用TP把参数显存摊掉再用SP或CP把激活和KV摊掉最后用DP把数据吞吐拉起来。4.2 显存收益的粗略估算7B模型跑32K序列的例子光说理论容易飘我算一个具体例子。假设7B模型bf16权重约14GBTP4时每卡参数3.5GBAdam优化器分片后每卡额外约多个GB。序列长度L32768hidden4096层数32。不切序列时单卡需要缓存一份完整的激活LayerNorm、Dropout、残差、注意力分数等粗略估算在高batch下会是几十GB量级KV cache完整版约16GB。这样合计单卡很容易冲破80GBA100都压不住。打开SPTP4的SP编排后LayerNorm和Dropout的激活变成原来的1/4注意力内部如果再叠加Ulysses式切分注意力分数矩阵也变成约1/4。打开CP假设CP4后KV cache每卡只存4.2GB左右单卡总显存能压到20GB以内。这个估算不追求精确但能让你明白收益的量级SP主要削的是“激活层”的峰值CP主要削的是“KV cache”的线性膨胀。两者叠加才会让超长序列在一个合理规模的多卡集群上真正跑得动。4.3 通信拓扑与硬件是选型的底层约束选SP还是CP、选Ring还是All-to-All最终都要落到你的硬件拓扑上。如果你有NVLink/NVSwitch级别的全互联All-to-All类方案Ulysses式SP、CP中的全局归约最合适。它们一步到位完成数据交换带宽利用率极高用环形反而浪费了高速互联。如果你是多机多卡只走普通以太网或PCIe环形通信Ring Attention式CP是更稳妥的选择。每步只有相邻设备的一次Send/Recv不依赖全互联。缺点是多步延迟。因此Ring类方案要尽量做通信和计算的重叠让当前设备的注意力计算和KV块传输同时进行理想状态下通信延迟被完全隐藏。还有一点容易被忽略decode阶段和prefill阶段的通信特性完全不同。长上下文推理里prefill阶段是计算密集一次要处理大量query通信占比不高decode阶段是每步只生成一个token跨卡通信的同步点会成为延迟的绝对瓶颈。此时如果CP切分粒度过细反而会拖垮单token生成速度。所以生产环境里CP并行度通常要结合batch大小来定batch越大CP的通信摊销越充分。5. 论文地图与上手路线附踩坑记录5.1 值得精读的论文清单与阅读顺序围绕SP和CP不同时间点的论文和对应实现陆续铺垫了一条很清晰的演进路径。下面是我按“从浅到深”排的必读清单里程碑论文/技术一句话贡献阅读建议基础Reducing Activation Recomputation in Large Transformer Models正式提出Sequence Parallelism把LayerNorm/Dropout按序列维切分通信量不变、激活减半必读理解SP怎么嵌入TP注意力内切分DeepSpeed Ulysses用All-to-All实现序列块到头块的置换真正把注意力内部并行化必读理解无近似的序列并行注意力长序列注意力Ring Attention环形KV块传递分布式softmax支持百万token的超长序列训练必读理解CP/Ring的最简形态负载均衡Striped Attention条纹状KV分配解决因果mask下的负载不均衡进阶长序列吞吐优化利器系统梳理Sequence Parallelism: Long Sequence Training from System Perspective从系统视角梳理SP的多种实现范式TP-based、ragged-based、ring-based高屋建瓴适合读完前面再来建立全局观我自己建议的阅读顺序是先读Megatron-SP那篇搞清楚“SPTP通信编排的再利用”再读Ulysses“原来注意力也能切”第三篇读Ring Attention“原来还能这样传KV”有时间再看Striped Attention和系统综述。按这个顺序读不啰嗦概念也不容易绕晕。5.2 实践中的常见问题与排查技巧读论文是一回事代码跑起来又是另一回事。我踩过的坑和排查思路整理一下异常表现可能原因排查方向开了SP后显存反而涨没关掉activation recomputation或SP切分和TP的AllReduce重复缓存检查临时张量生命周期确认LayerNorm输入是否只保留1/P分片通信耗时暴增TPSP在跨节点拓扑下使用全互联算子用NCCL的拓扑感知确认通信组是否跨节点必要时改为Ring执行序训练loss异常漂移分布式softmax合并时用了局部max或直接平均按全局max流程重写合并用数值稳定性测试比对单卡结果长序列训练吞吐上不去计算与通信没有重叠在Ring Attention里把KV块的发送、接收和当前块的注意力计算做成异步流水线decode阶段延迟极高CP并行度太高导致每步同步开销过大减小CP并行度或把CP与更大的batch绑定来摊销同步成本负载明显不均衡因果mask让靠后块计算量过大改用Striped Attention式的条纹KV分配5.3 我的一点实操体会最后聊几句我自己跑了这些策略后的体验。第一别一上来就追求把所有并行度都拉满。我刚开始做超长序列时把TP、PP、SP、CP全开结果通信开销和调度复杂性把吞吐拖垮了显存倒是省了时间却翻倍。后来退回“TPSP”的稳妥组合先把激活显存压住再按需加CP简单且可控。第二一定要先做小规模数值验证。每次改并行切分我都习惯先在2K序列长度、单机多卡上跑一个对照实验保证SP/CP开启前后的输出和梯度几乎一致。因为这类并行策略的bug往往不会直接报错而是静默地污染数值结果一旦上了大序列再排查成本极高。第三对自己的硬件拓扑要有清醒认知。我看到太多人拿着Ulysses的方案往廉价的万兆以太网机器上压结果All-to-All把网络打爆然后回头骂论文是“骗子”。其实不是论文不对是选型没匹配硬件。SP/CP的效果高度依赖互联带宽这一点比模型结构的选择更硬。这些策略本身还在快速演进很多新的负载均衡、异步通信、分块稀疏注意力技巧正在不断落地。如果你正在做长上下文方向建议盯住这篇系统综述后面的引用链按作者和机构去追最新的实现对比能少走很多弯路。
返回列表