
🧠 记忆锚点:OOM 先定位谁占显存;减批量不够时,再从激活、状态、参数和分片逐层处理。
💡 答案要点
显存占用分析:
总显存 = 模型 + 梯度 + 优化器状态 + 激活值 + 缓存
示例(7B 模型,FP16):
模型:14GB
梯度:14GB
优化器(Adam):28GB
激活值:10-20GB(取决于 batch size)
总计:66-76GB解决方案:
| 方法 | 显存节省 | 速度影响 | 实现难度 |
|---|---|---|---|
| 梯度累积 | 50-80% | 无 | ⭐ |
| 混合精度(FP16) | 50% | +20% | ⭐ |
| 梯度检查点 | 30-40% | -20% | ⭐⭐ |
| DeepSpeed ZeRO | 75-90% | -10% | ⭐⭐⭐ |
| LoRA/QLoRA | 80-95% | 无 | ⭐⭐ |
| 量化(8bit/4bit) | 75% | -15% | ⭐⭐ |
1. 梯度累积(Gradient Accumulation):
python
# 原来:batch_size=32,一次性计算
loss = model(batch_32)
loss.backward()
# 改进:分 4 次,每次 batch_size=8
accumulation_steps = 4
for micro_batch in split_batch(batch_32, 4):
loss = model(micro_batch) / accumulation_steps
loss.backward() # 梯度累积,不更新
optimizer.step() # 累积 4 次后统一更新2. 梯度检查点(Gradient Checkpointing):
python
# 不保存中间激活值,需要时重新计算
model.gradient_checkpointing_enable()
# 代价:训练时间增加 20%
# 收益:显存减少 30-40%3. DeepSpeed ZeRO:
ZeRO-1:分片优化器状态(节省 4x)
ZeRO-2:分片梯度(节省 8x)
ZeRO-3:分片模型参数(节省 N x,N=GPU数)综合方案(7B 模型,单卡 A100 40GB):
python
# 配置
model = AutoModelForCausalLM.from_pretrained(
"llama-7b",
load_in_4bit=True, # 4bit 量化
bnb_4bit_compute_dtype=torch.float16,
)
# LoRA
peft_config = LoraConfig(r=8, lora_alpha=16)
# 训练参数
training_args = TrainingArguments(
per_device_train_batch_size=1, # 小 batch
gradient_accumulation_steps=16, # 累积梯度
gradient_checkpointing=True, # 梯度检查点
fp16=True, # 混合精度
)
# 结果:显存占用 ~25GB,可以训练面试话术:
"遇到 OOM,我的解决流程是:先开梯度累积和 FP16(几乎无损),还不够就用 LoRA(轻量),实在不行就 QLoRA(最省)。曾经在单卡 24GB 上微调 13B 模型。"