「CUDA out of memory」可能是深度学习工程师看到最多的错误信息。很多人的第一反应是换更大的 GPU,但在此之前,你应该先检查是否有显存浪费。
以下 8 个技巧按从简单到复杂排列,通常前 3 个就能解决大部分 OOM 问题。
1. 减小 Batch Size
最简单但也最容易被忽略的方法。Batch size 从 32 减到 16,显存占用几乎减半。
但减小 batch size 会影响训练效果?不一定。配合 gradient accumulation,你可以保持等效的大 batch size:
# 等效 batch_size = 32,但每步只用 8 的显存
accumulation_steps = 4
optimizer.zero_grad()
for i, batch in enumerate(dataloader): # dataloader batch_size=8
loss = model(batch) / accumulation_steps
loss.backward()
if (i + 1) % accumulation_steps == 0:
optimizer.step()
optimizer.zero_grad()
2. 开启混合精度训练
FP16/BF16 混合精度训练将显存占用减少约 40%,同时还能加速训练。
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch in dataloader:
with autocast(dtype=torch.bfloat16):
loss = model(batch)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
BF16 比 FP16 更稳定,如果你的 GPU 支持(A100、H100),优先选 BF16。
3. 开启 Gradient Checkpointing
用计算换显存:前向传播时不保存中间激活值,反向传播时重新计算。显存占用可减少 60-70%,代价是训练速度慢约 20-30%。
model.gradient_checkpointing_enable()
一行代码,HuggingFace Transformers 模型直接支持。
4. 使用 LoRA 而不是全参微调
全参微调 7B 模型需要存储完整的优化器状态(Adam 需要 2 份参数副本),显存需求约 60GB+。LoRA 只训练极少量参数,显存需求降到 15-20GB。
from peft import LoraConfig, get_peft_model
config = LoraConfig(
r=8,
lora_alpha=32,
target_modules=["q_proj", "v_proj"],
lora_dropout=0.05,
)
model = get_peft_model(model, config)
model.print_trainable_parameters()
# 输出:trainable params: 4,194,304 || all params: 6,738,415,616 || trainable%: 0.0622%
5. 优化 DataLoader
DataLoader 的 pin_memory=True 和 num_workers 设置不当都会浪费显存:
pin_memory=True:预分配锁页内存,加速 CPU→GPU 传输,但会额外占用内存num_workers过多:每个 worker 都会预加载数据到内存prefetch_factor过大:预取太多 batch 占用额外显存
dataloader = DataLoader(
dataset,
batch_size=8,
num_workers=4, # 不要设太大
pin_memory=True,
prefetch_factor=2, # 默认值就够了
persistent_workers=True,
)
6. 及时释放不需要的 Tensor
训练循环中临时变量不及时释放,会导致显存碎片化:
# ❌ loss tensor 持续占用显存
losses = []
for batch in dataloader:
loss = model(batch)
losses.append(loss) # loss tensor 带有整个计算图
# ✅ 只保留标量值
losses = []
for batch in dataloader:
loss = model(batch)
losses.append(loss.item()) # .item() 取出 Python 数字,释放计算图
loss.backward()
7. 使用 DeepSpeed ZeRO
DeepSpeed ZeRO 将优化器状态、梯度、参数分片到多个 GPU 上,单卡显存占用大幅降低:
- ZeRO Stage 1:分片优化器状态,显存减少约 4x
- ZeRO Stage 2:+ 分片梯度,显存减少约 8x
- ZeRO Stage 3:+ 分片参数,显存减少约 N 倍(N = GPU 数量)
即使只有一张 GPU,ZeRO-Offload 也能把优化器状态卸载到 CPU 内存。
8. 量化训练(QLoRA)
4-bit 量化 + LoRA 是显存优化的终极方案。7B 模型的显存需求从 60GB+ 降到 6-8GB,一张 RTX 3090 就能跑。
from transformers import BitsAndBytesConfig
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
)
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7b",
quantization_config=bnb_config,
)
总结
显存优化的优先级建议:
- 先试 1-3(零代码或一行代码,效果明显)
- 再试 4(LoRA,如果是微调场景)
- 最后考虑 5-8(需要更多配置和理解)
在 RunWhere 上提交任务前,可以先在「价格」页按 GPU 型号对比各家云厂商的实时报价——显存需求算清楚之后,选刚好够用的机型,不多花一分冤枉钱。