已开启
[Usage]: API一致性:torch.jit.ScriptModule类下6个hook相关api一致性说明 #1723
黄小猛创建于 4月24日
4月24日 添加了label:usage
黄小猛
4月24日 评论:
4月24日 评论:
4月25日 添加了label:event: api-consistency
dinglaiping
4月25日 评论:
4月25日 评论:
新增用例文件,2.7.1直到master,共6个版本都要提交PR。新增代码中不要用中文。


黄小猛
4月27日 评论:
4月27日 评论:
register_full_backward_hook相关issue:https://github.com/pytorch/pytorch/issues/181473


5月15日 关联了pull request:[feat] Add test cases for verifying torch.jit.ScriptModule hook API on NPU
5月15日 修改了issue 的描述
5月16日 关联了pull request:[feat] Add test cases for verifying torch.jit.ScriptModule hook API on NPU
5月16日 关联了pull request:docs: fix description for api torch.jit.ScriptModule hook API on NPU
5月16日 关联了pull request:docs: fix description for api torch.jit.ScriptModule hook API on NPU
7月4日 issue类型由 Bug-Report 改变为 任务
7月8日 关联了看板:MindStudio ISSUE管理
7月28日 添加了label:bot-triaged
TorchNPU-Bot
7月28日 评论:
7月28日 评论:
检测到当前 issue 已关联 PR !34408,自动添加标签:bot-triaged


在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。
环境信息
对应版本:torch-npu 2.6.0、2.7.1、2.8.0、2.9.0、2.10.0、2.11.0、2.12.0
环境信息
操作系统:AlmaLinux 8.10
CANN 软件版本:8.5.0
安装的软件版本:torch、torch-npu 2.6.0~2.12.0
使用场景及问题
通过
torch.jit.script()创建的模型是RecursiveScriptModule实例,与测试目的不符,所以本次测试模型均通过继承torch.jit.ScriptModule创建。torch.jit.ScriptModule继承自torch.nn.Module,在__init__执行完毕后通过init_then_script将模块编译为 TorchScript,内部持有_actual_script_module(RecursiveScriptModule实例)作为 C++ 侧代理。编译过程只迁移参数/子模块/缓冲区的存储,不替换类上的方法,也不影响 hook 相关属性(_backward_hooks、_backward_pre_hooks等),因此大部分 hook API 在 ScriptModule 子类上行为与普通 Module 一致。测试涉及以下API
register_full_backward_hook例外的原因:该方法内部会设置self._is_full_backward_hook = True,该属性赋值被ScriptModule.__setattr__代理到 C++ 侧后,因类型不匹配(C++ 侧期望 NoneType)导致 RuntimeError。这是pytorch自身的bug。API 功能说明
register_full_backward_hook(hook, prepend=False) -> RemovableHandle
在模块上注册反向传播后置 hook。hook 签名为
hook(module, grad_input, grad_output) -> tuple[Tensor] or None,在模块梯度计算完成时被调用,可返回新的 grad_input 替代原有值。prepend=True 时 hook 在已有 hook 之前执行。返回 RemovableHandle 用于移除 hook。register_full_backward_pre_hook(hook, prepend=False) -> RemovableHandle
在模块上注册反向传播前置 hook。hook 签名为
hook(module, grad_output) -> tuple[Tensor] or None,在模块梯度计算之前被调用,可返回新的 grad_output 影响后续梯度计算。prepend=True 时 hook 在已有 hook 之前执行。返回 RemovableHandle 用于移除 hook。register_load_state_dict_pre_hook(hook) -> RemovableHandle
在 load_state_dict 调用前触发 hook。hook 签名为
hook(module, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) -> None,可用于在加载前对 state_dict 进行预处理。返回 RemovableHandle 用于移除 hook。register_load_state_dict_post_hook(hook) -> RemovableHandle
在 load_state_dict 调用后触发 hook。hook 签名为
hook(module, incompatible_keys) -> None,incompatible_keys 包含 missing_keys 和 unexpected_keys,可原地修改。返回 RemovableHandle 用于移除 hook。register_state_dict_pre_hook(hook) -> RemovableHandle
在 state_dict 调用前触发 hook。hook 签名为
hook(module, prefix, keep_vars) -> None,可用于在序列化前执行预处理。返回 RemovableHandle 用于移除 hook。register_state_dict_post_hook(hook) -> RemovableHandle
在 state_dict 调用后触发 hook。hook 签名为
hook(module, state_dict, prefix, local_metadata) -> None,可原地修改 state_dict。返回 RemovableHandle 用于移除 hook。社区用例现状
pytorch官方社区有
test/jit/test_hooks.py和test/jit/test_hooks_modules.py测试文件,但没有针对API的测试,需要补充相关测试。欢迎加入社区,感谢您对社区的贡献 🎉!