内容提要
本文介绍训练大模型时显存管理的通用记账方法。显存分为五类:权重、梯度、优化器状态、激活和通信缓冲。前三类由参数量和并行度决定,激活随输入长度增长,优化空间最大。省显存技巧本质是搬走某行并付出代价,如重算付FLOPs、量化付精度、offload付带宽。激活显存呈阶梯状,可跨rank搬移。静态行可借ZeRO调整,通信缓冲需避免按最坏情况预留。
延伸解读
显存分类与优化空间
训练时显存主要分为五类:权重、梯度、优化器状态、激活和通信缓冲。前三类由模型参数量和并行策略决定,相对静态;激活则随输入长度和流水线在途份数动态变化,是优化空间最大的部分。理解各类显存的特性,有助于针对性地选择优化手段。
省显存技巧的本质:搬走与付费
各种省显存技巧本质上都是将某类显存“搬走”,并付出相应代价:重算消耗额外计算量(FLOPs),量化降低精度,offload则占用PCIe或网络带宽。这些代价需根据硬件能力和模型特性权衡,例如带宽不足时,offload可能导致计算等待,此时重算或减小micro-batch可能更合适。
激活显存的阶梯效应与跨卡搬移
在流水线并行中,激活显存呈阶梯状分布:前面的rank占用多,后面的rank空闲。利用这一特性,可将前面rank的激活远程offload到后面rank的显存中,实现负载均衡。这种方式比CPU offload更快,但需消耗卡间网络带宽,适合带宽充足的环境。
通信缓冲的静态化设计
通信缓冲常被忽视,尤其在MoE模型中,all-to-all通信若按最坏情况预留,会浪费大量显存。通过设计使通信shape静态化,例如确保每个rank接收相同数量的token,可将缓冲缩至固定大小。这提示在模型设计阶段就应考虑通信模式,避免不必要的显存开销。
Q&A
训练大模型时,一张显卡的显存主要被哪几类数据占用?
训练时一张卡的显存主要分为五类:权重、梯度、优化器状态、激活和通信缓冲。其中前三类由参数量和并行度决定,激活由micro-batch长度和流水线在途份数决定,通信缓冲则是all-to-all、all-gather等操作的临时空间。
激活显存的大小如何计算?为什么它随输入长度增长?
激活显存的计算公式为:激活显存 = a × S × (1/P) × n_在途 × 每值字节数。其中a是每token保存的激活值数量,S是micro-batch的token数,1/P表示该卡只持有1/P的层,n_在途是流水线中同时存在的micro-batch份数。因为反向传播需要前向的中间结果,而输入序列越长,保存的中间激活就越多,所以激活显存随输入长度增长。
常见的省显存技巧有哪些?它们分别付出什么代价?
常见的省显存技巧包括:重算(反向时重新计算前向,付出FLOPs代价)、量化(减少字节数,付出精度代价)、offload(将张量搬到CPU或别的卡,付出PCIe或网络带宽代价)。此外,还可以利用激活的阶梯形状在流水线rank之间搬移,付出卡间网络带宽。
为什么激活显存呈阶梯状?如何利用这一点优化显存?
在流水线并行中,采用1F1B调度时,前面的rank(如rank 0)持有的在途micro-batch份数多,后面的rank少,因此激活显存呈阶梯状:前面的rank满,后面的rank空。可以利用这一点,将前面rank的激活远程offload到后面rank的空闲显存中,拉平各rank的显存占用,这比offload到CPU更快。
通信缓冲为什么容易被忽视?如何减小通信缓冲的显存占用?
通信缓冲容易被忽视是因为它不像其他类别那样直观,但MoE的all-to-all通信需要预留接收缓冲,如果每个rank会收到多少token事先未知,就得按最坏情况预留,导致缓冲大小可能达到R倍。任何能让通信shape静态化的设计(如让每个rank恰好收到相同份数的token)都能将通信缓冲缩到固定大小。
面对显存不足,应该按什么顺序排查和优化?
面对显存不足,可以按以下顺序排查:1. 静态三行(权重、梯度、优化器状态)是多少?能否靠ZeRO和并行度压到目标以下?2. 激活是多少?哪一层最大?能否重算、量化或offload?3. 流水线的阶梯有多陡?前面的rank能否借后面rank的显存?4. 通信缓冲是否按最坏情况预留?