从显存走到效率:AI内存优化笔记

AI内存优化笔记我没按教科书顺序做。

去年搞一个 7B 模型的 fine-tuning 项目时,显存崩溃成了日常。

场景和问题

当时的训练环境是这样的:

  • 硬件:4 × NVIDIA A100 80GB (PCIe)
  • 框架:PyTorch 2.1.0 + transformers 4.35.0
  • 模型:Qwen-7B-Chat,LoRA 微调
  • 训练数据:100万条中文对话样本
  • 目标:训练一个能够高质量回复的聊天模型

原始代码很简单,就是用 Trainer 包装一下模型和数据:

from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
from peft import LoraConfig, get_peft_model

model = AutoModelForCausalLM.from_pretrained(
    "Qwen/Qwen-7B-Chat",
    torch_dtype=torch.float16,
    device_map="auto"
)

tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen-7B-Chat")

peft_config = LoraConfig(
    r=8,
    lora_alpha=16,
    target_modules=["q_proj", "v_proj"],
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM"
)

model = get_peft_model(model, peft_config)

training_args = TrainingArguments(
    output_dir="./qwen-7b-lora",
    num_train_epochs=3,
    per_device_train_batch_size=4,
    gradient_accumulation_steps=4,
    learning_rate=2e-4,
    fp16=True,
    logging_steps=10,
    save_steps=1000,
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
)

trainer.train()

跑起来第一轮就崩了,显存占用 76GB/80GB,离崩溃只差一步。调小 batch size 到 2,然后梯度累积步数调到 8,总算能跑了,但训练时间直接拉长了 4 倍。这不是办法。

梯度检查点

第一个尝试的是梯度检查点(Gradient Checkpointing)。它的原理是:在反向传播时重新计算前向传播的中间激活值,而不是全部存下来。这用计算换空间,典型的时空权衡。

from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
from peft import LoraConfig, get_peft_model

model = AutoModelForCausalLM.from_pretrained(
    "Qwen/Qwen-7B-Chat",
    torch_dtype=torch.float16,
    device_map="auto",
    use_cache=False  # 禁用 KV cache 训练时不推理
)

model.gradient_checkpointing_enable()  # 开启梯度检查点

peft_config = LoraConfig(
    r=8,
    lora_alpha=16,
    target_modules=["q_proj", "v_proj"],
    lora_dropout=0.05,
    bias="none",
    task_type="CAUSAL_LM"
)

model = get_peft_model(model, peft_config)

这个改动很小,但效果立竿见影。显存占用从 76GB 降到了 58GB,batch size 可以恢复到 4 了。

但这里有个坑:梯度检查点需要模型支持,不是所有模型都能直接开。第一次在老版本的 transformers 上用,报了 AttributeError: 'Qwen2Model' object has no attribute 'gradient_checkpointing'。升级到 4.35.0 才解决。

另一个坑是训练时间会变长。反向传播时需要重新计算前向激活值,训练速度大约慢了 15-20%。这个 trade-off 要自己权衡,内存紧张时值得,内存够用时可以考虑不开。

混合精度训练

代码里已经开启了 fp16=True,但可以更进一步用 bf16。A100 支持 BF16,它比 FP16 的数值稳定性更好,而且不需要损失缩放。

training_args = TrainingArguments(
    output_dir="./qwen-7b-lora",
    num_train_epochs=3,
    per_device_train_batch_size=4,
    gradient_accumulation_steps=4,
    learning_rate=2e-4,
    bf16=True,  # 改用 BF16
    logging_steps=10,
    save_steps=1000,
    # 启用梯度检查点
    gradient_checkpointing=True,
    # 优化内存碎片
    optim="adamw_torch_fused",
    # 减少 checkpoint 保存频率
    save_total_limit=2,
)

改完之后显存占用没有明显变化,但训练更稳定了,不会因为数值溢出而崩溃。早期用 FP16 时经常遇到 NaN 损失,BF16 彻底解决了这个问题。

优化数据加载

显存不只是模型占的,数据加载也会占一块。特别是做文本任务时,长句子的 padding 会占用不少空间。

from transformers import DataCollatorForLanguageModeling

data_collator = DataCollatorForLanguageModeling(
    tokenizer=tokenizer,
    mlm=False,
    pad_to_multiple_of=8  # 填充到 8 的倍数,优化张量对齐
)

# 数据预处理时做动态截断
def preprocess_function(examples):
    # 限制最大长度,减少内存压力
    return tokenizer(
        examples["text"],
        truncation=True,
        max_length=2048,  # 根据实际场景调整
        padding=False,  # 不在这里做 padding,交给 data_collator
    )

这里有个容易被忽略的坑:padding=True 在预处理阶段会把所有样本都 pad 到 batch 里最长的那一个,导致大量无用填充。改成 padding=False,让 DataCollatorForLanguageModeling 在 batch 级别做动态 padding,节省不少内存。

深度模型并行

如果前面几步还不够,可以考虑模型并行。这里有几个方案:

ZeRO 优化

DeepSpeed 的 ZeRO 优化把模型参数、梯度、优化器状态切片到不同 GPU 上。

{
  "zero_optimization": {
    "stage": 2,
    "offload_optimizer": {
      "device": "cpu",
      "pin_memory": true
    },
    "offload_param": {
      "device": "cpu",
      "pin_memory": true
    },
    "overlap_comm": true,
    "contiguous_gradients": true,
    "sub_group_size": 1e9,
    "reduce_bucket_size": 5e8,
    "stage3_prefetch_bucket_size": 5e7,
    "stage3_param_persistence_threshold": 1e5,
    "stage3_max_live_parameters": 1e9,
    "stage3_max_reuse_distance": 1e9,
    "stage3_gather_16bit_weights_on_model_save": true
  }
}

