已合并
[feat] Add test cases for verifying torch.jit.ScriptModule hook API on NPU #34411
黄小猛创建于 4月25日
[feat] Add test cases for verifying torch.jit.ScriptModule hook API on NPU #34411
已合并
黄小猛创建于 4月25日
黄小猛
黄小猛
4月25日

【合入来源】

issue连接:https://gitcode.com/Ascend/pytorch/issues/1723

【修改方案】

一、背景说明

  1. 通过torch.jit.script()创建的模型是RecursiveScriptModule实例,与测试目的不符,所以本次测试模型均通过继承torch.jit.ScriptModule创建。

  2. torch.jit.ScriptModule 继承自 torch.nn.Module,在 __init__ 执行完毕后通过 torch.jit.ScriptModule .init_then_script() 将模块编译为 TorchScript,内部持有 torch.jit.ScriptModule ._actual_script_moduleRecursiveScriptModule 实例)作为 C++ 侧代理。编译过程只迁移参数/子模块/缓冲区的存储,不替换类上的方法,也不影响 hook 相关属性(torch.jit.ScriptModule ._backward_hookstorch.jit.ScriptModule ._backward_pre_hooks 等),因此大部分 hook API 在 ScriptModule 子类上行为与普通 Module 一致。

  3. torch.jit.ScriptModule .register_full_backward_hook 例外的原因:该方法内部会设置 self._is_full_backward_hook = True,该属性赋值被 ScriptModule.__setattr__ 代理到 C++ 侧后,因类型不匹配(C++ 侧期望 NoneType)导致 RuntimeError。这是pytorch自身的bug。

二、API 功能说明

1. torch.jit.ScriptModule.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。

2. torch.jit.ScriptModule.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。

3. torch.jit.ScriptModule.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。

4. torch.jit.ScriptModule.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。

5. torch.jit.ScriptModule.register_state_dict_pre_hook(hook) -> RemovableHandle

在 state_dict 调用前触发 hook。hook 签名为 hook(module, prefix, keep_vars) -> None,可用于在序列化前执行预处理。返回 RemovableHandle 用于移除 hook。

6. torch.jit.ScriptModule.register_state_dict_post_hook(hook) -> RemovableHandle

在 state_dict 调用后触发 hook。hook 签名为 hook(module, state_dict, prefix, local_metadata) -> None,可原地修改 state_dict。返回 RemovableHandle 用于移除 hook。

三、测试用例说明

1. torch.jit.ScriptModule.register_full_backward_hook

用例 说明
test_register_full_backward_hook_raises 验证调用 register_full_backward_hook 时抛出 RuntimeError

2. torch.jit.ScriptModule.register_full_backward_pre_hook

用例 说明
test_register_full_backward_pre_hook_called 验证 hook 在反向传播时被调用
test_register_full_backward_pre_hook_modify_grad 验证 hook 返回修改后的 grad_output 可影响梯度计算(返回全零梯度,输入梯度为零)
test_register_full_backward_pre_hook_prepend 验证 prepend=True 时 hook 在已有 hook 之前执行
test_register_full_backward_pre_hook_remove 验证 handle.remove() 后 hook 不再触发

3. torch.jit.ScriptModule.register_load_state_dict_pre_hook & torch.jit.ScriptModule.register_load_state_dict_post_hook

用例 说明
test_load_state_dict_pre_hook_fires_before_module_and_post_hook 验证完整时序 pre_hook -> module(load_state_dict) -> post_hook,post_hook 先注册证明调用顺序与注册顺序无关;pre_hook 修改 state_dict 为全零,post_hook 检查权重已加载为零,证明 module 在 pre 和 post 之间执行
test_register_load_state_dict_pre_hook_called 验证 pre_hook 在 load_state_dict 时被调用,接收到正确的 prefix 参数
test_register_load_state_dict_pre_hook_with_module 验证 pre_hook 接收到的 module 参数就是当前模型实例
test_register_load_state_dict_pre_hook_remove 验证 handle.remove() 后 pre_hook 不再触发
test_register_load_state_dict_post_hook_called 验证 post_hook 在 load_state_dict 后被调用
test_register_load_state_dict_post_hook_with_module 验证 post_hook 接收到的 module 参数就是当前模型实例
test_register_load_state_dict_post_hook_remove 验证 handle.remove() 后 post_hook 不再触发

