More than code

More Than Code
The efficiency of your iteration of reading, practicing and thinking decides your understanding of the world.
  1. 首页
  2. 未分类
  3. 正文

torch selective activation checkpoint

2026年7月12日 53点热度 0人点赞 0条评论

背景是之前看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可能不太友好

标签: 暂无
最后更新:2026年7月12日

sheep

think again

点赞
< 上一篇
下一篇 >

文章评论

取消回复

COPYRIGHT © 2021 heavensheep.xyz. ALL RIGHTS RESERVED.

THEME KRATOS MADE BY VTROIS