把大Batch拆开后如何保证梯度等效
前言
大 Batch 放不进显存时,最常见的做法是拆成几个 micro-batch,分别 forward/backward,最后只执行一次 optimizer step。
听起来就是“多调几次 backward”,但如果 micro-batch 大小不同、每个样本有效 token 数不同,或者 loss 本身按 token 归一化,直接累加得到的梯度并不一定等于完整 Batch。
Trio 还要把这件事拆过 SDK、Actor 和 Trainer 三层,因此除了数学等效,还要处理事件组的原子性、结果顺序和部分失败。
平均micro-batch loss是错的
假设完整逻辑 Batch 有 个样本,目标 loss 是:
现在拆成两个 micro-batch,大小分别为 和 。
如果每个 micro-batch 先算自己的平均 loss,再直接相加:
当 时,小 Batch 中的每个样本权重会更大。即使最后再除以 micro-batch 数,也不等于完整 Batch。
正确做法是让每片都使用逻辑 Batch 的全局分母:
然后:
不同loss的分母不一样
样本数只适用于 sample-mean 类型的 loss。
训练系统里至少还会遇到:
- Cross Entropy:可能按样本或有效 token 归一化
- PPO:通常按 active token 或 advantage mask 归一化
- CISPO、DRO、Importance Sampling:各自有不同的有效权重范围
因此 SDK 在拆 Batch 前要先统计逻辑 Batch 的全局量:
1 | total_samples |
每个 micro-batch 事件都携带相同的全局归一化元数据,Trainer 再根据当前 loss 类型选择正确分母。
如果只把 gradient_accumulation_steps=4 传下去,Trainer 并不知道四片各有多少样本和有效 token,只能假设它们完全等大,这在真实数据里通常不成立。
SDK规划逻辑Batch
SDK 先根据 Trainer 限制规划切片:
1 | 完整逻辑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 上逐字段完全相等。和另一套正式环境相比,相对误差约为 ,属于不同运行栈带来的浮点差异,不是分片误差。
最后
梯度累加不是简单把显存问题变成更多次 backward。想和完整 Batch 等效,必须先理解 loss 到底按什么归一化,并把逻辑 Batch 的全局统计穿过整条链路。
相关实现:
把大Batch拆开后如何保证梯度等效
https://blog.novashen.top/2026/07/14/tech/AI Infra/把大Batch拆开后如何保证梯度等效/