微调显存优化栈

OOM 是微调最常见阻塞。本篇按显存组成 → 单卡手段 → 多卡手段 → 调参顺序组织,与 02-02 QLoRA05-04 全参 vs PEFT 配合使用。

段末注释:Gradient Checkpointing(梯度检查点)以前向重算换显存,不存全部中间激活;Gradient Accumulation(梯度累积)用小 micro-batch 模拟大 batch。

系列索引:微调技术路线导读


一、显存去哪了

组件 全参 LoRA QLoRA
模型权重 基座 + 小 adapter 4bit 基座 + adapter
优化器状态 全量 mainly adapter 8bit paged 常见
梯度 全量 adapter adapter
激活 batch×长度×层数 同左(可 checkpoint) 同左

反向传播时激活常是 LoRA 场景的大头 → 优先 gradient_checkpointing


二、单卡优化(按性价比排序)

手段 配置 效果
LoRA peft_config 砍掉大部分梯度/优化器
QLoRA BitsAndBytesConfig 权重 4bit
Gradient Checkpointing gradient_checkpointing=True 激活 ↓,速度略慢
Gradient Accumulation gradient_accumulation_steps=8 等效大 batch,micro-batch ↓
降 batch per_device_train_batch_size=1~2 直接降激活
降 max_length max_length=256 线性影响激活
BF16 bf16=True 较 FP32 省一半权重激活
收窄 target_modules 仅 q/v proj 可训练参数 ↓
Packing SFTConfig(packing=True) 提吞吐,间接省 step
1
2
3
4
5
6
7
8
SFTConfig(
per_device_train_batch_size=2,
gradient_accumulation_steps=8,
gradient_checkpointing=True,
bf16=True,
max_length=512,
optim="adamw_torch", # ROCm 见 AMD 实战
)

三、调参顺序(OOM 时)

1
2
3
4
5
6
7
1. gradient_checkpointing=True
2. per_device_train_batch_size ÷ 2
3. max_length 降低
4. gradient_accumulation_steps × 2(保 effective batch)
5. 收窄 LoRA target_modules
6. QLoRA(CUDA 环境)
7. 减样本或 packing

03 SFTTrainer §6.2 一致。


四、多卡与分布式

策略 适用 说明
DDP 每卡完整模型 数据并行;LoRA 常用
ZeRO-2/3 全参 / 大模型 DeepSpeed;分片优化器/梯度
FSDP PyTorch 原生 accelerate / Trainer 集成

微调 7B LoRA 多数 单卡 + checkpointing 即可;全参 7B+ 考虑 ZeRO-3 / FSDP。公式背景 Math-04/20 §3


五、加速(非省显存但省时间)

工具 说明
Unsloth Triton 内核;与 TRL 集成
Liger Kernel use_liger_kernel=True
Flash Attention 2 长上下文 attention 省显存

详见 03 §十一5023 Unsloth


六、平台注意

平台 建议
NVIDIA BF16 + LoRA/QLoRA
AMD ROCm BF16 LoRA;adamw_torch;避免 bnb 4bit
消费级 24GB 7B LoRA max_length 512 常见可训

七、小结

显存优化 = PEFT 优先 + checkpointing + 小 batch 累积 + 控长度;仍 OOM 上 QLoRA(CUDA)或多卡。验收用 04 任务指标,勿为省显存把 max_length 截断到丢 label。

下一步:08-02 遗忘与版本管理 | 07-02 分布式 | 07-03 加速

-------------本文结束感谢您的阅读-------------