Skip to content
🔗 分享本题
查看我的学习进度 →

训练显存账本、OOM 诊断和参数激活状态分片优化图

🧠 记忆锚点:OOM 先定位谁占显存;减批量不够时,再从激活、状态、参数和分片逐层处理。

💡 答案要点

显存占用分析:

总显存 = 模型 + 梯度 + 优化器状态 + 激活值 + 缓存

示例(7B 模型,FP16):
  模型:14GB
  梯度:14GB
  优化器(Adam):28GB
  激活值:10-20GB(取决于 batch size)
  总计:66-76GB

解决方案:

方法显存节省速度影响实现难度
梯度累积50-80%
混合精度(FP16)50%+20%
梯度检查点30-40%-20%⭐⭐
DeepSpeed ZeRO75-90%-10%⭐⭐⭐
LoRA/QLoRA80-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 模型。"

📚 参考:Hugging Face Accelerate(显存不足的分布式方案)