微调显存估算

微调与对齐深水4讲解 2

显存三档账:权重 + 梯度 + 优化器状态 + 激活。全参微调 7B 光优化器状态就是权重的两倍以上, LoRA 把梯度与优化器状态压到可训参数那一小份,QLoRA 再把权重压到 4 bit;激活随 batch 与序列长度涨,梯度检查点拿算力换显存。

也叫:显存估算 · 全参微调 · 优化器状态 · 激活显存 · 梯度检查点

显存三档账出自 T4-2

先记一句总纲:三种方案的差别,几乎全在「梯度和优化器状态跟谁走」。

全参微调时,每个参数都要带:权重、梯度、优化器状态。用 Adam 且混合精度训练:

经典口径(含 fp32 主权重副本,ZeRO 论文的算法):
  BF16 权重 2 + BF16 梯度 2 + FP32 主权重 4 + FP32 一阶动量 4 + FP32 二阶动量 4 = 16 bytes/param

现代纯 BF16 口径(不保 fp32 主权重):
  BF16 权重 2 + BF16 梯度 2 + FP32 双动量 8                              = 12 bytes/param

面试就答"12 到 16 之间,取决于优化器实现有没有留 fp32 主权重副本"——能说清这句话本身就是分水岭。

LoRA / QLoRA 时,基座冻结,所以基座只占"权重"那一项,没有梯度也没有优化器状态;那 12~16 bytes 的账只按可训参数算。

7B 模型的完整对照(r=64,all-linear):

基座 可训参数三件套 激活 合计
全参(12B/param) 14 GB 70 GB 2–20 GB 约 88 GB
全参(16B/param) 14 GB 98 GB 2–20 GB 约 116 GB
LoRA r=64 14 GB(BF16 冻结) 约 1.9 GB(160M × 12) 2–20 GB 约 19 GB
QLoRA r=64 3.5 GB(NF4 冻结) 约 1.9 GB 2–20 GB 约 9 GB

那个最容易被漏掉的加项:激活显存。 它跟 batch × seq_len × hidden × 层数 走,是 2–20 GB 的浮动项,也是"我算出来 20 G,结果 80 G 卡还是 OOM"的头号原因。压它的标准手段是梯度检查点(gradient checkpointing)——前向时只存少数几个中间结果,反向时重算,拿时间换显存,典型代价是训练慢 20%–30%。

以上节选自T4-2 参数高效微调:LoRA、QLoRA 与一笔算得清的显存账,读全文能看到前后语境。

延伸阅读

考这个知识点的题1

会连带问到3