拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

训练侧显存测量与优化:从账单拆解到预算决策实战

训练侧显存测量与优化:从账单拆解到预算决策实战

1. 训练侧显存测量到底在测什么

显存优化这件事,很多人一上来就想着怎么省,结果省了半天发现根本没省到点子上。问题出在哪?出在没搞清楚显存到底被谁吃掉了。训练侧的显存测量,核心目标就一个:把显存账单拆开,看清楚每一笔开销的去向,然后才能做预算决策。

我见过太多人拿着nvidia-smi看一眼显存占用,发现快满了就开始慌,然后盲目上梯度检查点、盲目降batch size,最后训练速度掉了一半,显存也没省下多少。这种做法的问题在于,nvidia-smi看到的只是一个总数,它不会告诉你这8GB里有3GB是模型参数、2GB是优化器状态、1.5GB是激活值、剩下的是碎片和临时缓冲区。你不知道钱花在哪,就没办法做预算。

训练侧显存测量要回答的问题很具体:模型参数占多少、梯度占多少、优化器状态占多少、激活值占多少、临时缓冲区占多少。这五块加起来才是真实账单。而且这五块的比例关系会随着模型规模、batch size、序列长度、精度策略的变化而剧烈变化。一个7B模型在FP16下参数占14GB,但如果用Adam优化器,优化器状态就要占28GB(FP32的一阶动量和二阶动量各14GB),再加上梯度14GB,光这三项就56GB了,还没算激活值。所以为什么大家说7B模型全量微调至少要80GB显存,账就是这么算出来的。

测量方法上,最直接的是用PyTorch的显存分析工具。torch.cuda.memory_allocated()能拿到当前分配的显存,torch.cuda.max_memory_allocated()能拿到峰值。但这两个数字只反映PyTorch分配器层面的情况,不包括CUDA上下文本身占用的那几百MB。更细的拆解需要用torch.cuda.memory_summary(),它会按分配块大小分类列出。不过这个输出比较原始,我一般会自己写个hook,在模型forward前后、backward前后分别打点,算出每个阶段的增量。

还有一个容易被忽略的点:显存碎片。PyTorch的缓存分配器会预留一些显存不还给系统,导致nvidia-smi看到的占用比memory_allocated()高不少。这个差值在长时间训练中会逐渐增大,尤其是当你有动态shape的输入时。测量的时候要把这个差值也记录下来,否则你按memory_allocated()做的预算到了实际训练时就会OOM。

注意:测量一定要在真实训练循环里做,不能只跑一个forward就完事。因为激活值的峰值出现在backward阶段,只测forward会严重低估。

2. 显存账单的五大部分与计算逻辑

2.1 模型参数与梯度的显存占用

模型参数的显存占用最好算:参数量乘以每个参数的字节数。FP32是4字节,FP16和BF16是2字节,INT8是1字节。一个7B模型在FP16下就是7B乘2等于14GB。梯度占用的字节数和参数一致,因为梯度需要和参数同样的精度来保证更新时不丢信息。所以FP16训练时,参数加梯度就是28GB。

但这里有个坑:很多人以为用FP16训练,参数就是FP16。实际上PyTorch的AMP(自动混合精度)会保留一份FP32的master weight。也就是说,参数实际上占了两份:一份FP16用于前向和反向计算,一份FP32用于优化器更新。这样算下来,7B模型的参数占用是14GB(FP16)加28GB(FP32 master),总共42GB。梯度也是FP16一份,但优化器更新时需要FP32的梯度,所以梯度实际占用也是14GB(FP16)加28GB(FP32),又是42GB。这就是为什么AMP能省显存但省不了太多——它省的是激活值和部分计算缓冲区,参数和优化器状态的大头省不掉。

2.2 优化器状态的显存黑洞

优化器状态是显存占用里最容易被低估的部分。以Adam为例,它需要为每个参数维护一阶动量m和二阶动量v,都是FP32。所以优化器状态的显存等于参数量乘以4字节乘以2,也就是参数量乘以8字节。7B模型就是56GB。加上前面的参数和梯度,已经98GB了。这就是为什么全量微调7B模型至少需要8张80GB的卡来做数据并行——单卡根本放不下。

AdamW稍微好一点,它把权重衰减和梯度更新解耦了,但动量部分还是一样的。Adafactor通过分解二阶动量矩阵来省显存,能把优化器状态降到参数量乘以4字节左右,但收敛性会受一些影响。Sophia优化器用对角Hessian估计来替代二阶动量,也能省不少,但实现复杂度和调参难度都上去了。

