把大Batch拆开后如何保证梯度等效

前言

大 Batch 放不进显存时,最常见的做法是拆成几个 micro-batch,分别 forward/backward,最后只执行一次 optimizer step。

听起来就是“多调几次 backward”,但如果 micro-batch 大小不同、每个样本有效 token 数不同,或者 loss 本身按 token 归一化,直接累加得到的梯度并不一定等于完整 Batch。

Trio 还要把这件事拆过 SDK、Actor 和 Trainer 三层,因此除了数学等效,还要处理事件组的原子性、结果顺序和部分失败。

平均micro-batch loss是错的

假设完整逻辑 Batch 有 NN 个样本,目标 loss 是:

L=1Ni=1NliL = \frac{1}{N}\sum_{i=1}^{N} l_i

现在拆成两个 micro-batch,大小分别为 n1n_1n2n_2

如果每个 micro-batch 先算自己的平均 loss,再直接相加:

L=1n1iB1li+1n2iB2liL' = \frac{1}{n_1}\sum_{i\in B_1}l_i + \frac{1}{n_2}\sum_{i\in B_2}l_i

n1n2n_1\neq n_2 时,小 Batch 中的每个样本权重会更大。即使最后再除以 micro-batch 数,也不等于完整 Batch。

正确做法是让每片都使用逻辑 Batch 的全局分母:

Lk=1NiBkliL_k = \frac{1}{N}\sum_{i\in B_k}l_i

然后:

L=kLk\nabla L = \sum_k \nabla L_k

不同loss的分母不一样

样本数只适用于 sample-mean 类型的 loss。

训练系统里至少还会遇到:

  • Cross Entropy:可能按样本或有效 token 归一化
  • PPO:通常按 active token 或 advantage mask 归一化
  • CISPO、DRO、Importance Sampling:各自有不同的有效权重范围

因此 SDK 在拆 Batch 前要先统计逻辑 Batch 的全局量:

1
2
3
total_samples
total_active_tokens
total_advantage_tokens

每个 micro-batch 事件都携带相同的全局归一化元数据,Trainer 再根据当前 loss 类型选择正确分母。

如果只把 gradient_accumulation_steps=4 传下去,Trainer 并不知道四片各有多少样本和有效 token,只能假设它们完全等大,这在真实数据里通常不成立。

SDK规划逻辑Batch

SDK 先根据 Trainer 限制规划切片:

1
2
3
4
5
完整逻辑Batch
→ 检查样本数、序列长度和总token
→ 生成若干不一定等大的micro-batch
→ 计算全局归一化元数据
→ 按事件组顺序提交

拆分必须保留原始样本顺序,返回结果时也要按原下标重新聚合。否则 loss 输出和用户提交的样本无法对应。

这里还要区分“单次 Event 成功”和“事件组完整成功”。如果四片中前三片已经 backward,第四片提交失败,当前 Trainer 中只留下了 3/4 梯度。

这份梯度不能继续用于 optimizer step。

因此事件组只要出现部分提交、结果状态不确定或中途失败,就终止当前训练 Session,避免用户误用不完整梯度。看起来比较严格,但比悄悄更新一次错误参数安全得多。

等效测试

同一份数据分别执行:

  • 完整Batch一次forward/backward
  • 不均匀micro-batch多次forward/backward
  • 一次optimizer step

然后比较:

  • loss_sum
  • loss_mean
  • token_count
  • 每层梯度
  • optimizer state
  • 更新后的 LoRA 参数

一次真实测试中,32 样本的结果为:

方式 loss_sum loss_mean token_count
单片32 799.5701904296875 6.246642112731934 128
16+16 799.5701904296875 6.246642112731934 128

同一 Trainer 上逐字段完全相等。和另一套正式环境相比,相对误差约为 7.63×1077.63\times10^{-7},属于不同运行栈带来的浮点差异,不是分片误差。

最后

梯度累加不是简单把显存问题变成更多次 backward。想和完整 Batch 等效,必须先理解 loss 到底按什么归一化,并把逻辑 Batch 的全局统计穿过整条链路。

相关实现:

作者

Noah Shen

发布于

2026-07-14

更新于

2026-07-14

许可协议

评论