训练一个7B参数的大模型,用FP32做全参数微调,理论上光是模型本体、梯度和优化器状态就需要超过280GB显存,这还不算前向传播过程中产生的中间激活值。而一张消费级显卡只有24GB显存,差距显而易见。很多团队在动手微调之前都会低估这个数字,结果跑到一半OOM直接退出。想要在有限硬件上完成全参数微调,必须先搞清楚显存到底花在了哪里,再选择合适的优化手段把显存压下来。本文围绕显存占用的构成和几类主流优化方案展开,配合代码示例说明具体实现方式。

全参数微调的显存到底花在哪里
搞优化之前先算账。全参数微调时,显存主要由四部分组成:模型参数、梯度、优化器状态和中间激活值。假设模型有N个参数,用FP32训练,Adam优化器会为每个参数维护一阶动量和二阶动量两组状态,加上参数本身和梯度,静态部分就是4N+4N+4N+4N=16N字节。一个7B模型就是大约112GB,这只是理论下限,实际还有框架开销。
如果改用混合精度训练,参数用FP16存储,梯度也是FP16,但优化器状态仍然需要FP32的参数副本和动量,静态部分变成2N+2N+4N+4N+4N=16N字节,看起来没省,但实际上激活值的占用减半了,这才是混合精度的主要收益来源。激活值的显存和批次大小、序列长度、模型层数成正比,长文本训练时激活值往往是最大的显存杀手。
明白了这个构成,优化思路就清晰了:要么减少静态部分的冗余(比如把优化器状态量化存储),要么用计算换显存(比如重算激活值),要么把状态切分到多张卡上(数据并行分片)。下面的方案都是围绕这三条路展开的。
混合精度训练与梯度检查点
混合精度是最容易落地的优化。PyTorch原生提供了torch.cuda.amp模块,前向和反向传播用FP16或BF16计算,权重更新时用FP32累积,既省显存又提速。需要注意的是,如果显卡支持BF16(Ampere架构以后),优先用BF16,它的数值范围更大,不需要额外的loss scaling,训练更稳定。
import torch
from torch.cuda.amp import autocast, GradScaler
model = get_model().cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5)
scaler = GradScaler() # FP16需要,BF16可以省略
for step, batch in enumerate(dataloader):
optimizer.zero_grad()
with autocast(dtype=torch.bfloat16):
loss = model(**batch).loss
loss.backward()
optimizer.step()
梯度检查点(Gradient Checkpointing)解决的是激活值占用问题。正常反向传播需要保留所有中间激活值,开启检查点后只保留少数几个边界点的激活,反向传播到某一层时再重新计算这一层的激活值。代价是增加大约30%的计算时间,换来激活值显存降低到原来的几分之一。在Hugging Face的Transformers库里,一行配置就能开启:
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen2-7B",
torch_dtype=torch.bfloat16,
)
model.gradient_checkpointing_enable()
# 使用重入式实现时需要开启这个,否则梯度无法回传
model.config.use_cache = False
这两个手段几乎没有副作用,建议无条件开启。单靠它们,7B模型在80GB的A100上已经可以做全参数微调,但在24GB的消费级卡上还差得远,需要更激进的手段。
ZeRO分片与8bit优化器
数据并行的传统做法是每张卡都复制完整的模型、梯度和优化器状态,ZeRO(Zero Redundancy Optimizer)的核心思想是把这些状态切分到不同卡上:每张卡只存自己负责的那一份分片,需要时再临时聚合。ZeRO分三个阶段,Stage 1切分优化器状态,Stage 2额外切分梯度,Stage 3连模型参数也切分。对全参数微调来说,Stage 2性价比最高,Stage 3适合单卡放不下模型的极端情况。
用DeepSpeed启用ZeRO Stage 2配合梯度检查点,7B模型的微调显存可以从80GB以上压到30GB左右每卡。配置文件示例如下:
{
"bf16": {"enabled": true},
"zero_optimization": {
"stage": 2,
"offload_optimizer": {
"device": "cpu",
"pin_memory": true
},
"allgather_partitions": true
},
"gradient_accumulation_steps": 4,
"train_micro_batch_size_per_gpu": 1
}
注意配置里的offload_optimizer,它把优化器状态卸载到CPU内存,进一步把每卡显存压到20GB以内。代价是优化器步骤变慢,如果CPU内存够大(比如128GB以上),这个交换通常划算。另外还可以把Adam换成8bit量化版本,bitsandbytes提供的AdamW8bit把动量状态从FP32压成int8,优化器状态显存直接降为四分之一,训练效果基本没有损失:
import bitsandbytes as bnb
optimizer = bnb.optim.AdamW8bit(
model.parameters(),
lr=1e-5,
betas=(0.9, 0.999),
)
需要注意,8bit优化器和ZeRO的优化器切分不能同时叠加使用,二者都是针对优化器状态做文章,选一个即可。单卡场景选8bit优化器,多卡场景优先ZeRO。
不同硬件条件下的方案组合建议
把上面的手段组合起来,不同显卡各有最优解。单卡24GB(如RTX 4090)微调7B模型:混合精度BF16加梯度检查点加8bit优化器加优化器CPU卸载,再用1到2的批次大小配合梯度累积,勉强可以跑通,但训练速度较慢,属于能跑但不舒服的状态。单卡48GB(如A6000)或双卡24GB:去掉CPU卸载,体验会好很多。
如果是13B以上模型,单卡基本不现实,要么上多卡ZeRO Stage 3,要么认真考虑是否真的需要全参数微调。很多场景下LoRA等参数高效微调方法能拿到接近的效果,显存只需全参数微调的几分之一。判断标准是任务是否要求模型深度改变行为模式,比如领域知识注入通常全参数微调更好,风格和能力微调用LoRA往往足够。
最后提醒几个实操细节:训练前用torch.cuda.memory_reserved()监控真实占用,PyTorch的缓存分配器会预留比实际更多的显存,排查OOM时要区分开;梯度累积的步数要和学习率缩放策略匹配;开启CPU卸载时确认内存充足,否则会触发swap导致训练极慢。把显存构成算清楚,按需组合优化手段,全参数微调并没有想象中那么遥不可及。