已开启
[Usage]: API一致性:torch.jit.ScriptModule类下6个hook相关api一致性说明 #1723
黄小猛创建于  4月24日
黄小猛
黄小猛
4月24日 创建

在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。

环境信息

对应版本: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_moduleRecursiveScriptModule 实例)作为 C++ 侧代理。编译过程只迁移参数/子模块/缓冲区的存储,不替换类上的方法,也不影响 hook 相关属性(_backward_hooks_backward_pre_hooks 等),因此大部分 hook API 在 ScriptModule 子类上行为与普通 Module 一致。

测试涉及以下API

  • torch.jit.ScriptModule.register_full_backward_hook(不可用)
  • torch.jit.ScriptModule.register_full_backward_pre_hook
  • torch.jit.ScriptModule.register_load_state_dict_post_hook
  • torch.jit.ScriptModule.register_load_state_dict_pre_hook
  • torch.jit.ScriptModule.register_state_dict_post_hook
  • torch.jit.ScriptModule.register_state_dict_pre_hook

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.pytest/jit/test_hooks_modules.py测试文件,但没有针对API的测试,需要补充相关测试。

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
ascend-robotascend-robot成员
4月24日 添加了label:usage
黄小猛
黄小猛
4月24日 评论:
Ddinglaiping成员
4月25日 添加了label:event: api-consistency
dinglaiping成员
4月25日 评论:

新增用例文件,2.7.1直到master,共6个版本都要提交PR。新增代码中不要用中文。

likedislike
黄小猛
黄小猛
4月27日 评论:

register_full_backward_hook相关issue:https://github.com/pytorch/pytorch/issues/181473

likedislike
黄小猛黄小猛
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
chenrayraychenrayray成员
7月4日 issue类型由 Bug-Report 改变为 任务
ascend-robotascend-robot成员
7月8日 关联了看板:MindStudio ISSUE管理
TorchNPU-BotTorchNPU-Bot成员
7月28日 添加了label:bot-triaged
TorchNPU-Bot
TorchNPU-Bot成员
7月28日 评论:

检测到当前 issue 已关联 PR !34408,自动添加标签:bot-triaged

likedislike