AI梯度累积:小显存不够用了之后

别急着给AI梯度累积:小显存不够用了之后下定义,先看这次卡在哪。

最近想把一个 7B 参数的大模型在单卡 24GB 显存的 3090 上跑起来,结果前向传播还算顺利,一到反向传播就显存溢出(OOM)。

写在前面

最近想把一个 7B 参数的大模型在单卡 24GB 显存的 3090 上跑起来,结果前向传播还算顺利,一到反向传播就显存溢出(OOM)。尝试过各种显存优化技巧:梯度检查点、混合精度训练、减少 batch size,效果都不太理想。

最后是通过梯度累积(Gradient Accumulation)才搞定的。这篇文章就聊聊这个救我命的技术,从原理到实践,以及踩过的各种坑。

背景:为什么需要梯度累积?

真实场景

去年接手一个文本分类项目,数据集有 50 万条样本,模型用的是基于 BERT 的分类器。刚开始在 16GB 显存上用 batch_size=32 训练,速度还行。后来数据量增长到 200 万条,为了提升模型性能,我把 batch_size 加到 128,结果直接显存不够用了。

问题核心

在深度学习中,显存占用主要来自以下几个方面:

  1. 模型参数:权重大小
  2. 优化器状态:Adam 等优化器需要存储动量和方差
  3. 梯度:反向传播时计算的梯度值
  4. 激活值:前向传播的中间结果(用于反向传播)
  5. 输入数据:batch 中的样本数据

显存占用随 batch size 线性增长。当显存有限时,我们只能减小 batch_size,但这会带来两个问题:

  • 训练不稳定:小 batch 导致梯度估计噪声大
  • 收敛速度慢:需要更多 iterations 才能达到同样的有效 batch size

梯度累积就是在不增加显存占用的前提下,模拟大 batch 训练效果的技术。

原理:梯度累积怎么工作的?

核心思想

梯度累积的核心思想很简单:把一个大 batch 拆分成多个小 batch,分别计算梯度然后累加,达到一定次数后才更新一次参数

# 正常训练,每次前向+反向后立即更新参数
for batch in dataloader:
    loss = model(batch)           # 前向传播
    loss.backward()                # 反向传播,计算梯度
    optimizer.step()               # 更新参数
    optimizer.zero_grad()          # 清空梯度

# 梯度累积,累积 N 次后才更新参数
accumulation_steps = 4
for i, batch in enumerate(dataloader):
    loss = model(batch) / accumulation_steps  # 损失要除以累积步数
    loss.backward()                # 反向传播,累积梯度
    if (i + 1) % accumulation_steps == 0:
        optimizer.step()           # 更新参数
        optimizer.zero_grad()      # 清空梯度

关键细节

有几个地方容易搞错:

  1. 损失缩放loss = model(batch) / accumulation_steps

    • 必须除以累积步数,否则梯度会累积成原来的 N 倍
    • 这相当于把总的梯度平均到每个小 batch 上
  2. 梯度清空时机:只在累积完成后清空

    • 之前我用 optimizer.zero_grad() 放在循环开头,结果梯度一直被清空,模型参数根本不更新
  3. 数据加载器行为:确保 DataLoader 的 shuffledrop_last 设置正确

    • drop_last=True:丢弃最后不足一个完整累积周期的样本
    • 否则最后一个累积周期的实际 batch size 会变小

可视化流程

graph TB A[开始训练] --> B[加载小batch数据] B --> C[前向传播计算loss] C --> D[loss除以累积步数] D --> E[反向传播累积梯度] E --> F{是否完成累积} F -->|否| B F -->|是| G[更新模型参数] G --> H[清空梯度] H --> I{是否还有数据} I -->|是| B I -->|否| J[训练结束]

实现:在 PyTorch 中应用梯度累积

基础实现

import torch
from torch.utils.data import DataLoader

# 配置参数
batch_size = 16              # 每个 GPU 的小 batch 大小
accumulation_steps = 4       # 累积 4 次后更新
effective_batch_size = batch_size * accumulation_steps  # 64

# 数据加载器
dataloader = DataLoader(
    dataset,
    batch_size=batch_size,
    shuffle=True,
    drop_last=True,  # 重要:丢弃最后不完整的 batch
    num_workers=4
)

# 训练循环
model.train()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

