技术指南

梯度检查点

梯度检查点(也称为激活检查点)是一种节省内存的技巧,它在前向传播期间丢弃大多数中间激活,并在反向传播期间动态重新计算它们。

阅读时间:2分钟最后更新

概述

It lets you train deeper, larger networks by trading extra compute for much lower memory use.

深入探讨

训练神经网络通常会在前向传播过程中存储每一层的激活,因为反向传播需要它们来计算梯度。对于深度模型来说,这些激活支配着记忆。相反,梯度检查点仅在一组稀疏的“检查点”层上保存激活,并丢弃其余部分。当反向传播到达激活被删除的区域时,它会重新运行该段的前向计算以重新生成所需的内容,然后继续。大约每个 N 平方根层都放置检查点,激活内存从 N 阶下降到 N 平方根阶,而计算量仅增加大约一次额外的前向传递(大约慢 20-30%)。这使得在同一 GPU 上适应更大的批量大小或更深的转换器成为可能。

技术洞察

该技术利用了时间与内存的权衡。存储所有激活速度很快,但很耗内存;相对于耗尽内存的成本,在现代加速器上重新计算它们的成本很低。像 PyTorch (torch.utils.checkpoint) 这样的框架包装了一个模块,因此它的前向输出被保存,但它的内部结构在后向过程中被重新计算。选择检查点放置很重要:大约 sqrt(N) 段的均匀间距可以最大限度地减少总内存,同时仅添加一个额外的前向计算整体。

战略影响

成本与预算

多年来,架构决策决定着性能和运营成本。

更清晰的判决

技术教育帮助团队选择正确的堆栈,而不仅仅是最新的堆栈。

质量控制

更好的工程选择可以减少生产中的可靠性事故。

梯度检查点的未来

梯度检查点现在是大型模型训练的标准,并且越来越自动化,库会为您选择最佳检查点位置。它与 FSDP、混合精度和卸载自然配合,以提高模型尺寸。期望“选择性”检查点仅重新计算廉价的操作,同时缓存昂贵的操作(如注意力矩阵),加上 PyTorch 的 torch.compile 等工具中的编译器驱动方法,自动决定保存哪些内容与重新计算以获得最佳速度内存平衡。

现实世界的实施

通过丢弃和重新计算层激活,在单个 GPU 上训练具有更大批量大小的深度转换器。

在高分辨率图像上微调视觉模型,否则激活图会溢出 GPU 内存。

拥抱 Face Transformers 启用gradient_checkpointing=True,以在微调期间适应十亿参数模型。

将检查点与 FSDP 相结合,使参数和激活都保持较小,从而能够训练非常大的语言模型。

风险与防护栏

优化一项基准测试可以隐藏更广泛的系统弱点。

基础设施和维护成本常常被低估。

随着系统变得更加复杂,安全性和可观察性差距可能会扩大。

实施路线图

1

在实施之前定义延迟、质量和成本目标。

2

在实际负载和数据条件下进行基准测试。

3

仪器监控错误、漂移和用户影响。

4

在扩展之前准备回滚和事件响应路径。

不断探索

Free newsletter

Get the daily AI briefing

Three verified AI stories every weekday morning, written in plain English. Free forever, no ads.

One email each weekday. Unsubscribe in one click. We never sell or share your address.

Test yourself

Take the Gradient Checkpointing quiz

Instant feedback on every answer, and a shareable certificate with a verifiable ID once you pass a course.

开始测验

Support free AI education. AI Understanding is a 501(c)(3) nonprofit — no ads, no paywall, ever. Make a donation

常见问题

What is Gradient Checkpointing?

梯度检查点(也称为激活检查点)是一种节省内存的技巧,它在前向传播期间丢弃大多数中间激活,并在反向传播期间动态重新计算它们。它可以让您通过用额外的计算换取更低的内存使用来训练更深、更大的网络。

为了节省内存,梯度检查点主要交换什么?

梯度检查点在向后传递过程中重新计算丢弃的激活,花费额外的计算来换取显着减少的内存。

为什么激活通常在前向传递过程中存储?

反向传播使用前向传递的中间激活来计算梯度,因此它们必须可用,除非重新计算。

如果在 N 层网络中的每个 sqrt(N) 层放置检查点,那么激活内存大致如何扩展?

每个 N 平方根层的间隔检查点将存储的激活内存从 N 阶减少到 sqrt(N) 阶。

适当放置的梯度检查点通常会增加大约多少额外计算?

如果检查点位置良好,开销大约是一次额外的前向传递,通常会减速 20-30% 左右。

在 PyTorch 中,哪个实用程序通常用于将梯度检查点应用于模块?

torch.utils.checkpoint 包装一个模块,以便在向后期间重新计算其内部激活而不是存储。