背景是之前看torchtitan中用了很多这个SAC,然后看实现比较复杂。最近把dispatch/autograd差不多读了读有了基本的概念,这里来看一下SAC的实现,希望可以更好的整合对dispatch/autograd的理解
实现主要在torch/utils/checkpoint.py中。
这里有一段claude帮忙整理的基本思路,可以对SAC的实现思路有一个了解:
普通(vanilla)非重入 checkpoint 的做法是:forward 时用 saved_tensors_hooks 把所有本该被 autograd 保存的 activation 全部"劫持"掉不存,backward 需要时整段重算。这是"要么全存、要么全不存"。
SAC 想要的是逐 op 粒度:某些贵的 op(比如 matmul)的输出直接存下来,backward 时不重算;其余便宜的 op 才重算。
关键 insight:vanilla checkpoint 是在 autograd 层(saved_tensors_hooks)运作的,而"哪个 op 要存"的决策发生在 dispatch 层(torch_dispatch)。SAC 就是把这两层缝在一起:
- 上层(autograd)仍然用 vanilla 那套:forward 不真正存,backward 触发整段 recompute。
-
下层(dispatch)用一个共享的
storage字典:forward 时记录"每个 op 的哪次调用要存、存的值是什么";recompute 时拦截这些 op,对"要存"的直接返回缓存值(短路,不真算),对"要重算"的才真正执行。
这样 recompute 走的还是完整的 fn,但被 save 的 op 全被缓存短路了,只有 recompute 的 op 真正跑,重算成本就降下来了。
这里SAC是建立在vanilla checkpoint之上,用torch dispatch来决策是否做SAC的。所以需要先把vanilla checkpoint的实现看懂。
这块我建议是先把use_reentrant=true的看懂,然后看一下_checkpoint_without_reentrant_generator链路,但是先不用关注forward_context相关的。
_checkpoint_without_reentrant_generator我个人理解应该是一个compile friendly的实现,这块我没有细抠,只是看里面有很多dynamoc/compile相关的字段
简单说一下这里的实现:
- 一个checkpoint fn对应一个checkpoint frame,里面存这个算子对应的meta
-
forward时候,进入_checkpoint_hook,会把save for backward的内容替换成一个holder。释放掉原本的tensor
-
backward的时候,第一次调用unpack hook的时候,开始进行recompute。recompute就是用之前保存的ctx,比如rng state, autocast ctx等,重新调用一次forward fn
- 这里recompute的时候会进入_recomputation_hook,作用就是pack tensor的时候,把他保存下来。同时还有一个优化就是做early return,因为recompute不一定需要所有的forward intermediate,只要把需要的计算完,就可以结束recompute了。
- unpack_hook结束后,再正常调用backward fn即可。
然后再看SAC的实现,就是引入了一个forward_context, recompute_context,在torch dispatch层发生作用。
用户可以给定一个policy fn,policy fn的作用是说明一个op是做recompute/save
_CachingTorchDispatchMode,在执行算子之前,会进入到这个hook中:
- 这里会记录每个算子的执行次数,根据op, counter来确定每一个算子的顺序(避免缓存的结果错位)
-
根据policy fn来决定是否进行重算,或者是做save。
-
如果是save的话,这里会把算子的输出保存到一个dict中。
-
这里也处理了compile相关的。对于compile来说,(应该是)在aot autograd的时候执行,会把SAC的决策放到node meta中,compile后续在AOT autograd的partition的时候就会处理
_CachedTorchDispatchMode,对应的是recompute ctx,作用就是把之前save的tensor回填:
- 判断如果是save的op的话,直接从之前forward中的结果取出来,就不用再重算了。
不过这里也可以看到,因为作用粒度是在op的层面,可能有一些matmul在backward中不需要,也不需要save。比如MoE中的fc2,应该在recompute的时候是不需要重算的,但是如果这里只写matmul的话就去分不出来了。所以理论上是可以做的更细一点的。
autograd.Function的处理
还有一个点,就是可能legacy的代码的自定义算子都是通过autograd.Function做的,而不是custom op。
先看一下两种定义的区别:
import torch
class MyMul(torch.autograd.Function):
@staticmethod
def forward(ctx, x, y):
# 保存需要在 backward 中用到的张量
ctx.save_for_backward(x, y)
# 非张量的东西直接挂在 ctx 上
ctx.some_flag = True
return x * y
@staticmethod
def backward(ctx, grad_out):
x, y = ctx.saved_tensors
grad_x = grad_out * y
grad_y = grad_out * x
return grad_x, grad_y # 对应 forward 每个输入的梯度
out = MyMul.apply(a, b)
import torch
from torch import Tensor
@torch.library.custom_op("mylib::my_mul", mutates_args=())
def my_mul(x: Tensor, y: Tensor) -> Tensor:
return x * y
@my_mul.register_fake
def _(x, y):
return torch.empty_like(x)
def setup_context(ctx, inputs, output):
x, y = inputs
# 在这里保存反向需要的张量
ctx.save_for_backward(x, y)
def backward(ctx, grad):
x, y = ctx.saved_tensors
return grad * y, grad * x
my_mul.register_autograd(backward, setup_context=setup_context)
其中方法二custom op的方法对compile更加友好,也是发生在op粒度,可以适配上面的SAC功能。
不过两种方法其实逻辑比较简单,可以通过一个简单的wrapper来处理:
- 写一个custom op的wrapper,forward就是执行forward,save tensor到ctx中后,在wrapper中,可以把ctx的saved tensor取出来,作为output输出。此时forward就是纯函数了
-
setup_context中把output中的saved tensor取出来返回
-
backward的时候,把saved tensor赛回去。然后再做autograd.Function的backward即可
不过这种方法好像会额外做一次unpack/pack。对offload可能不太友好
文章评论