已合并
[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
已合并
Pull Request已成功合入, 合并人@ascend-robot
(感谢 黄小猛 的贡献)4月25日 添加了label:stat/needs-squash
ascend-robot
4月25日 评论:
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
/approveor/lgtm- Commenting
/approveimplies 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. 👍


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


5月18日 添加了label:approvedlgtm
ascend-robot
5月18日 评论:
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.


5月18日 合入了pull request
【合入来源】
issue连接:https://gitcode.com/Ascend/pytorch/issues/1723
【修改方案】
一、背景说明
通过
torch.jit.script()创建的模型是RecursiveScriptModule实例,与测试目的不符,所以本次测试模型均通过继承torch.jit.ScriptModule创建。torch.jit.ScriptModule继承自torch.nn.Module,在__init__执行完毕后通过torch.jit.ScriptModule .init_then_script()将模块编译为 TorchScript,内部持有torch.jit.ScriptModule ._actual_script_module(RecursiveScriptModule实例)作为 C++ 侧代理。编译过程只迁移参数/子模块/缓冲区的存储,不替换类上的方法,也不影响 hook 相关属性(torch.jit.ScriptModule ._backward_hooks、torch.jit.ScriptModule ._backward_pre_hooks等),因此大部分 hook API 在 ScriptModule 子类上行为与普通 Module 一致。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
2. torch.jit.ScriptModule.register_full_backward_pre_hook
3. torch.jit.ScriptModule.register_load_state_dict_pre_hook & torch.jit.ScriptModule.register_load_state_dict_post_hook
4. torch.jit.ScriptModule.register_state_dict_pre_hook & torch.jit.ScriptModule.register_state_dict_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_hooktorch.jit.ScriptModule.register_full_backward_pre_hooktorch.jit.ScriptModule.register_load_state_dict_pre_hooktorch.jit.ScriptModule.register_load_state_dict_post_hooktorch.jit.ScriptModule.register_state_dict_pre_hooktorch.jit.ScriptModule.register_state_dict_post_hook测试日志
pytorch 2.12
【CheckList】