显存瓶颈在训练或推理大模型时经常是最先撞到的一面墙。GPU算力明明还有富余,显存却先被参数、梯度和中间激活塞满,导致批次只能一降再降,甚至模型根本启动不了。要突破这个限制,不能只等着换更大显存的显卡,更需要从内存管理和模型卸载两个方向同时入手,减少常驻显存、把暂时不用的数据换到成本更低的存储中去。

显存都消耗在哪里:先拆解占用模型
模型训练时的显存占用可以拆成四块:参数、梯度、优化器状态和激活值。参数本身在推理阶段就存在,一个70亿参数的模型如果使用FP32存储,仅参数就占掉约28GB。训练阶段还需要保存同样大小的梯度,又是约28GB。Adam优化器还要为每个参数维护一阶动量和二阶方差,这两项合计是参数显存的2倍,因此70亿参数模型仅参数、梯度、优化器三项就接近112GB。这里还没有计算前向传播过程中每一层产生的激活值,激活值大小与批次、序列长度、隐藏层宽度直接挂钩。批次为1的短文本可能只占几个GB,长文本或者大batch下会迅速膨胀,成为显存不足的主要来源。理解这个构成是后续优化的前提,几乎所有显存优化手段都是在减少其中某一块。
可以用一个简单的估算函数做显存预算,下面这段代码按精度、是否使用Adam优化器和激活值占用做粗略计算。实际工程里还要叠加通信缓冲区、临时张量和框架自身开销,但作为显存规划的第一版已经足够。估算时还要注意,数据并行会在每张卡上复制参数、梯度和优化器状态,只有模型并行或分片策略才能把这些状态分摊到多张卡。
def estimate_memory(num_params, bytes_per_param=4, use_adam=True, activation_gb=0):
param_gb = num_params * bytes_per_param / 1e9
grad_gb = param_gb
if use_adam:
optimizer_gb = param_gb * 2
else:
optimizer_gb = 0
total_gb = param_gb + grad_gb + optimizer_gb + activation_gb
return total_gb
# 70亿参数FP32训练,Adam优化器,未包含激活值
print(estimate_memory(7e9, 4, True, 0))
推理场景也有显存压力,主要来源是参数常驻和KV Cache。长序列、多并发时KV Cache增长非常快,它的管理与卸载同样是内存优化的重要延伸。先厘清训练和推理的差异,才能判断当前应该先压缩激活值,还是先把参数和优化器状态移出显存。
内存管理:把常驻和瞬时占用都压下来
激活重计算也被称为梯度检查点,是压缩训练显存最直接的手段之一。前向传播默认会保存所有中间激活供反向使用,这一块最容易随着批次和序列长度急剧膨胀。开启梯度检查点后,只保留部分层的激活,反向传播时从最近的检查点重新执行前向计算,临时重建缺失的激活。这样做通常会用20%到30%的额外算力换回大量显存,在Transformer结构上尤其划算,因为注意力层的中间矩阵往往很大,而矩阵乘法重算相对便宜。PyTorch中可以用 torch.utils.checkpoint.checkpoint 包裹需要重计算的层。重计算并不是无条件最优,如果某些层的激活本身很小、重计算成本又很高,就应该保留这些激活。
混合精度训练是另一个关键策略。FP16或BF16会让前向和反向过程中的激活、梯度减半,而FP32主参数副本仍然被保留用于优化器更新。PyTorch的自动混合精度通过 torch.cuda.amp.autocast 和 GradScaler 实现,既降低显存又提高计算速度。BF16的指数范围与FP32一致,更不容易上溢,适合大模型训练;FP16则需要通过动态损失缩放来避免小梯度下溢。开启混合精度后再配合梯度累积,可以在显存约束下模拟更大的全局批次,减少因单卡batch过小造成的训练不稳定。
import torch
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
model = model.cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
for data, target in dataloader:
optimizer.zero_grad()
with autocast():
output = model(data)
loss = loss_fn(output, target)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
注意力计算本身也有很大的优化空间。标准注意力会生成形状为批次数×头数×序列长度×序列长度的分数矩阵,长序列时这个矩阵非常庞大。FlashAttention等融合内核避免显式生成完整矩阵,通过分块计算并利用高速SRAM,把注意力显存从平方级降到接近线性,同时减少读写延迟。它不改变数学结果,只改变计算过程,与混合精度和梯度检查点可以叠加使用。较新的PyTorch版本可以通过 torch.nn.functional.scaled_dot_product_attention 自动选择高效实现,调用方不需要大幅改动代码。
模型卸载:把放不下的部分挪到CPU或磁盘
当内存管理手段用尽后,显存仍然可能不够,因为参数、梯度和优化器状态这三块的体量被模型尺寸锁死。模型卸载的思路很直接:把暂时不参与计算的部分放到GPU之外。CPU内存容量通常比GPU显存大一个数量级,成本也更低。显存中只保留当前层需要的参数和激活,CPU负责存放完整参数、优化器状态或梯度,计算时按需搬回GPU。这个过程会增加PCIe传输开销,所以卸载粒度非常关键。如果按整层粒度卸载,计算第L层时再加载该层参数,传输延迟有机会与前向计算重叠;如果频繁做小块搬移,吞吐会急剧下降。
DeepSpeed的ZeRO策略可以理解为从分片到卸载的渐进方案。Stage 1分片优化器状态,Stage 2进一步分片梯度,Stage 3再分片参数。ZeRO-Offload在Stage 2或Stage 3基础上把优化器状态和梯度下放到CPU,GPU只保留前向反向必需的参数和激活。ZeRO-Infinity则支持把参数卸载到NVMe磁盘,适合单卡甚至消费级设备。配置上通常先启用Stage 2,如果仍然OOM再上Stage 3和参数卸载,因为Stage越深,通信复杂度越高。
{
"train_batch_size": 16,
"gradient_accumulation_steps": 4,
"fp16": {
"enabled": true,
"loss_scale": 0
},
"zero_optimization": {
"stage": 3,
"offload_optimizer": {
"device": "cpu",
"pin_memory": true
},
"offload_param": {
"device": "cpu",
"pin_memory": true
},
"stage3_max_live_parameters": 100000000,
"stage3_max_reuse_distance": 100000000
}
}
除了DeepSpeed,PyTorch FSDP也支持CPU offload,Hugging Face Accelerate的 device_map 可以自动切分模型到不同设备。推理场景下,llama.cpp、ExLlama等工具可以把部分层放在CPU内存,实现消费级设备运行大模型。卸载的本质是用带宽换空间。PCIe Gen4 x16单向带宽约32GB/s,与GPU内部带宽差距悬殊,所以卸载不适合每步大量随机访问。要提高效率,应尽量让数据流按层顺序推进,减少换入换出,并避免把频繁更新的优化器状态放到NVMe磁盘。
实战顺序与调优建议
调优时建议按照先测量、后压缩、再卸载的顺序推进。先用 torch.cuda.memory_allocated 或 nvidia-smi 观察显存是在哪个阶段突然飙升。如果激活值明显偏大,优先考虑降低batch、开启梯度检查点或使用FlashAttention。如果优化器状态占大头,直接启用混合精度和ZeRO Stage 1/2。如果仍然OOM,再考虑参数或优化器卸载。盲目先上卸载可能引入大量PCIe延迟,而真正的显存压力只是激活值造成的话,重计算往往比卸载更划算。
import torch
def print_memory(tag):
allocated = torch.cuda.memory_allocated() / 1024 / 1024
reserved = torch.cuda.memory_reserved() / 1024 / 1024
print(f"{tag}: allocated={allocated:.1f}MB reserved={reserved:.1f}MB")
print_memory("before forward")
# 在前向、反向、step前后分别调用,观察显存峰值变化
print_memory("after forward")
针对不同瓶颈可以套用一个简单判断:激活值占大头时,优先使用FlashAttention、梯度检查点和更小的序列长度;参数和优化器状态占大头时,优先使用混合精度、ZeRO Stage 2/3和CPU offload;想用单卡跑超大模型时,再考虑ZeRO-Infinity或者量化与卸载结合。不同模型结构对激活值的敏感度不同,视觉Transformer的长图像块序列与语言模型的长文本类似,都应先处理注意力中间激活,而不是一开始就急着卸载参数。
模型卸载还会明显影响训练速度。CPU卸载后单步时间通常增加30%到2倍不等,取决于通信与计算重叠程度。可以适当增大梯度累积步数,让每次优化器更新前在GPU上保留更少的优化器状态,减少CPU与GPU之间的优化器同步频率。同时要保证CPU内存足够大,优化器状态一旦移出GPU,CPU内存占用可能比原显存占用还要高。如果使用NVMe卸载,要注意磁盘顺序读写和随机读写之间的巨大差异,尽量使用高速NVMe并控制写入放大。
显存优化不是一次性配置,而是一个动态权衡过程。训练开始时显存够用不代表训练过程中不会增长,某些动态结构、缓存增长以及梯度累积中梯度延迟释放都可能造成后期OOM。建议在配置里预留10%到15%的显存余量,并持续监控峰值内存。把内存管理和模型卸载结合起来使用,才更有可能在有限硬件上稳定运行更大的模型。