实际做预算的时候,我一般按这个公式估算:总显存等于参数量乘以(2加2加8)加激活值加缓冲区。前面的2是FP16参数,第二个2是FP16梯度,8是Adam的优化器状态。这个公式在AMP加AdamW的场景下比较准,误差在10%以内。

2.3 激活值的动态波动

激活值是显存占用里最动态的部分,它跟batch size、序列长度、模型层数、隐藏维度都相关。粗略估算的话,激活值显存约等于batch size乘以序列长度乘以隐藏维度乘以层数乘以一个系数。这个系数取决于具体的网络结构,Transformer里主要是注意力矩阵和FFN的中间激活。

注意力矩阵的显存是序列长度的平方乘以batch size乘以头数乘以2字节。序列长度2048时,这个平方项就是4M,乘以batch size和头数后很容易上GB。序列长度拉到8192,平方项变成64M,直接爆炸。这就是长上下文训练显存吃紧的核心原因。

FFN的中间激活是batch size乘以序列长度乘以4倍隐藏维度乘以2字节。这个和序列长度是线性关系,比注意力矩阵温和一些。但层数一多,累积起来也很可观。

梯度检查点(gradient checkpointing)就是针对激活值的优化手段。它不保存中间激活,而是在backward时重新计算。代价是计算量增加约30%,但激活值显存能降到原来的平方根级别。对于层数很深的模型,这个 trade-off 非常划算。

2.4 临时缓冲区与碎片

临时缓冲区包括CUDA kernel执行时的workspace、通信操作的缓冲区、以及PyTorch分配器的预留空间。这部分很难精确测量,但可以通过对比nvidia-smi和memory_allocated()的差值来估算。一般来说,这个差值在500MB到2GB之间,模型越大、并行策略越复杂,差值越大。

碎片问题在动态shape场景下特别严重。比如你训练时序列长度不固定,PyTorch的缓存分配器会按最大shape预留块,导致实际占用远高于理论值。解决办法是设置PYTORCH_CUDA_ALLOC_CONF环境变量,启用expandable_segments,让分配器能合并碎片。这个设置在新版PyTorch里效果很明显,我实测能把碎片率从20%降到5%以下。

3. 预算决策:从测量结果到优化方案

3.1 显存预算的分配原则

测完显存账单后,下一步是做预算决策。预算的核心原则是:先保训练稳定性,再保吞吐量,最后才考虑省显存。很多人搞反了顺序,为了省显存把batch size降到1,结果训练不稳定,收敛慢得要命,省下来的显存也没换来什么好处。

我的预算分配一般是这样的:参数和优化器状态是刚性支出,没法省,必须留足。激活值是弹性支出,可以通过梯度检查点、序列并行、FlashAttention等手段压缩。临时缓冲区留10%到15%的余量,防止峰值OOM。

具体到数字上,假设你有80GB显存,7B模型AMP加AdamW,参数加梯度加优化器状态大约98GB,单卡放不下。这时候你有几个选择:一是上ZeRO Stage 2,把优化器状态和梯度分片到多卡,每卡只存一部分;二是上LoRA,只训练低秩适配器,参数量降到原来的百分之一;三是上QLoRA,把基座模型量化到4bit,进一步压缩。

3.2 ZeRO与FSDP的显存账

ZeRO Stage 1只分片优化器状态,每卡显存等于参数量乘以2加梯度乘以2加优化器状态除以N。7B模型8卡,每卡优化器状态7GB,加上参数14GB和梯度14GB,总共35GB,80GB卡能放下。Stage 2再分片梯度,每卡梯度降到1.75GB,总共约23GB。Stage 3分片参数,每卡参数1.75GB,总共约10GB,但通信开销大幅增加。

FSDP本质上是ZeRO Stage 3的PyTorch原生实现,它把参数、梯度、优化器状态都分片,每卡只存1/N。但FSDP在forward和backward时需要all-gather参数,通信量很大。实际用下来,FSDP在8卡A100上训练7B模型,每卡显存约12GB,但吞吐量比DDP低20%左右。这个 trade-off 要看你的瓶颈是显存还是算力。

3.3 AMP的显存收益与代价

AMP(自动混合精度)是显存优化的第一板斧,但它的收益经常被高估。AMP省的主要是激活值和部分计算缓冲区,参数和优化器状态的大头省不掉。实测下来,AMP能把激活值显存降到FP32的60%左右,总体显存节省约20%到30%。

AMP的代价是数值稳定性。FP16的动态范围窄,梯度容易下溢或上溢。PyTorch的GradScaler通过动态调整loss scale来缓解这个问题,但在某些模型结构上仍然会出现NaN。BF16的动态范围和FP32一样,不需要loss scaling,但精度比FP16低。A100及以上支持BF16,V100只支持FP16。选哪个取决于你的硬件和模型对精度的敏感度。

