从显存走到效率: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,训练总时间反而比优化前短了。
踩坑总结
梯度检查点不是万能的:老版本 transformers 不支持,模型代码改动大时需要重新验证。有些自定义层没有正确实现梯度检查点,开了也没用。
混合精度要看硬件:A100 用 BF16,V100 只能用 FP16,老显卡可能都不支持。强行开启会报错或者训练不稳定。
ZeRO 有通信开销:多卡网络带宽不够时,ZeRO-3 可能比单卡还慢。最好先在测试环境跑一遍,确认性能收益。
量化要评估任务类型:推理时 4-bit 没问题,训练时会影响收敛。特别是对精度敏感的任务(比如数学计算、代码生成),建议谨慎使用。
内存碎片要清理:长时间训练后,显存碎片会严重,
torch.cuda.empty_cache()能临时缓解,但更好的做法是定期重启训练进程。监控要到位:盲目调参不如先监控,知道瓶颈再优化。很多看似内存问题,其实是数据加载或者预处理写得烂。
原理补几句
梯度检查点的原理是反直觉的:通常理解是前向传播存激活值,反向传播用这些激活值算梯度。但梯度检查点故意只存一部分激活值,剩下的反向时重新算。
这像学生做作业:可以把每一步都记下来,或者只记关键节点,做题时重新推导一遍。前者省脑子但费纸,后者省纸但费脑子。内存就是纸,计算就是脑子。
ZeRO 的本质是把大家都能看到的"公共数据"切片,每个 GPU 只存自己那一份,需要时再通信取。这像多人合作读书,把书拆成几部分每人读一段,要交叉引用时再传阅。
混合精度更直白:有些数不需要那么高精度,用 16 位存就够了。就像数字照片,不是所有像素都要 16-bit 色深,8-bit 肉眼看不出区别,但省了一半空间。
技术大多是 trade-off,没有银弹。内存优化的本质是在计算、通信、精度之间找平衡点,找到那个"刚好够用"的状态。
折腾完这轮,对显存的恐惧少了一些,但对平衡的敬畏多了一分。工具是拿来用的,不是拿来迷信的。知道边界在哪里,才敢在安全区里大胆折腾。
版权声明: 本文首发于 指尖魔法屋-从显存走到效率:AI内存优化笔记(https://blog.thinkmoon.cn/post/236-ai-memory-optimization-gpu-efficiency-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。