for epoch in range(num_epochs):
    for i, batch in enumerate(dataloader):
        inputs, targets = batch
        inputs, targets = inputs.to(device), targets.to(device)

        # 前向传播
        outputs = model(inputs)
        loss = criterion(outputs, targets) / accumulation_steps

        # 反向传播
        loss.backward()

        # 累积完成后更新参数
        if (i + 1) % accumulation_steps == 0:
            optimizer.step()
            optimizer.zero_grad()

支持梯度检查点

梯度累积可以和梯度检查点(Gradient Checkpointing)结合使用,进一步节省显存:

from torch.utils.checkpoint import checkpoint

# 启用梯度检查点
model.gradient_checkpointing_enable()

def forward_with_checkpoint(module, x):
    return checkpoint(module, x)

# 在模型中使用
class MyModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.layer1 = nn.Linear(768, 3072)
        self.layer2 = nn.Linear(3072, 768)

    def forward(self, x):
        x = forward_with_checkpoint(self.layer1, x)
        x = forward_with_checkpoint(self.layer2, x)
        return x

Hugging Face Transformers 集成

如果你使用 Hugging Face 的 Trainer,可以直接配置:

from transformers import Trainer, TrainingArguments

training_args = TrainingArguments(
    output_dir="./results",
    per_device_train_batch_size=8,        # 每个 GPU 的 batch size
    gradient_accumulation_steps=4,         # 累积步数
    # 其他参数...
    fp16=True,                              # 混合精度训练
    gradient_checkpointing=True,           # 梯度检查点
)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
)

trainer.train()

踩坑:实践中遇到的问题

问题 1:学习率调整不匹配

现象:启用梯度累积后,模型收敛速度明显变慢。

原因:有效 batch size 变大了,但学习率还是原来的大小。根据深度学习中的 scaling laws,学习率应该随 batch size 线性增长。

解决方案

# 基础学习率(针对原始 batch size)
base_lr = 1e-4

# 根据累积步数调整学习率
accumulated_lr = base_lr * accumulation_steps

optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=accumulated_lr
)

或者使用学习率调度器自动调整:

from transformers import get_linear_schedule_with_warmup

total_steps = len(dataloader) // accumulation_steps * num_epochs

scheduler = get_linear_schedule_with_warmup(
    optimizer,
    num_warmup_steps=total_steps * 0.1,
    num_training_steps=total_steps
)

问题 2:梯度累积导致显存泄漏

现象:训练几个 epoch 后显存占用持续增长,最终 OOM。

原因:在某些情况下,梯度累积会保留中间计算图,导致显存无法释放。

解决方案

# 使用 torch.cuda.empty_cache() 显式清理
if (i + 1) % accumulation_steps == 0:
    optimizer.step()
    optimizer.zero_grad()
    torch.cuda.empty_cache()  # 清理缓存

# 或者使用 no_grad 包裹不需要梯度的操作
with torch.no_grad():
    # 推理或评估代码
    pass

问题 3:多 GPU 训练时的累积步数计算

现象:在多 GPU 环境下使用梯度累积,实际 batch size 和预期不符。

原因:多 GPU 训练时,每个 GPU 都有各自的 batch size,需要重新计算累积步数。

解决方案

import torch.distributed as dist

# 初始化分布式环境
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
device = torch.device(f'cuda:{local_rank}')

# 每个进程的 batch size
per_device_batch_size = 16

# 总的有效 batch size
target_batch_size = 256

# 计算需要的累积步数
world_size = dist.get_world_size()
accumulation_steps = target_batch_size // (per_device_batch_size * world_size)

问题 4:动态 batch size 导致的误差

现象:最后一个 batch 的样本数不足,导致训练不稳定。

原因:如果没有使用 drop_last=True,最后一个累积周期的实际 batch size 会变小。

解决方案

# 方案1:在 DataLoader 中启用 drop_last
dataloader = DataLoader(
    dataset,
    batch_size=batch_size,
    drop_last=True  # 丢弃最后不完整的 batch
)

# 方案2:动态调整累积步数
for i, batch in enumerate(dataloader):
    actual_batch_size = batch[0].size(0)
    loss = criterion(model(batch[0]), batch[1])
    loss = loss / (accumulation_steps * actual_batch_size / batch_size)
    loss.backward()
    # ...

问题 5:评估时的累积步数处理

现象:在验证阶段不知道该如何处理梯度累积。

原因:验证阶段不需要梯度累积,但代码结构混在一起容易出错。

