这篇文章来简单整理一下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彻底没关系了,就不提了
文章评论