普通训练:
梯度累积的做法:
N 个小 batchN 次后,再一次性更新✅ 显存峰值 ≈ 小 batch 的显存
batch_size=4batch_size=32accumulation_steps=8loss.backward() 多次,再 optimizer.step()optimizer.zero_grad()梯度累积是小显存服务器的“穷人版大 batch”神器,但代价是训练更慢。
如果你愿意,我可以给你一段 PyTorch 梯度累积标准写法 或针对你具体模型/显存给建议。