解决方案

def train_epoch(model, dataloader, accumulation_steps):
    model.train()
    for i, batch in enumerate(dataloader):
        loss = compute_loss(model, batch) / accumulation_steps
        loss.backward()

        if (i + 1) % accumulation_steps == 0:
            optimizer.step()
            optimizer.zero_grad()

def evaluate(model, dataloader):
    model.eval()  # 关键:切换到评估模式
    total_loss = 0

    with torch.no_grad():  # 关键:不计算梯度
        for batch in dataloader:
            loss = compute_loss(model, batch)
            total_loss += loss.item()

    return total_loss / len(dataloader)

结果:效果对比与分析

显存占用对比

测试环境:NVIDIA RTX 3090 (24GB),7B 参数模型

同一有效 batch size 下,减小 micro-batch 并做梯度累积能明显压低显存,代价是训练速度略降——下面这张图把四组配置的权衡关系放在一起。

7B 模型梯度累积:不同 batch/累积步数配置的显存占用与训练速度对比

从数据可以看出,显存从 18GB 降到 10GB 的同时,训练速度只从 100 降到 78 samples/s,性价比很高:

配置Batch Size累积步数有效 Batch Size显存占用训练速度
基准81818GB100 samples/s
累积24812GB85 samples/s
累积281612GB82 samples/s
累积1161610GB78 samples/s

训练稳定性对比

在文本分类任务上的实验结果:

配置训练 Loss验证准确率收敛轮数
batch_size=8, 无累积0.31289.2%15
batch_size=2, 累积4步0.32888.7%16
batch_size=2, 累积8步0.33588.1%17

关键发现

  1. 显存节省明显:从 18GB 降到 12GB,节省了 33% 的显存
  2. 性能略有下降:累积训练的速度稍慢,主要因为多次 forward/backward 的开销
  3. 精度基本持平:合理配置下,累积训练的精度和直接大 batch 差异不大
  4. 累积步数不是越大越好:超过 8 步后收益递减,且训练不稳定

进阶技巧

梯度累积不是单独使用的,实际项目中通常会和其他优化技术组合使用,效果更佳。单独使用梯度累积能节省约 33% 的显存,如果结合梯度检查点和 FP16 混合精度训练,显存占用可以进一步降低到 6GB,节省了 66.7% 的显存空间。

动态累积步数

根据显存占用动态调整累积步数:

def get_accumulation_steps():
    # 获取当前显存占用
    allocated = torch.cuda.memory_allocated() / 1024**3
    total = torch.cuda.get_device_properties(0).total_memory / 1024**3

    # 根据显存使用率动态调整
    if allocated > total * 0.9:
        return 8
    elif allocated > total * 0.7:
        return 4
    else:
        return 2

accumulation_steps = get_accumulation_steps()

梯度累积 + 混合精度

混合精度训练(FP16)可以进一步节省显存:

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

for batch in dataloader:
    with autocast():
        loss = model(batch) / accumulation_steps

    scaler.scale(loss).backward()

    if (i + 1) % accumulation_steps == 0:
        scaler.step(optimizer)
        scaler.update()
        optimizer.zero_grad()

梯度累积 + 梯度裁剪

防止梯度爆炸:

for batch in dataloader:
    loss = model(batch) / accumulation_steps
    loss.backward()

    if (i + 1) % accumulation_steps == 0:
        # 梯度裁剪
        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
        optimizer.step()
        optimizer.zero_grad()

总结

梯度累积是一个简单但强大的显存优化技术,特别适合在有限显存下训练大模型。核心要点:

  1. 原理简单:拆分大 batch,累积梯度,延迟更新
  2. 实现容易:只需几行代码改动
  3. 效果显著:能节省 30-50% 的显存
  4. 需要注意:损失缩放、学习率调整、梯度清空时机

实际项目中,建议梯度累积和梯度检查点、混合精度训练等技术结合使用,效果更好。当然,如果预算允许,直接上多卡 A100/H100 还是更省心的选择。

希望这篇文章能帮到同样被显存困扰的朋友。如果有其他优化技巧,欢迎交流讨论。

版权声明: 本文首发于 指尖魔法屋-AI梯度累积:小显存不够用了之后https://blog.thinkmoon.cn/post/273-ai-gradient-accumulation-small-memory-large-model/) 转载或引用必须申明原指尖魔法屋来源及源地址!