返回关卡地图
预判 0 / 2 闯关题:0 / 2 XP 0 / 200

Level 19 | 零基础导学关卡

激活检查点

Activation Checkpointing

先通过图解和小例子理解原理,再做预测与练习;不懂的地方可以反复尝试,最后把理解写进官方 Notebook。

本课官方 Notebook ↗
19_Activation_Checkpointing.ipynb 显存优化激活值Checkpointing
Mission 1 概念练习

先读调用契约:把每个 Block 的前向交给 checkpoint

🎯 先猜一猜

循环已经来到第 3 个 Transformer Block,当前变量 x 是前两个 Block 更新后的状态。

哪种写法既执行当前 Block,又把结果继续传给下一层?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • blocks 是 nn.ModuleList;循环中的 block 是一个可调用的 nn.Module,调用 block(x) 会返回下一层输入。
  • x 是需要梯度的浮点张量。Checkpoint 不改变它的逻辑 shape;本题的测试输入是 [2, 2048, 2048],device 是 CUDA。
  • checkpoint 的第一个位置参数是要执行的函数,后面的位置参数是传给该函数的输入;use_reentrant=False 是本 Notebook 要求的关键字参数。
  • 循环必须把返回值重新赋给 x,否则后面的 Block 仍会收到旧状态,网络就没有按层向前传播。

图解原理

普通循环写的是 x = block(x)。Checkpoint 版本没有换模型,也没有跳过计算,只是在 block 外面加了一个“反向时可以重算”的包装器。因此最稳的解题方式是先写出普通调用,再把 block(x) 改写成 checkpoint(block, x, ...)。

输入检查点保存这一段开始时的 x 前向执行仍然调用当前 block 少存激活中间结果不全部长期保留 反向重算需要梯度时再次执行该段前向

普通版本

xblock(x)新的 x

Autograd 为反向传播保留 Block 内部需要的中间激活。

Checkpoint 版本

xcheckpoint(block, x, ...)同形状的新 x

数学前向不变,变化的是 Autograd 保存与重算中间量的策略。

函数run_with_checkpointing(blocks: nn.ModuleList, x: torch.Tensor)循环输入当前 block 和上一层返回的 x循环输出新的 x,继续传给下一个 block关键参数use_reentrant=False最终返回最后一个 Block 的输出张量
不要把 checkpoint 写成 checkpoint(block(x), ...)。第一个参数必须是“稍后可以重新调用的函数”,而不是已经算完的张量。

语法热身:用一个不同的小函数练“函数和参数分开传入”

import torch
from torch.utils.checkpoint import checkpoint

def add_bias(values, bias):
    return values + bias

values = torch.randn(4, requires_grad=True)
bias = torch.ones(4, requires_grad=True)
result = checkpoint(add_bias, values, bias, use_reentrant=False)
result.sum().backward()

从语法例子迁移到 Notebook

  • add_bias:对应循环中的 block;二者都是可以稍后再次调用的 callable
  • values:对应当前 x;它是传给 callable 的输入张量
  • result:对应 checkpoint 返回的新 x,必须接住并继续向后传
  • use_reentrant=False:原样映射到 TODO,不能误当成 block 的参数

巩固一下

为什么 TODO 中不能只写 checkpoint(block, x, use_reentrant=False) 而不赋值?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

Mission 2 概念练习

再读测试:结果要能 backward,收益看 CUDA 峰值

🎯 先猜一猜

你只运行了 run_with_checkpointing(blocks, x),没有调用 backward。

此时为什么还不能完整观察 checkpoint 的机制?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • Checkpoint 是训练期优化,只有执行 backward 才会触发重算;只看一次 forward 不能验证完整机制。
  • 测试先运行普通版本并记录 mem_normal,再清理输出、梯度和 CUDA 缓存,随后记录 mem_ckpt。
  • 输入 x_input 设置 requires_grad=True,最终对 out.sum() 调用 backward;这保证梯度链路必须保持完整。
  • 本节的显存测试明确要求 NVIDIA GPU;没有 CUDA 时 Notebook 会跳过,而不是说明实现已经通过。

图解原理

这道题不是测试“输出数值变小”,而是测试“同一条可求导计算,在少保存激活的情况下仍能完成反向”。评价顺序应是:先保证梯度链路正确,再比较峰值显存,最后讨论时间换空间是否值得。

1run_without_checkpointing完成 forward + backward,记录 mem_normal
2清理状态删除旧输出、清空 x_input.grad、重置峰值统计
3run_with_checkpointing再次完成 forward + backward,记录 mem_ckpt
4比较峰值可见断言接受 mem_ckpt 不大于 mem_normal

空间从哪里省

前向不长期保存每个 Block 内的所有中间激活,只保留重算所需的边界输入。

时间花到哪里

反向传播需要恢复激活时,会重新执行对应 Block 的前向计算。

shape
前后都保持 [B,S,D]
device
测试张量和 Blocks 都在 CUDA
梯度
x_input.requires_grad=True
指标
峰值 allocated memory

语法热身:用小网络练峰值统计的调用顺序

import torch
import torch.nn as nn

tiny_stage = nn.Linear(32, 32).cuda()
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()

sample = torch.randn(8, 32, device='cuda', requires_grad=True)
result = tiny_stage(sample)
result.sum().backward()

peak_mb = torch.cuda.max_memory_allocated() / (1024 ** 2)
print(peak_mb)

从语法例子迁移到 Notebook 测试

  • tiny_stage(sample):对应 run_without_checkpointing 或 run_with_checkpointing 的一次完整前向
  • result.sum().backward():对应测试中触发 Autograd 和 checkpoint 重算的动作
  • reset_peak_memory_stats:必须在每种策略开始前调用,避免沿用上一轮峰值
  • peak_mb:对应 mem_normal 或 mem_ckpt,单位从 byte 换算为 MB

巩固一下

本地没有 CUDA,测试打印“忽略测试”时,最准确的结论是什么?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

把理解变成自己的代码

准备好,去官方题目试一试

在本页用小例子建立直觉,再去官方 Notebook 完成实现。先运行你自己的测试,遇到困难时再查看官方提示与参考答案。

前往本课官方 Notebook ↗

这节课会遇到的代码对象

SimpleTransformerBlock · forward · run_without_checkpointing · run_with_checkpointing · build_checkpoint_segments · run_with_segment_checkpointing

先找题目里的输入、输出与 TODO,再把本课的手算过程对应进去;以官方题目中的函数说明和测试为准。

练习来源:Datawhale 官方仓库 · 518cc45。这里的入口直接打开官方版本,不读取或分享你的本地 Notebook。

完成 2 道闯关题后,本关即算完成;作业 checklist 用来辅助你回 notebook 练习。
进度只保存在当前浏览器 localStorage,分享 HTML 不会带走你的记录。
上一关:L18 下一关:L20