4. torch.jit.ScriptModule.register_state_dict_pre_hook & torch.jit.ScriptModule.register_state_dict_post_hook

用例 说明
test_state_dict_pre_hook_fires_before_module_and_post_hook 验证完整时序 pre_hook -> module(state_dict) -> post_hook,post_hook 先注册证明调用顺序与注册顺序无关;post_hook 检查 state_dict 已包含模型参数,证明 module 在 pre 和 post 之间执行
test_register_state_dict_pre_hook_called 验证 pre_hook 在 state_dict 时被调用,接收到正确的 prefix 参数
test_register_state_dict_pre_hook_with_module 验证 pre_hook 接收到的 module 参数就是当前模型实例
test_register_state_dict_pre_hook_remove 验证 handle.remove() 后 pre_hook 不再触发
test_register_state_dict_post_hook_called 验证 post_hook 在 state_dict 后被调用,接收到正确的 prefix 参数
test_register_state_dict_post_hook_with_module 验证 post_hook 接收到的 module 参数就是当前模型实例
test_register_state_dict_post_hook_remove 验证 handle.remove() 后 post_hook 不再触发

【资料变更】

torch.jit.ScriptModule.register_load_state_dict_pre_hook、torch.jit.ScriptModule.register_state_dict_post_hook API昇腾支持,但是在文档中并无描述,因此需要增加。
其他API昇腾支持,且在文档中有描述,无需添加。
文档修改PR链接:https://gitcode.com/Ascend/pytorch/pull/35840

【接口变更】

不涉及

【功能验证】

接口 状态 说明
torch.jit.ScriptModule.register_full_backward_hook 不可用 调用时抛出 RuntimeError
torch.jit.ScriptModule.register_full_backward_pre_hook 可用 正常注册并触发,支持修改 grad_output、prepend 顺序控制、handle 移除
torch.jit.ScriptModule.register_load_state_dict_pre_hook 可用 在 load_state_dict 前触发,可接收 module 参数,支持 handle 移除
torch.jit.ScriptModule.register_load_state_dict_post_hook 可用 在 load_state_dict 后触发,可接收 module 参数,支持 handle 移除
torch.jit.ScriptModule.register_state_dict_pre_hook 可用 在 state_dict 前触发,可接收 module 参数,支持 handle 移除
torch.jit.ScriptModule.register_state_dict_post_hook 可用 在 state_dict 后触发,可接收 module 参数,支持 handle 移除

测试日志

========================================
Python代码多环境测试开始
测试脚本: /root/torch-2.11/test_script_module_hooks.py
测试时间: Fri Apr 24 06:37:57 AM UTC 2026
========================================

检测到已激活的虚拟环境,先取消激活...
----------------------------------------
测试环境: torch-2.7.1
虚拟环境: /root/torch-2.7.1/.venv/bin/activate
日志文件: /root/torch-2.11/logs/test_2.7.1_20260424_063757.log
----------------------------------------
测试环境: torch-2.7.1
退出码: 0
状态: 成功 ✅
详细日志: /root/torch-2.11/logs/test_2.7.1_20260424_063757.log

完成测试: torch-2.7.1 (退出码: 0)

----------------------------------------
测试环境: torch-2.8
虚拟环境: /root/torch-2.8/.venv/bin/activate
日志文件: /root/torch-2.11/logs/test_2.8_20260424_063757.log
----------------------------------------
测试环境: torch-2.8
退出码: 0
状态: 成功 ✅
详细日志: /root/torch-2.11/logs/test_2.8_20260424_063757.log

完成测试: torch-2.8 (退出码: 0)