实操心得:用AMP时,把LayerNorm和softmax强制保留在FP32,能显著提升训练稳定性。PyTorch的torch.cuda.amp.autocast默认已经这么做了,但如果你自己写kernel,要注意手动指定。

4. 实操测量流程与工具链

4.1 测量脚本的编写要点

我一般会写一个独立的测量脚本,不掺在训练代码里,这样干净、可复现。脚本的核心结构是:初始化模型和优化器,构造一个真实batch的输入,然后分阶段打点。

import torch from torch.cuda import memory_allocated, max_memory_allocated, reset_peak_memory_stats def measure(model, optimizer, input_ids, labels): reset_peak_memory_stats() base = memory_allocated() # forward outputs = model(input_ids, labels=labels) loss = outputs.loss after_forward = memory_allocated() # backward loss.backward() after_backward = memory_allocated() # optimizer step optimizer.step() optimizer.zero_grad() after_step = memory_allocated() return { 'base': base, 'forward_delta': after_forward - base, 'backward_delta': after_backward - after_forward, 'step_delta': after_step - after_backward, 'peak': max_memory_allocated() }

这个脚本跑一次就能拿到四个关键数字。forward_delta主要是激活值,backward_delta是梯度加激活值的峰值,step_delta是优化器状态的增量。peak是整个过程的最大值,用来做OOM判断。

4.2 不同配置的对比测量

测量不能只测一个配置,要测一组配置做对比。我一般会测这几组:FP32基线、AMP、AMP加梯度检查点、AMP加梯度检查点加ZeRO。每组跑三次取平均,排除冷启动的影响。

对比的时候重点看两个指标:峰值显存和吞吐量。峰值显存决定你能不能跑起来,吞吐量决定你跑得多快。有时候省显存的方案会把吞吐量砍半,这时候就要算一笔账:省下来的显存能不能换来更大的batch size,如果能,吞吐量可能反而更高。

举个例子,7B模型FP32训练峰值显存60GB,吞吐量100 samples/s。AMP后峰值45GB,吞吐量130 samples/s。AMP加梯度检查点后峰值30GB,吞吐量90 samples/s。虽然梯度检查点让吞吐量降了,但省下的15GB显存可以让你把batch size翻倍,实际吞吐量变成180 samples/s。这就是预算决策的价值。

4.3 显存碎片的手动清理

测量过程中如果发现nvidia-smi和memory_allocated()差值越来越大,说明碎片在累积。这时候可以手动调torch.cuda.empty_cache(),但注意这个操作会释放缓存分配器预留的块,可能导致后续分配变慢。更好的办法是设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,让分配器自己管理碎片。

还有一个技巧是在训练循环里定期调用torch.cuda.reset_peak_memory_stats(),把峰值统计重置,这样能更准确地看到每个step的显存波动。如果不重置,峰值会一直保留历史最大值,掩盖了后期的显存增长。

5. 常见问题与排查技巧实录

5.1 测量结果与nvidia-smi对不上

这是最常见的问题。memory_allocated()只统计PyTorch分配器管理的显存,不包括CUDA上下文、cuDNN workspace、NCCL通信缓冲区。这些加起来可能有1到2GB。所以nvidia-smi看到的数字总是比memory_allocated()大。

解决办法是用torch.cuda.memory_snapshot()拿到完整的内存快照,它会列出所有分配块,包括那些不在PyTorch管理范围内的。不过这个输出很大,一般只在排查问题时用。

5.2 激活值峰值出现在意想不到的地方

有时候你会发现backward阶段的显存峰值比forward高很多,甚至高出一倍。这通常是因为某些操作的backward需要保存forward的中间结果,而这些结果在forward结束后并没有释放。比如attention的softmax输出,forward时算完就丢了,但backward需要它来计算梯度,所以PyTorch会把它保留到backward。

排查方法是给每个模块注册forward hook和backward hook,记录每个模块的输入输出显存。这样能精确定位到哪个模块的backward显存开销异常。

5.3 OOM发生在optimizer.step()

optimizer.step()本身不分配大块显存,但它会触发参数更新,而参数更新需要读取梯度、写入参数。如果梯度是FP16而参数是FP32,这里会有一个隐式的类型转换,产生临时缓冲区。Adam的动量更新还会产生中间变量。这些加起来可能几百MB,在显存已经接近上限时就是压死骆驼的最后一根稻草。

解决办法是在step之前手动torch.cuda.empty_cache(),或者把优化器状态分片到多卡。另一个办法是用foreach实现的优化器,它把多个参数的更新合并成一个kernel,减少临时缓冲区的分配次数。

5.4 梯度检查点与AMP的兼容问题

