激活重新计算权衡
激活重新计算(梯度或激活检查点)通过丢弃前向传递中的中间激活并在后向传递中重新计算它们,在训练期间节省 GPU 内存。
概述
It trades extra compute for the ability to train larger models or longer sequences on the same hardware.
深入探讨
反向传播需要前向传递激活来计算梯度,因此默认情况下存储每一层的输出——巨大的内存成本随着模型大小、批量大小和序列长度的增加而增长。激活重新计算仅保留一些“检查点”张量(通常只是层边界)并丢弃其余部分。在向后传递期间,它重新运行检查点之间的前向计算,以根据需要重新生成丢弃的激活。经典的结果是,在每 sqrt(N) 层放置检查点时,内存会下降到大约 O(sqrt(N)),同时增加大约一次额外的前向传递(计算量增加约 33%)。选择性变体仅重新计算廉价但占用大量内存的操作(例如注意力或丢失),同时缓存昂贵的操作,以少得多的重新计算开销获得大部分内存节省。
技术洞察
基本的权衡是内存与 FLOPs。完全重新计算大约每步增加一次额外的前向传递(慢约 30-40%),但可以将激活内存减少一个数量级。明智之举是选择性检查点:识别内存大但计算成本低的操作(softmax、layernorm、GELU、注意力分数)并仅重新计算这些操作,同时缓存昂贵的 GEMM 的结果 - 最大限度地减少计算浪费。
战略影响
成本与预算
多年来,架构决策决定着性能和运营成本。
更清晰的判决
技术教育帮助团队选择正确的堆栈,而不仅仅是最新的堆栈。
质量控制
更好的工程选择可以减少生产中的可靠性事故。
激活重新计算权衡的未来
重新计算变得越来越自动化和选择性。现在,框架会分析每个操作的内存和 FLOP 成本,以选择最佳检查点,并将重新计算与激活卸载到 CPU/NVMe 以及并行策略相结合。随着上下文长度和模型大小不断增长,预计编译器驱动的策略(在 PyTorch、JAX/XLA 中)会自动选择每个操作的重新计算决策,再加上重新计算与通信的更紧密重叠,因此额外的 FLOP 被部分隐藏。
现实世界的实施
通过检查每个层块来训练一个大型变压器,否则该变压器无法适应
使用 PyTorch 的 torch.utils.checkpoint 包装变压器块并切割激活内存
Megatron-LM 中注意力/softmax 的选择性重新计算,以节省内存且速度减慢最小
通过重新计算激活而不是存储它们,在固定的 GPU 预算上实现更长的序列长度
风险与防护栏
优化一项基准测试可以隐藏更广泛的系统弱点。
基础设施和维护成本常常被低估。
随着系统变得更加复杂,安全性和可观察性差距可能会扩大。
实施路线图
在实施之前定义延迟、质量和成本目标。
在实际负载和数据条件下进行基准测试。
仪器监控错误、漂移和用户影响。
在扩展之前准备回滚和事件响应路径。
不断探索
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 Activation Recomputation Tradeoffs 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 Activation Recomputation Tradeoffs?
激活重新计算(梯度或激活检查点)通过丢弃前向传递中的中间激活并在后向传递中重新计算它们,在训练期间节省 GPU 内存。它以额外的计算能力换取在相同硬件上训练更大模型或更长序列的能力。
激活重新计算会牺牲什么来节省内存?
重新计算会丢弃存储的激活并在向后传递中重新生成它们,从而花费额外的计算来减少内存使用。
为什么前向传递激活通常会被存储?
向后传递使用前向激活来计算梯度,因此默认情况下它们会保留在内存中,直到向后传递运行。
完全激活重新计算通常会增加大约多少额外计算?
完全重新计算会在后向传递过程中重新运行前向计算,大约增加一次额外的前向传递,计算量增加了 30-40% 左右。
选择性(非完整)重新计算背后的想法是什么?
选择性重新计算的目标是使用大量内存但计算量很少的操作(例如 Softmax 或 Layernorm),同时缓存昂贵的 GEMM 结果以最大程度地减少浪费的 FLOP。
哪种补充技术经常与重新计算相结合以节省更多内存?
激活卸载将一些激活移动到 CPU/NVMe 存储,并经常与重新计算和并行性相结合,以进一步节省内存。