梯度检查点(Gradient Checkpointing) 是一种在训练大模型时常用的 显存优化技术,其目标是在不牺牲模型性能的前提下,节省 GPU 显存,以便能训练更大的模型或使用更大的 batch size。


🧠 原理概述

在标准的训练过程中(使用自动微分),前向传播(forward)时每层的中间激活值都会保留,以便在反向传播(backward)时计算梯度。

问题是:

这些中间值占用了大量显存,尤其是在深层模型中。


✅ 梯度检查点怎么做?

梯度检查点的核心思想是:

在前向传播时只保留部分关键节点(称为“检查点”)的中间激活值,其余的在反向传播时再重新计算。

也就是说:

  • 少存一点中间结果(省显存)。

  • 回头再算一次前向(多花一点算力)。

是一种 以计算换显存 的策略。


🧮 举个例子

假设你有一个模型分成了 10 层:

  • 正常训练时:forward 会保存 10 层的中间激活,占用大量显存。

  • 用梯度检查点:比如只保存第 0 层、第 5 层、第 10 层的结果。

    • backward 时,先从第 5 层重算 6~10 层,再计算这些层的梯度。

    • 这样最多只需要保存一小部分层的中间结果,节省大量显存。


📉 优缺点

项目优点缺点
显存使用显著减少—
计算量增加(重复前向)更慢一些
适用场景大模型、内存受限时训练对推理没用

🔧 在 PyTorch 中的用法示例

import torch
from torch.utils.checkpoint import checkpoint

def custom_forward(*inputs):
    return model_layer(*inputs)

# 使用 checkpoint 包裹某一层
output = checkpoint(custom_forward, input_tensor)

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