Transformers 梯度检查点实战:用约 20% 训练开销换取激活显存的大幅下降
梯度检查点(Gradient Checkpointing)是 Transformers 中控制训练显存的核心手段之一:前向传播不再缓存所有中间激活,而是只保留检查点处的激活,反向传播时按需重算被丢弃的部分。阅读本文后,你将掌握如何在 TrainingArguments 与 Trainer 中启用该特性、它如何在源码层面改造模型的每层 __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`."
},
)
两个要点:
gradient_checkpointing默认False,不会偷偷生效;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 的核心步骤:
- 校验模型支持:若
self.supports_gradient_checkpointing为假,直接抛出ValueError,避免在不支持的架构上静默失败; - 绑定 checkpoint 函数:默认
gradient_checkpointing_kwargs为{"use_reentrant": False},并用functools.partial(checkpoint, **gradient_checkpointing_kwargs)封装成_gradient_checkpointing_func。默认use_reentrant=False与现代 PyTorch 的 checkpoint 实现保持一致,兼容性更好; - 遍历模块打标记:调用
_set_gradient_checkpointing(enable=True, ...),为模型中所有带有gradient_checkpointing布尔属性的模块写入 checkpoint 函数并置位开关。若有模块不带该属性,会抛出兼容性错误; - 兼容旧格式:对 Hub 上 transformers < 4.35.0 时代的旧式
_set_gradient_checkpointing(value=...)签名,会自动回退并打印弃用警告; - 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 属性只要有一个子模块开启即返回 True;gradient_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__而不是forward:gradient_checkpointing_enable 的 docstring 明确说明,传__call__是为了让模块的 forward hook 依然生效; - KV cache 自动关闭:反向重放会二次写入 cache,因此训练态下
use_cache、past_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()会自动处理嵌入层梯度传播。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0623
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00