----------------------------------------
测试环境: torch-2.9
虚拟环境: /root/torch-2.9/.venv/bin/activate
日志文件: /root/torch-2.11/logs/test_2.9_20260424_063757.log
----------------------------------------
测试环境: torch-2.9
退出码: 0
状态: 成功 ✅
详细日志: /root/torch-2.11/logs/test_2.9_20260424_063757.log

完成测试: torch-2.9 (退出码: 0)

----------------------------------------
测试环境: torch-2.10
虚拟环境: /root/torch-2.10/.venv/bin/activate
日志文件: /root/torch-2.11/logs/test_2.10_20260424_063757.log
----------------------------------------
测试环境: torch-2.10
退出码: 0
状态: 成功 ✅
详细日志: /root/torch-2.11/logs/test_2.10_20260424_063757.log

完成测试: torch-2.10 (退出码: 0)

----------------------------------------
测试环境: torch-2.11
虚拟环境: /root/torch-2.11/.venv/bin/activate
日志文件: /root/torch-2.11/logs/test_2.11_20260424_063757.log
----------------------------------------
测试环境: torch-2.11
退出码: 0
状态: 成功 ✅
详细日志: /root/torch-2.11/logs/test_2.11_20260424_063757.log

完成测试: torch-2.11 (退出码: 0)

========================================
测试完成总结
总测试环境数: 5
完成时间: Fri Apr 24 06:39:05 AM UTC 2026

成功测试数: 5
失败测试数: 0
所有详细日志保存在: /root/torch-2.11/logs
主日志文件: /root/torch-2.11/logs/test_results_20260424_063757.log
========================================

pytorch 2.12

../root/.local/share/uv/python/cpython-3.12.13-linux-aarch64-gnu/lib/python3.12/multiprocessing/popen_fork.py:66: DeprecationWarning: This process
     (pid=3985006) is multi-threaded, use of fork() may lead to deadlocks in the child.
       self.pid = os.fork()
     [W515 03:02:03.790319610 NPUCachingAllocator.cpp:199] Warning: The current CANN and Soc versions require processing for 32 padding size, with memory
     allocation. (function operator())
     [W515 03:02:11.756102926 IscloseKernelNpu.cpp:34] Warning: Device do not support double dtype of rtol and atol now, dtype cast repalce with float.
     (function operator())
     ....................
     ----------------------------------------------------------------------
     Ran 21 tests in 18.055s

     OK

【CheckList】

PR提交人对以下CheckList自检项进行全量自检,自检通过或不涉及,均修改 [ ] 为 [x]

likedislike
Pull Request已成功合入, 合并人@ascend-robot
(感谢 黄小猛 的贡献)
黄小猛黄小猛
4月25日 创建了 pull request,commit 8362d9d9
ascend-robotascend-robot成员
4月25日 添加了label:stat/needs-squash
ascend-robot
ascend-robot成员
4月25日 评论:

Thanks for your pull-request.
The full list of commands accepted by me can be found at here
You can get sig-info at here


PR Approval Progress

Congratulations! All modules have met the lgtm and approve requirements.

Module Approval Details

module lgtm status approve status
test sunyu-xuan, 李伟 (2/2) 李伟 (1/1)

💡 Tip:

  • Committer can comment /approve or /lgtm
  • Commenting /approve implies both code review (lgtm) and intent to merge (approve)

CLA Signature Pass

hxm_, thanks for your pull request. All authors of the commits have signed the CLA. 👍

likedislike
ascend-robotascend-robot成员
4月25日 添加了label:ascend-cla/yes
此处折叠了137条消息 查看更多
sunyu-xuan成员
5月18日 评论:

/lgtm

likedislike
liwei386成员
5月18日 评论:

/approve

likedislike
ascend-robotascend-robot成员
5月18日 添加了label:approvedlgtm
ascend-robot
ascend-robot成员
5月18日 评论:

Review Guide

This pull-request passes review.
Committers who wrote a comment of /approve are: 李伟.
Reviewers who wrote a comment of /lgtm are: 李伟, sunyu-xuan.

likedislike
ascend-robotascend-robot成员
5月18日 合入了pull request