ZeRO-2 会把优化器状态和梯度切片,ZeRO-3 甚至会把参数也切片。不过 ZeRO-3 的通信开销比较大,训练速度会慢 30-40%。我的经验是先试 ZeRO-2,不够再考虑 ZeRO-3。

模型分片

from transformers import AutoModelForCausalLM, BitsAndBytesConfig

quantization_config = BitsAndBytesConfig(
    load_in_4bit=True,
    bnb_4bit_compute_dtype=torch.float16,
    bnb_4bit_use_double_quant=True,
    bnb_4bit_quant_type="nf4"
)

model = AutoModelForCausalLM.from_pretrained(
    "Qwen/Qwen-7B-Chat",
    quantization_config=quantization_config,
    device_map="auto"
)

4-bit 量化能把显存占用再降一半,但精度损失是真实的。对于推理场景没问题,但训练时需要评估精度影响。我的经验是:如果是知识性任务(问答、摘要),4-bit 量化可以用;如果是生成性任务(创意写作、代码生成),建议谨慎使用。

监控和调优

优化的前提是知道瓶颈在哪里。可以用 torch.cuda.memory_summary()nvidia-smi 实时监控:

import torch

def print_memory_usage():
    allocated = torch.cuda.memory_allocated() / 1024**3
    reserved = torch.cuda.memory_reserved() / 1024**3
    max_allocated = torch.cuda.max_memory_allocated() / 1024**3
    print(f"Allocated: {allocated:.2f} GB")
    print(f"Reserved: {reserved:.2f} GB")
    print(f"Max Allocated: {max_allocated:.2f} GB")
    torch.cuda.reset_peak_memory_stats()

# 在训练循环里定期调用
for step, batch in enumerate(train_dataloader):
    outputs = model(**batch)
    loss = outputs.loss / gradient_accumulation_steps
    loss.backward()

    if (step + 1) % gradient_accumulation_steps == 0:
        optimizer.step()
        optimizer.zero_grad()
        print_memory_usage()  # 每 N 步打印一次内存占用

这样可以定位到底是参数、梯度还是激活值占了大头。有时候瓶颈不在模型,而在数据加载或者预处理代码写得有问题。

最终效果

折腾完一圈,最后的效果是这样的:

  • 梯度检查点:显存占用降了 24%
  • 混合精度 + BF16:没有省显存,但训练更稳定
  • 动态 padding:数据占用降了 30%
  • ZeRO-2 优化:优化器状态占用降了 40%
  • 4-bit 量化:模型参数占用降了 75%(但最终没采用)

最终方案:梯度检查点 + BF16 + 动态 padding + ZeRO-2。显存占用从最初的 76GB 降到了 42GB,batch size 可以从 4 提升到 8,训练总时间反而比优化前短了。

踩坑总结

  1. 梯度检查点不是万能的:老版本 transformers 不支持,模型代码改动大时需要重新验证。有些自定义层没有正确实现梯度检查点,开了也没用。

  2. 混合精度要看硬件:A100 用 BF16,V100 只能用 FP16,老显卡可能都不支持。强行开启会报错或者训练不稳定。

  3. ZeRO 有通信开销:多卡网络带宽不够时,ZeRO-3 可能比单卡还慢。最好先在测试环境跑一遍,确认性能收益。

  4. 量化要评估任务类型:推理时 4-bit 没问题,训练时会影响收敛。特别是对精度敏感的任务(比如数学计算、代码生成),建议谨慎使用。

  5. 内存碎片要清理:长时间训练后,显存碎片会严重,torch.cuda.empty_cache() 能临时缓解,但更好的做法是定期重启训练进程。

  6. 监控要到位:盲目调参不如先监控,知道瓶颈再优化。很多看似内存问题,其实是数据加载或者预处理写得烂。

原理补几句

梯度检查点的原理是反直觉的:通常理解是前向传播存激活值,反向传播用这些激活值算梯度。但梯度检查点故意只存一部分激活值,剩下的反向时重新算。

这像学生做作业:可以把每一步都记下来,或者只记关键节点,做题时重新推导一遍。前者省脑子但费纸,后者省纸但费脑子。内存就是纸,计算就是脑子。

ZeRO 的本质是把大家都能看到的"公共数据"切片,每个 GPU 只存自己那一份,需要时再通信取。这像多人合作读书,把书拆成几部分每人读一段,要交叉引用时再传阅。

混合精度更直白:有些数不需要那么高精度,用 16 位存就够了。就像数字照片,不是所有像素都要 16-bit 色深,8-bit 肉眼看不出区别,但省了一半空间。

技术大多是 trade-off,没有银弹。内存优化的本质是在计算、通信、精度之间找平衡点,找到那个"刚好够用"的状态。

折腾完这轮,对显存的恐惧少了一些,但对平衡的敬畏多了一分。工具是拿来用的,不是拿来迷信的。知道边界在哪里,才敢在安全区里大胆折腾。

版权声明: 本文首发于 指尖魔法屋-从显存走到效率:AI内存优化笔记https://blog.thinkmoon.cn/post/236-ai-memory-optimization-gpu-efficiency-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!