已关闭
[Usage]: test目录下面的test_models.py继承自pytorch社区,但是里面并没有专门验证torch.nn.Module,torch.nn.Module.state_dict,torch.nn.ModuleDict,torch.nn.ModuleList的测试用例,需要补充 #1575
dinglaiping创建于  3月13日关闭于  3月19日
dinglaiping成员
3月13日 创建

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

环境信息

  • 操作系统 乌班图
  • 昇腾硬件信息 910B4
  • CANN软件版本 8.5.0
  • 安装的对应软件版本 pytorch torch-npu,2.6.0及以上版本

API功能

1.1.1 torch.nn.Module
PyTorch 所有神经网络模块的基类,所有网络层、模型均继承该类实现:

  1. 封装参数(Parameter)、缓冲区(Buffer)的注册与管理,自动跟踪可训练参数;
  2. 支持to()、npu()等方法,实现模块参数 / 缓冲区的设备(CPU/CUDA/NPU)与数据类型统一迁移;
  3. 提供forward()方法抽象,支撑自定义计算逻辑与前后向传播联动;
  4. 为状态序列化提供底层支撑,适配模型断点续训、部署等场景;
  5. 支持子模块嵌套注册,自动递归管理子模块的参数、设备与状态。
    1.1.2 torch.nn.Module.state_dict
    模块状态管理核心方法,用于提取模块及子模块的状态键值对字典:
  6. 完整收集可训练参数(weight/bias)、非训练缓冲区(如BatchNorm运行均值/方差)
  7. 提取的张量与模块当前设备、数据类型完全一致,无隐式转换;
  8. 与load_state_dict()双向兼容,支持跨设备、跨环境的状态迁移;
  9. 轻量无侵入,提取后模块可正常执行前后向传播;
  10. 键名按模块嵌套层级命名(如sub_module.linear.weight),确保状态与子模块一一对应。
    1.1.3 torch.nn.ModuleDict
    键值对型子模块容器,以字符串为索引管理子模块:
  11. 支持字符串索引、新增(md["key"] = sub_module)、删除(del)、遍历(items())等核心操作;
  12. 容器内子模块自动注册到主模块,参与参数管理、设备迁移与状态序列化;
  13. 主模块执行设备 / 类型迁移时,容器内子模块自动递归迁移,保持状态一致;
  14. 支持嵌套其他容器,适配复杂层级化网络构建。
    1.1.4 torch.nn.ModuleList
    有序型子模块容器,以数字为索引管理子模块:
  15. 支持数字索引、append()新增、insert()插入、pop()删除、enumerate()遍历等列表式操作;
  16. 容器内子模块自动注册到主模块,参与统一参数、设备、状态管理;
  17. 主模块的设备迁移、状态序列化自动覆盖容器内子模块,无需单独操作;
  18. 支持直接 for 循环遍历,无缝集成到前向传播逻辑。

社区用例现状

torch-npu 官方社区的test_models.py通用测试用例,不包含这4个API测试的:torch.nn.Module(基础模块类),torch.nn.Module.state_dict(状态字典方法),torch.nn.ModuleDict(模块字典容器),torch.nn.ModuleList(模块列表容器)。官网测试文件测试的是:通过@modules(module_db)装饰器参数化测试 module_db中定义的神经网络层(如Conv2d、Linear、ReLU等),测试内容包括:test_forward - 前向传播,test_factory_kwargs - 工厂参数,test_pickle - 序列化/反序列化,test_grad / test_gradgrad - 梯度检查test_cpu_gpu_parity - CPU/GPU结果一致性。

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

likedislike
ascend-robotascend-robot成员
3月13日 添加了label:usage
Ddinglaiping成员
3月13日 修改了issue 的描述
Ddinglaiping成员
3月13日 关联了pull request:test(nn): add test for models api:torch.nn.Module, torch.nn.Module.state_dict, torch.nn.ModuleDict, torch.nn.ModuleList
Ddinglaiping成员
3月13日 关联了pull request:test(nn): add test for models api:torch.nn.Module, torch.nn.Module.state_dict, torch.nn.ModuleDict, torch.nn.ModuleList
Ddinglaiping成员
3月13日 关联了pull request:test(nn): add test for models api:torch.nn.Module, torch.nn.Module.state_dict, torch.nn.ModuleDict, torch.nn.ModuleList
Ddinglaiping成员
3月13日 关联了pull request:test(nn): add test for models api:torch.nn.Module, torch.nn.Module.state_dict, torch.nn.ModuleDict, torch.nn.ModuleList
Ddinglaiping成员
3月13日 关联了pull request:test(nn): add test for models api:torch.nn.Module, torch.nn.Module.state_dict, torch.nn.ModuleDict, torch.nn.ModuleList
Ddinglaiping成员
3月13日 关联了pull request:test(nn): add test for models api:torch.nn.Module, torch.nn.Module.state_dict, torch.nn.ModuleDict, torch.nn.ModuleList
huangyunlong成员
3月18日 评论:
ascend-robotascend-robot成员
3月19日 关闭了 issue
ascend-robotascend-robot成员
3月19日 添加了label:resolved
Ddinglaiping成员
4月10日 添加了label:event: api-consistency