已合并
docs: fix description for api torch.jit.ScriptModule hook API on NPU #35840
黄小猛创建于 5月16日
docs: fix description for api torch.jit.ScriptModule hook API on NPU #35840
已合并
黄小猛创建于 5月16日
黄小猛
黄小猛
5月16日

【合入来源】

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_hooktorch.jit.ScriptModule.register_state_dict_post_hook API昇腾支持,但是在文档中并无描述,因此需要增加。
其他API昇腾支持,且在文档中有描述,无需添加。
所有API均在以下PR中通过测试:

【接口变更】

不涉及

【功能验证】

接口 状态 说明
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
(感谢 黄小猛 的贡献)
黄小猛黄小猛
5月16日 创建了 pull request,commit 26552e9b
黄小猛黄小猛
5月16日 关联了issue:[Usage]: API一致性:torch.jit.ScriptModule类下6个hook相关api一致性说明
ascend-robotascend-robot成员
5月16日 添加了label:stat/needs-squash
ascend-robot
ascend-robot成员
5月16日 评论:

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
docs 李伟, molly123321, lyx324521 (3/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成员
5月16日 添加了label:ascend-cla/yes
ascend-robot
ascend-robot成员
5月16日 评论:

当前仓库存在以下 保护分支

Protected Branch Version Release
master
v2.7.1
v2.9.0
v2.10.0
v2.11.0
v2.12.0

评论 /sync <branch1> <branch2> ... 可将当前 PR 修改同步到其它分支(创建同步 PR):
a) 如果当前 PR 是 Open 状态,同步操作将延迟到 PR 被合并时执行
b) 如果当前 PR 已经 Merged,将立即执行同步操作

注意:

  1. /sync 命令可以指定同步到多个分支,仅最后一个 /sync 命令生效
  2. 如果创建的同步 PR 不正确,可通过向同步 PR 的源分支提交轻量级 PR 完善,或使用 /close 命令关闭
likedislike
ascend-robot
ascend-robot成员
5月16日 评论:

Ascend docs pipeline is running...

likedislike
ascend-robotascend-robot成员
5月16日 添加了label:docs-ci-pipeline-running
ascend-robot
ascend-robot成员
5月16日 评论:

✅ 文档门禁通过!

检查项 检查结果 详情
markdownlint ✅ 已通过 查看详情
link-validity-check ✅ 已通过 查看详情
resource-existence-check ✅ 已通过 查看详情
tag-closed-check ✅ 已通过 查看详情
likedislike
ascend-robotascend-robot成员
5月16日 删除了label:docs-ci-pipeline-running
ascend-robotascend-robot成员
5月16日 添加了label:docs-ci-pipeline-success
黄小猛黄小猛
5月16日 update merge request[project id: 7404318, iid: 35840, commit_id: 38db3b7ab8904e712ae3a334d221fdfd6cd90841] virtual merging success
黄小猛黄小猛
5月16日 强制推送  82 个提交:d6f3efa0-81 commits from branch v2.7.148acb95b-discard extra file
黄小猛黄小猛
5月16日 update merge request[project id: 7404318, iid: 35840, commit_id: ecfe957503404e49f1e166a5903072e3e3ccbc31] virtual merging success
ascend-robotascend-robot成员
5月16日 删除了label:stat/needs-squash
AtlasAccountAtlasAccount成员
5月16日 添加了label:ci-pipeline-failed
ascend-robot
ascend-robot成员
5月16日 评论:

Ascend docs pipeline is running...

likedislike
ascend-robotascend-robot成员
5月16日 删除了label:docs-ci-pipeline-success
ascend-robotascend-robot成员
5月16日 添加了label:docs-ci-pipeline-running
黄小猛
黄小猛
5月16日 评论:

compile

likedislike
ascend-robot
ascend-robot成员
5月16日 评论:

✅ 文档门禁通过!

检查项 检查结果 详情
markdownlint ✅ 已通过 查看详情
link-validity-check ✅ 已通过 查看详情
resource-existence-check ✅ 已通过 查看详情
tag-closed-check ✅ 已通过 查看详情
likedislike
ascend-robotascend-robot成员
5月16日 删除了label:docs-ci-pipeline-running
ascend-robotascend-robot成员
5月16日 添加了label:docs-ci-pipeline-success
ascend-robotascend-robot成员
5月16日 删除了label:ci-pipeline-failed
ascend-robotascend-robot成员
5月16日 添加了label:ci-pipeline-running
ascend-robot
ascend-robot成员
5月16日 评论:

Ascend docs pipeline is running...

likedislike
ascend-robotascend-robot成员
5月16日 删除了label:docs-ci-pipeline-success
ascend-robotascend-robot成员
5月16日 添加了label:docs-ci-pipeline-running
ascend-robot
ascend-robot成员
5月16日 评论:

✅ 文档门禁通过!

检查项 检查结果 详情
markdownlint ✅ 已通过 查看详情
link-validity-check ✅ 已通过 查看详情
resource-existence-check ✅ 已通过 查看详情
tag-closed-check ✅ 已通过 查看详情
likedislike
ascend-robotascend-robot成员
5月16日 删除了label:docs-ci-pipeline-running
ascend-robotascend-robot成员
5月16日 添加了label:docs-ci-pipeline-success
ascend-robotascend-robot成员
5月16日 删除了label:ci-pipeline-running
ascend-robotascend-robot成员
5月16日 添加了label:ci-pipeline-passed
ascend-robot
ascend-robot成员
5月16日 评论:
流水线 PR-pipeline_pytorch#22956 已完成
阶段 任务名 状态 详情
编译构建 Build_X86 >>>
Build_ARM >>>
Build_LibTorch_x86 >>>
Build_LibTorch_ARM >>>
Build_X86_torchair 🛑 >>>
Build_ARM_torchair 🛑 >>>
patch_test 🛑 >>>
恶意代码检查 Antipoison >>>
编码安全与规范检查 CodeCheck >>>
check_error >>>
开源片段检查 SCA >>>
开发者测试 UT_X86_Part_01 >>>
UT_X86_Part_02 >>>
UT_ARM_A3_Part_01 >>>
UT_ARM_A3_Part_02 >>>
UT_ARM_A2_Part_01 >>>
UT_ARM_A2_Part_02 >>>
UT_ARM_A2_Part_03 >>>
UT_inductor_Part_01_pool 🛑 >>>
UT_inductor_Part_02_pool 🛑 >>>
UT_inductor_Part_03_pool 🛑 >>>
UT_inductor_Part_04_pool 🛑 >>>
UT_DIST_ARM_Part_01 🛑 >>>
UT_DIST_ARM_Part_02 🛑 >>>
UT_DIST_ARM_Part_03 🛑 >>>
UT_DIST_ARM_Part_04 🛑 >>>
流水线 PR-pipeline_pytorch >>>
likedislike
molly123321成员
5月18日 评论:

/lgtm

likedislike
lyx324521
lyx324521成员
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: 李伟, molly123321, lyx324521.

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