首页
/ Transformers 梯度检查点实战:用约 20% 训练开销换取激活显存的大幅下降

Transformers 梯度检查点实战:用约 20% 训练开销换取激活显存的大幅下降

2026-09-06 15:15:51作者:江焘钦

梯度检查点(Gradient Checkpointing)是 Transformers 中控制训练显存的核心手段之一:前向传播不再缓存所有中间激活,而是只保留检查点处的激活,反向传播时按需重算被丢弃的部分。阅读本文后,你将掌握如何在 TrainingArgumentsTrainer 中启用该特性、它如何在源码层面改造模型的每层 __call__,以及如何用 every_n_layers 在“省显存”与“省时间”之间做精细权衡。

原理:前向只存检查点,反向按需重算

标准前向传播会把每一层的中间激活都缓存下来供反向传播使用,而激活大小随 batch size 和序列长度线性增长,往往是训练显存的最大头之一。梯度检查点改变了这一策略:

Normal training:
  Forward:   [L1]→[L2]→[L3]→[L4]   (save ALL activations)
  Backward:  ←uses cached activations everywhere

Gradient checkpointing:
  Forward:   [L1]→[L2]→[L3]→[L4]   (save only at checkpoints, discard the rest)
  Backward:  ←reaches L2, recomputes L2→L3 from scratch, uses it, discards it

代价是:部分激活需要在反向传播到达检查点时重新前向计算一次,因此训练速度约慢 20%,换来的是激活显存的显著下降。这正是“以计算换显存”(trade compute for memory)的典型做法——当你的瓶颈是 GPU 显存而不是时间时,这是一个几乎零风险的开关。

官方文档 docs/source/en/grad_checkpointing.md 还建议把它和梯度累积配合使用,进一步压低显存占用;并指引读者通过GPU 显存构成混合精度训练与Kernels文档继续优化。

启用方式:一个 TrainingArguments 开关

最直接的方式是在 TrainingArguments 中打开 gradient_checkpointing

from transformers import TrainingArguments

args = TrainingArguments(
    ...,
    gradient_checkpointing=True,
)

TrainingArguments 的定义可以看到这个参数的全部信息:

# src/transformers/training_args.py
gradient_checkpointing: bool = field(
    default=False,
    metadata={
        "help": "Enable gradient checkpointing to trade compute for memory. "
                "Reduces memory at the cost of ~20%% slower training."
    },
)
gradient_checkpointing_kwargs: dict[str, Any] | str | None = field(
    default=None,
    metadata={
        "help": "Keyword arguments passed to `gradient_checkpointing_enable()`. `every_n_layers` "
                "checkpoints only every n-th decoder layer instead of all of them; `1` is the usual "
                "all-or-nothing behavior, and larger values give some memory back to speed. "
                "Any other key is forwarded to `torch.utils.checkpoint.checkpoint`."
    },
)

两个要点:

  1. gradient_checkpointing 默认 False,不会偷偷生效;
  2. gradient_checkpointing_kwargs 是高级配置入口:其中 every_n_layers 被单独取出用于控制“每隔几层才设一个检查点”,其余键值会原样转发给 PyTorch 的 torch.utils.checkpoint.checkpoint 函数。

进阶:every_n_layers 控制检查点密度

every_n_layers 的语义与 PreTrainedModel.gradient_checkpointing_enable 的文档一致:1 表示逐层全检查点(即常规的全有或全无行为),更大的值则让部分层保留激活、只检查点每隔第 n 层——当全量检查点省下的显存“用不完”时,用它换回一些训练速度是值得的。用法示例:

args = TrainingArguments(
    ...,
    gradient_checkpointing=True,
    gradient_checkpointing_kwargs={"every_n_layers": 2},  # 只检查点每隔一层
)

源码链路:Trainer 如何把开关落到每一层模型

Trainer 开始训练时,它会读取上述参数并显式调用模型的开启方法。相关逻辑在 Trainer.train

# src/transformers/trainer.py
# Activate gradient checkpointing if needed
if args.gradient_checkpointing:
    # `every_n_layers` selects which layers are checkpointed; the remaining keys are
    # forwarded to `torch.utils.checkpoint.checkpoint`, so it has to come out of the
    # dict before that happens.
    gc_kwargs = dict(args.gradient_checkpointing_kwargs or {})
    every_n_layers = gc_kwargs.pop("every_n_layers", 1)
    self.model.gradient_checkpointing_enable(
        gradient_checkpointing_kwargs=gc_kwargs or None, every_n_layers=every_n_layers
    )

也就是说,参数解析顺序是:先把 every_n_layers 从 kwargs 字典中剥离出来,剩下的键(如 use_reentrant)才作为 checkpoint 函数的参数传递。

PreTrainedModel.gradient_checkpointing_enable 做了什么

gradient_checkpointing_enable 的核心步骤:

  1. 校验模型支持:若 self.supports_gradient_checkpointing 为假,直接抛出 ValueError,避免在不支持的架构上静默失败;
  2. 绑定 checkpoint 函数:默认 gradient_checkpointing_kwargs{"use_reentrant": False},并用 functools.partial(checkpoint, **gradient_checkpointing_kwargs) 封装成 _gradient_checkpointing_func。默认 use_reentrant=False 与现代 PyTorch 的 checkpoint 实现保持一致,兼容性更好;
  3. 遍历模块打标记:调用 _set_gradient_checkpointing(enable=True, ...),为模型中所有带有 gradient_checkpointing 布尔属性的模块写入 checkpoint 函数并置位开关。若有模块不带该属性,会抛出兼容性错误;
  4. 兼容旧格式:对 Hub 上 transformers < 4.35.0 时代的旧式 _set_gradient_checkpointing(value=...) 签名,会自动回退并打印弃用警告;
  5. PEFT 场景的特殊处理:若主输入是 input_ids 或已加载 PEFT 配置,会额外调用 enable_input_require_grads()

