PyTorch 梯度累积怎么写?处理小批次与最后不足一组的数据

把十条样本的 4、4、2 小批次按真实样本数累积成 8、2 两次更新,核对最后不足一组的数据未遗漏,并与真实大批次对照参数。

梯度累积的基本做法是:连续几个小批次执行 backward(),在一组结束时才 optimizer.step() 和清梯度。最容易漏掉的是最后不足一组的数据,以及不同批次样本数导致的损失权重错误。本文用 10 条标量回归样本,将 4、4、2 条的三个小批次分成 8 条和 2 条两个有效批次,逐组核对与真正的大批次更新是否一致。

PyTorch 梯度累积说明介绍了有效批次和累积边界;以下先使用普通精度单进程 CPU 示例,把样本数与更新次数检查清楚,未测量显存或速度。

PyTorch 梯度累积怎么写?处理小批次与最后不足一组的数据

为什么不能每次都简单除以同一个累积次数

当每个小批次同样大,而且一组恰好齐全时,把每次平均损失除以累积次数可得到整组平均梯度。但最后一组可能只有一个小批次,甚至最后一个小批次还更小;继续除以固定次数,会把这组的梯度缩小。对于不等大小的批次,要明确整组损失的实际分母。

本例每条样本只有一个回归输出:先对每个小批次计算求和损失,再除以这一整组的真实样本数。于是 8 条那组除以 8,最后 2 条那组除以 2。每个小批次完成反向后释放自己的计算图,不把多次前向的 loss 张量攒起来再一次反向。

完整代码:保留最后两条并做大批次对照

以下完整脚本在 Windows、Python 3.11.15、PyTorch 2.8.0+cpu、NumPy 2.3.3 环境实际运行通过,只验证这个 CPU 玩具任务。它不下载模型权重,也不代表真实模型的精度、速度或显存表现。在已启用的独立 Python 环境中可安装相同依赖:

python -m pip install torch==2.8.0 --index-url https://download.pytorch.org/whl/cpu
python -m pip install numpy==2.3.3

保存为 accumulate_demo.py,运行 python accumulate_demo.py。为减小浮点归约差异,核对使用 float64 和小型线性层;读取器一次只保留最多两个小批次。reference 分支只用来验证,接入真实训练时可移除。

import copy
from itertools import islice
import torch
from torch.utils.data import DataLoader, TensorDataset

torch.set_num_threads(1)
torch.manual_seed(7)
x = torch.linspace(-1, 1, 10).reshape(-1, 1)
y = 2 * x + 0.5
base = torch.nn.Linear(1, 1).double()
accumulated = copy.deepcopy(base)
reference = copy.deepcopy(base)
x, y = x.double(), y.double()
loader = DataLoader(TensorDataset(x, y), batch_size=4, shuffle=False)
accumulation_steps = 2
opt = torch.optim.SGD(accumulated.parameters(), lr=0.1)
ref_opt = torch.optim.SGD(reference.parameters(), lr=0.1)

updates = 0
batches = iter(loader)
while True:
    group = list(islice(batches, accumulation_steps))
    if not group:
        break
    group_samples = sum(len(xb) for xb, _ in group)
    opt.zero_grad(set_to_none=True)
    for xb, yb in group:
        # For scalar regression, sum loss / group sample count = mean loss.
        loss = torch.nn.functional.mse_loss(accumulated(xb), yb, reduction="sum")
        (loss / group_samples).backward()
    opt.step()
    updates += 1

    # A true larger batch with exactly the same samples is the reference.
    all_x = torch.cat([xb for xb, _ in group])
    all_y = torch.cat([yb for _, yb in group])
    ref_opt.zero_grad(set_to_none=True)
    torch.nn.functional.mse_loss(reference(all_x), all_y).backward()
    ref_opt.step()
    print("update_samples:", group_samples)

max_diff = max((a - b).abs().max().item()
               for a, b in zip(accumulated.parameters(), reference.parameters()))
print("optimizer_updates:", updates)
print("max_parameter_diff:", max_diff)
assert updates == 2 and max_diff < 1e-12

核对样本数、更新数和参数差

本次运行结果:

update_samples: 8
update_samples: 2
optimizer_updates: 2
max_parameter_diff: 0.0

第一次更新使用 8 条,第二次使用最后 2 条,优化器确实更新 2 次;与逐组真正拼成大批次的对照相比,参数差为 0.0。读者应同时确认最后不足一组的数据触发了一次更新、总更新数正确,并在同一初始权重下比较参数差,而不是只看 loss 能打印。

实际设备或算子上的归约顺序可能产生很小的浮点差异,因此代码用小于 1e-12 的阈值检查这个 float64 玩具任务,而不是要求所有实际模型逐位相等。

接到大模型训练前,先检查这些边界

  • 本例是标量 MSE。多维回归的 mean 还会除以输出元素数;语言模型的掩码 token 损失应按有效 token 数,带权分类损失可能按权重和。先核对自己的损失定义,再确定分母,不能一律除以样本数。
  • 学习率调度器通常跟随实际优化器更新,而不是每次小批次反向;核对原调度器设计的步数单位。
  • 梯度裁剪应放在本组梯度累积完、更新之前;不要每个小批次先独立裁剪。
  • 加入 AMP 后,一组内保持梯度缩放因子一致,组末才 unscale、step 和 update;具体按前述官方文档接入。

梯度累积经常用来让有效批次大于单次前向的批次,但模型权重、优化器状态和其他常驻内存仍存在,所以不能保证一个原本加载不下的模型因此能运行。本文没有验证 GPU、分布式同步或任何显存节省比例。

常见问题:为什么和大批次训练仍不同?

先检查归一化分母、遗漏的最后一组和 zero_grad 的位置,再看 BatchNorm 等依赖单次批次统计的层、Dropout 随机性与数据增强。梯度累积只合并梯度,不会自动让每次前向的批次统计变成整组统计;本例刻意使用不含这些模块的线性层来检查基础逻辑。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30321.html

赞 (0)
AI小管家的头像AI小管家
CPU 跑 AI 推理有多快?用固定模型测延迟和吞吐量
上一篇 1天前
用 DeepSeek 学机器学习数学:从线性打分推导 sigmoid 概率
下一篇 1天前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部