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. 正文

Pytorch autograd notes 2:hooks

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

这篇文章来简单整理一下autograd相关的hook,这些hook在优化模型训练方便有比较多的贡献,比如:

  • fsdp所依赖的各种hook,做pre-backward, post-backward。
    • forward相关的hook是在python的module层,不在autograd
  • save_tensor hook,做activation offload,activation checkpoint都依赖这些hook的实现

Node level

Node的定义上有这四种hook。区别是执行的时机和顺序不同。代码中有一块Hook ordering的注释解释了他们的执行顺序

```C++
std::vector<std::unique_ptr<FunctionPreHook>> pre_hooks_;
std::vector<std::unique_ptr<FunctionPreHook>> tensor_pre_hooks_;
std::unordered_map<size_t, std::unique_ptr<FunctionPreHook>>
retains_grad_hooks_;
std::vector<std::unique_ptr<FunctionPostHook>> post_hooks_;
```

字段 触发条件 对应用户 API
pre_hooks_ 只有该 Node 真正被执行时 Node.register_prehook
tensor_pre_hooks_ 只要引擎遍历到该 Node(哪怕不执行) Tensor.register_hook
retains_grad_hooks_ 同 tensor_pre_hook,但永远排在最后 Tensor.retain_grad()
post_hooks_ Node 执行之后 Node.register_hook

有关这个pre_hook_和tensor_pre_hooks_的区别可以看Engine::evaluate_function的代码。在没有剪枝的情况下,hook都是在call_function中执行的:

  • 先调用call_tensor_pre_hooks(input)
    • 对应先调用tensor_pre_hooks_,然后是retains_grad_hooks_
  • call_pre_hooks(input)

  • 调用node的fn,即grad fn,计算output

  • call_post_hooks(input, output)

然后API这块让AI帮忙整理了一下,简单看看就好:

Tensor 级 hook

Tensor.register_hook(fn)

  • 语义:每次算出该 tensor 的梯度时调用,hook(grad) -> Tensor | None,返回值可替换梯度。

  • 实现:Python 入口 torch/_tensor.py → C++ VariableHooks::_register_hook,包装成 CppFunctionTensorPreHook(cpp_hook.h/.cpp)挂到 grad_fn 的 tensor_pre_hooks_(非叶子)或 AccumulateGrad 上。

Tensor.register_post_accumulate_grad_hook(fn)

  • 语义:仅对叶子张量,在梯度累加进 .grad 之后调用,hook(param) -> None,常用于 optimizer-in-backward。

  • 实现:PostAccumulateGradHook,存在 AccumulateGrad 节点。

Tensor.retain_grad()

  • 语义:让非叶子张量在 backward 后保留 .grad。本质是注册一个特殊的 retains_grad hook。

  • 实现:VariableHooks::retain_grad(variable.cpp:531),用 CppFunctionSingleTensorPreHook 挂到 retains_grad_hooks_。

Node 级 hook(torch.autograd.graph.Node)

Node.register_prehook(fn):hook(grad_outputs) -> tuple | None

  • Node 执行前,挂 pre_hooks_。

Node.register_hook(fn):hook(grad_inputs, grad_outputs) -> tuple | None

  • Node 执行后,挂 post_hooks_。

多张量 hook

torch.autograd.graph.register_multi_grad_hook(tensors, fn, mode="all"|"any")

  • 语义:等所有(all)或任一(any)张量的梯度就绪时触发一次。

  • 实现:纯 Python,构建在 Tensor.register_hook 之上

Saved Tensor

torch.autograd.graph.saved_tensors_hooks对应这个API
这个API的作用是在tls的栈上放一对hook,即pack_hook, unpack_hook。因为是栈维护的有嵌套的能力,但是只有最里面的一层生效,而非整个栈都按顺序调用一遍

forward阶段,每个算子需要构造SavedVariable,用来保存到grad_fn中。在构造SavedVariable的时候就会去检查tls看是否有设置pack_hook,然后set_hooks_and_pack_data,调用hook,改变保存的tensor。
注意这个SavedVariable是per tensor级别的。所以是每一个save for backward的tensor都会调用这个hook
hook的调用链路就是:

  • set_hooks_and_pack_data

  • call_pack_hook

  • 然后调用到python层的函数

unpack对应SavedVariable::unpack

  • call_unpack_hook,通过unpack_hook把tensor取出来

  • 注意这个需要和上面的pack是一对的,所以forward的时候把hook保存下来了,backward就没有读tls了。

至于unpack的调用点,不在autograd engine中,而是在每个算子(grad_fn)中,因为保存就是在算子内部自行保存的,所以unpack也在这个层面。

  • 另一个调用点是custom function,对应custom_function.cpp中的get_saved_variables,在这里先都unpack出来。调用点则是ctx.saved_tensors

  • 还有一个要注意的就是如果在python层主动获取这个saved tensor,比如tensor.grad_fn._saved_xxx的时候,也会先unpack再返回

除去上面说的全局的通过tls来pack每个tensor,还有另一个针对单个tensor做的hook:

  • SavedVariable::register_hooks对应 Python 侧 grad_fn._raw_saved_.register_hooks(pack, unpack),只对某一个已保存张量生效,且一个张量只能设一对。

  • 调用时立刻生效

Graph level

之前介绍过GraphTask是一次autograd一个,里面会有final_callback
通过Engine::queue_callback来注册。

  • python层面是:torch.autograd.Variable._execution_engine.queue_callback(fn)

在整个Autograd执行完之后会调用。代码在exec_post_processing中

Module Level

这里主要看一下API即可,因为和Autograd的关系不大

Forward hook

接口 签名 触发时机
register_forward_pre_hook(hook, *, prepend=False, with_kwargs=False) hook(module, args) 或 hook(module, args, kwargs) forward() 之前,可修改输入
register_forward_hook(hook, *, prepend=False, with_kwargs=False, always_call=False) hook(module, args, output) 或 hook(module, args, kwargs, output) forward() 之后,可修改输出

Backward hook

接口 签名 触发时机
register_full_backward_pre_hook(hook, prepend=False) hook(module, grad_output) -> tuple/None 计算该 module 梯度前,可替换 grad_output
register_full_backward_hook(hook, prepend=False) hook(module, grad_input, grad_output) -> tuple/None 计算完梯度后,可替换 grad_input

module 本身没有 grad_fn,这些 hook 靠 torch.utils.hooks.BackwardHook 实现:它给 forward 的输入/输出 tensor 套一层 view。然后用tensor级别的hook挂到这个view的输出上面。
这样当输入/输出的梯度算好时,也会流到对应的view的grad fn上
正常在算完「对 module 输入的梯度」时触发;若输入都不需要梯度,则在算完「对输出的梯度」时触发;若输出都不需要梯度则不触发。所以这里的view可能发生在输入/输出上面

Module 还有一些state_dict的hook,和autograd彻底没关系了,就不提了

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

sheep

think again

点赞
< 上一篇
下一篇 >

文章评论

取消回复

COPYRIGHT © 2021 heavensheep.xyz. ALL RIGHTS RESERVED.

THEME KRATOS MADE BY VTROIS