跳转到内容

输入关键词开始搜索

    梯度累积与梯度检查点

    概念更新 2026-08-02置信度 high#概念#基础#长青#训练#显存

    两种常见显存优化手段:梯度累积用多个 micro-batch 合成一次参数更新;梯度检查点少保存激活、反向时重新计算。前者换更新频率,后者换计算时间。

    若每张设备的 micro-batch 为 bb,累积步数为 aa,设备数为 nn,则常用的有效 batch size 为:

    Beffective=b×a×n.B_{effective}=b\times a\times n.

    手写训练循环时,如果每个 micro-batch 的 loss 默认取 mean,常需除以 aa,并且只在累积边界调用 optimizer.step()

    optimizer.zero_grad()
    for i, batch in enumerate(loader):
    loss = model(**batch).loss / accumulation_steps
    loss.backward()
    if (i + 1) % accumulation_steps == 0:
    optimizer.step()
    optimizer.zero_grad()

    Trainer 一类框架通常会处理 loss 缩放、同步和最后一个不满周期的 batch;不能机械重复手写逻辑。

    累积允许用较小 micro-batch 完成较大的有效 batch,但不会消除模型权重、优化器状态或单个 micro-batch 激活。一次 optimizer update 需要更多前后向步骤,吞吐和调参语义也会改变。

    梯度检查点(Gradient Checkpointing)

    Section titled “梯度检查点(Gradient Checkpointing)”

    普通反向传播保存大量中间激活。梯度检查点只保存部分边界张量,反向经过被检查点包围的区段时重新执行前向:

    更少激活显存 ↔ 更多重计算时间

    它保存的是 activation checkpoint,不是训练断点文件,也不等于 save_checkpoint()

    • accumulation 控制一次参数更新由多少 micro-batch 组成;
    • checkpointing 控制单个 micro-batch 前向要保留多少激活;
    • AMP 改变部分算子的 dtype;
    • LoRA 减少可训练参数与优化器状态。

    四者可组合,但每一种优化影响的显存组成不同,不能简单把节省比例相乘。

    第 11 讲给出 Trainer 的 gradient_accumulation_steps 和显存观察;第 13 讲展示 checkpointing 的激活重算及保存输出;第 14 讲在 Llama 2 SFT 中组合这些配置。

    课程实验的对照有时同时改变 micro-batch、累积步数或其他设置,因此具体 MB 数不应当作严格消融结论。概念层只保留机制与可复核公式。

    • 累积步数翻倍,显存必减半:单个 micro-batch 未变时,峰值激活可能几乎不变。
    • 有效 batch 相同,训练必完全等价:随机层、归一化、调度器与分布式同步会造成差异。
    • checkpointing 不影响速度:它以额外前向重算换显存。