导读:本期聚焦于张立峰创作的《大模型全参数微调显存爆满怎么办?几种实用的显存优化方案详解》,敬请观看详情。全参数微调时显存不够用是训练大模型最常见的拦路虎。单卡24G显卡想微调一个7B模型,往往连前向传播都撑不过去。其实显存占用是有明确构成的:模型参数、梯度、优化器状态、中间激活值各有各的账,算清楚了才能对症下药。本文从显存占用的计算公式入手,分析为什么7B模型全参数微调需要上百G显存,再依次讲解混合精度训练、梯度检查点、ZeRO分片优化、8bit优化器等主流优化手段的原理和代码实现,并给出不同显卡条件下的方案组合建议,帮你用有限的硬件跑通全参数微调。

训练一个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导致训练极慢。把显存构成算清楚,按需组合优化手段,全参数微调并没有想象中那么遥不可及。

全参数微调显存优化混合精度训练修改时间:2026-09-16 09:36:42

免责声明:​ 已尽一切努力确保本网站所含信息的准确性。网站内容多为原创整理与精心编撰,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们处理。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。