梯度检查点在backward时会重新计算forward,如果和AMP一起用,重计算时的精度策略要和原forward一致,否则梯度会对不上。PyTorch的torch.utils.checkpoint默认会保留AMP的autocast状态,但如果你自己实现了checkpoint逻辑,要注意手动传递autocast上下文。

还有一个坑是梯度检查点不能和某些自定义autograd函数一起用,因为重计算时这些函数的forward可能不是确定性的。遇到这种情况,要么把自定义函数排除在checkpoint范围外,要么确保它是确定性的。

5.5 多卡训练时的显存不均衡

数据并行时,每卡的显存占用应该基本一致。如果发现某张卡显存特别高,通常是数据加载不均衡或者通信操作卡住了。检查DataLoader的num_workers和pin_memory设置,确保每个进程拿到的batch大小一致。通信方面,NCCL的all-reduce是同步操作,如果某张卡算得慢,其他卡会在通信点等待,显存占用会暂时升高。

排查方法是打印每张卡的memory_allocated(),看差异是否超过5%。如果超过,检查数据分片逻辑和通信配置。

问题现象可能原因排查方法解决手段
nvidia-smi比memory_allocated高2GBCUDA上下文和通信缓冲区memory_snapshot预留2GB余量
backward峰值远高于forward中间激活未释放模块级hook梯度检查点
step时OOM类型转换临时缓冲区逐步打点foreach优化器
多卡显存不均衡数据或通信不均衡逐卡打印调整DataLoader
碎片率持续增长动态shape监控差值expandable_segments

6. 预算决策的实战案例

6.1 单卡24GB训练7B模型的可行性分析

有人问单卡24GB能不能训7B模型。按前面的公式算,7B模型AMP加AdamW,参数14GB(FP16)加梯度14GB(FP16)加优化器状态56GB(FP32),总共84GB,远超24GB。所以全量微调不可能。

但LoRA可以。LoRA只训练低秩矩阵,参数量通常是原模型的0.1%到1%。7B模型的LoRA参数量约7M到70M,FP16下占14MB到140MB。优化器状态按Adam算是参数量乘以8,也就是112MB到1.1GB。加上基座模型的14GB(FP16),总共约15GB到16GB。24GB卡能放下,还能留出8GB给激活值和缓冲区。

QLoRA更进一步,把基座模型量化到4bit,7B模型只占3.5GB。加上LoRA参数和优化器状态,总共约5GB。24GB卡能跑得很宽裕,甚至能上更大的batch size。

6.2 多卡场景下的并行策略选择

如果你有4张24GB卡,总共96GB显存,想训7B模型全量微调。DDP每卡都要存完整的参数、梯度、优化器状态,每卡84GB,放不下。ZeRO Stage 2把优化器状态和梯度分片到4卡,每卡参数14GB加梯度3.5GB加优化器状态14GB,总共31.5GB,还是放不下。ZeRO Stage 3把参数也分片,每卡参数3.5GB加梯度3.5GB加优化器状态14GB,总共21GB,勉强能放下,但激活值还没算。

这时候要么上梯度检查点把激活值压到2GB以内,要么上CPU offload把优化器状态放到内存。CPU offload的代价是step速度慢3到5倍,但能省下14GB显存。实测下来,4卡24GB加ZeRO Stage 3加梯度检查点加CPU offload,能训7B模型,但吞吐量只有DDP的十分之一。所以如果追求速度,还是建议上80GB卡。

6.3 预算决策的决策树

我把预算决策整理成一个简单的决策树,方便快速判断:

第一步,算刚性支出:参数量乘以(2加2加8)除以并行度。如果这个数字超过单卡显存的80%,考虑LoRA或QLoRA。

第二步,算激活值:batch size乘以序列长度乘以隐藏维度乘以层数乘以系数。如果超过剩余显存的50%,上梯度检查点。

第三步,算碎片余量:留10%到15%的显存给临时缓冲区和碎片。如果不够,调expandable_segments或减小batch size。

第四步,测吞吐量:在满足显存约束的前提下,找吞吐量最大的配置。有时候大batch加梯度检查点比小batch不加检查点更快。

这个决策树不是绝对的,但能帮你快速缩小选择范围。实际做的时候还是要跑测量脚本验证,因为不同模型结构、不同框架版本的显存行为差异很大。

最后分享一个小技巧:在训练脚本里加一个显存监控回调,每个step记录峰值显存和吞吐量,输出到TensorBoard。这样你能看到显存随训练进程的变化趋势,及时发现碎片累积或激活值增长的问题。我靠这个回调抓到过好几次隐蔽的显存泄漏,都是自定义层里缓存了不该缓存的东西。

返回列表