第 5 点是很多 LoRA 微调踩坑的地方:enable_input_require_grads 会在输入嵌入层上注册 forward hook,把嵌入输出标记为 requires_grad_(True)。原因在于微调时只有 adapter 层可训练,冻结层的输出如果不携带梯度信息,梯度就无法穿过 checkpoint 区段传播到 adapter 权重。

_set_gradient_checkpointing:every_n_layers 的落点

_set_gradient_checkpointing 对模型所有子模块遍历:

  • 凡是 hasattr(module, "gradient_checkpointing") 的模块,都会被赋予 _gradient_checkpointing_func
  • 只有属于 GradientCheckpointingLayer 的重复块(decoder 层)才参与 every_n_layers 计数:layer_index % every_n_layers == 0 的层被置 True,其余层保持 False。注释特别说明:即使其他模块也带有 gradient_checkpointing 标志,计数也只统计逐层重复块,因此“每 n 层”语义准确;
  • 若整个模型没有任何带该属性的模块,抛出错误提示架构不兼容。

此外还有配套的状态查询与关闭接口:is_gradient_checkpointing 属性只要有一个子模块开启即返回 Truegradient_checkpointing_disable() 则负责反向复位,并在 PEFT 场景下调用 disable_input_require_grads()

每层的运行时行为:GradientCheckpointingLayer

真正在训练时“拦截”前向的是 GradientCheckpointingLayer。它的 __call__training 且本层开关打开时执行三件事:

# src/transformers/modeling_layers.py(节选)
def __call__(self, *args, **kwargs):
    if self.gradient_checkpointing and self.training:
        # 1. 训练态下 KV cache 与 checkpoint 不兼容:强制关闭并警告
        if "use_cache" in kwargs and kwargs["use_cache"]:
            kwargs["use_cache"] = False
        if not self._can_checkpoint_with_cache:
            if "past_key_value" in kwargs and kwargs["past_key_value"] is not None:
                kwargs["past_key_value"] = None
            # ... past_key_values / layer_past 同理
        # 2. 用 checkpoint 包装本层 forward
        return self._gradient_checkpointing_func(partial(super().__call__, **kwargs), *args)
    return super().__call__(*args, **kwargs)

几个值得注意的实现细节:

  • 为什么包 __call__ 而不是 forwardgradient_checkpointing_enable 的 docstring 明确说明,传 __call__ 是为了让模块的 forward hook 依然生效;
  • KV cache 自动关闭:反向重放会二次写入 cache,因此训练态下 use_cachepast_key_value(s)layer_past 会被强制置空并给出一次性警告;只有“只读 cache”的层可以通过 _can_checkpoint_with_cache = True 豁免;
  • use_reentrant=True 的参数约束:docstring 特别警告,若使用 reentrant 模式,需要梯度的输入(如 hidden states)必须以位置参数而非关键字参数传入,否则梯度无法正确传播。这也是默认 use_reentrant=False 的原因之一;
  • 推理不受影响:所有拦截逻辑都在 self.training 分支内,eval 模式下模型走普通 __call__ 路径。

与分布式/FSDP 的交互

在分布式训练时,参数间存在互斥与冗余,源码中有两处明确的防御性检查:

  • Trainer:当 FSDP 插件同时配置了 activation_checkpointing 且用户又打开了 gradient_checkpointing 时会警告,因为两者语义重叠;
  • TrainingArguments 校验:明确提示 FSDP 场景下应优先使用 FSDP 的 activation_checkpointing,因为 Transformers 的 gradient_checkpointing 在反向阶段会引入一次冗余的 AllGather。

从源码结构看,建议在使用 FSDP 时二选一,避免“双重检查点 + 额外通信”的组合。

测试层面的验证

仓库为各模型提供了统一的检查点行为回归测试。例如 GPT-2 的建模测试 中:

self._test_lm_generate_gpt2_helper(gradient_checkpointing=True)

大量模型(GPT-Neo、GPT-J、CodeGen、MPT、ESM、BioGPT 等,见 tests/models/*/test_modeling_*.py)都调用了 create_and_check_forward_and_backwards(..., gradient_checkpointing=True),验证开启检查点后前向/反向输出仍然正确。若要自行验证某个模型是否支持,可以直接检查 model.supports_gradient_checkpointing 或调用 model.gradient_checkpointing_enable() 观察是否抛出异常。

小结与适用建议

  • 显存紧张(尤其是长序列、大 batch 场景)时,打开 TrainingArguments(gradient_checkpointing=True) 即可,无需改模型代码;
  • 若全量检查点省下的显存有富余,用 gradient_checkpointing_kwargs={"every_n_layers": n} 提升速度;需要调整底层行为时,字典中的其他键会转发给 torch.utils.checkpoint.checkpoint
  • 搭配梯度累积可进一步降低显存;配合 bf16 混合精度与 fused kernels(参见 mixed precision training、Kernels)能同时兼顾速度与显存;
  • FSDP 环境下优先使用框架自带的 activation checkpointing,避免重复机制叠加;
  • PEFT/LoRA 微调无需额外操作,enable_input_require_grads() 会自动处理嵌入层梯度传播。
登录后查看全文
热门项目推荐
相关项目推荐