已关闭
[Usage]: API一致性说明:torch.nn.Module.load_state_dictapi的一致性检测 #2266
zjucn创建于  6月5日关闭于  6月15日
zjucn
zjucn
6月5日 创建

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

环境信息

环境信息

  • 操作系统:乌班图
  • 昇腾硬件信息:800I A2
  • CANN软件版本:8.5.0
  • 安装的对应软件版本:torch torch-npu 2.7.1、2.9.0、2.10.0、2.11.0、2.12.0

使用场景及问题

【涉及 API】

  • torch.nn.Module.load_state_dict

1.1 API 功能说明

torch.nn.Module.load_state_dict(state_dict, strict=True, assign=False) 用于将 state_dict 中保存的 parameter 和 persistent buffer 加载到当前 module 及其所有子 module。state_dict 需要是 Mapping 类型,key 与 Module.state_dict() 生成的参数名、buffer 名保持同一命名规则,子模块参数通过 . 分隔的前缀递归匹配。

该接口执行时会先复制输入 state_dict,保留其中的 _metadata 信息,然后从当前 module 开始递归调用各级子模块的 _load_from_state_dict。每个子模块只处理自身 prefix 范围内的 parameter、persistent buffer 和 extra state;加载前会执行 load state dict pre-hook,递归加载完成后会执行 post-hook。

加载 parameter 和 buffer 时,接口会校验 checkpoint 中对应值是否为 Tensor 或 Tensor-like 对象,并检查 shape 是否与当前 module 中的目标对象一致。默认 assign=False 时,接口通过 copy_ 将 checkpoint tensor 的值写入当前 module 已有的 parameter 或 buffer,保留当前 module 中 tensor 对象及其属性。assign=True 时,接口将 state_dict 中的 tensor 赋给 module 中对应的 parameter 或 buffer,保留 state_dict tensor 的属性;Parameter.requires_grad 仍以当前 module 中的设置为准。

strict=True 时,state_dict 的 key 需要与当前 module 的 state_dict() key 完全匹配。缺失 key、未预期 key、shape mismatch、非 Tensor 输入、copy/assign 过程异常等问题会汇总后抛出 RuntimeError。strict=False 时,missing key 和 unexpected key 不作为加载错误抛出,但仍会在返回值中记录。

接口返回 _IncompatibleKeys,其中包含:

  • missing_keys:当前 module 期望存在但输入 state_dict 中缺失的 key。
  • unexpected_keys:输入 state_dict 中存在但当前 module 不期望的 key。

该接口本身不负责处理 DDP module. 前缀,前缀消费由 torch.nn.modules.utils.consume_prefix_in_state_dict_if_present 等调用方逻辑完成;完成前缀处理后的 state_dict 再进入 load_state_dict 的正常 key 匹配和递归加载流程。

1.2 社区测试完备性说明

PyTorch 官方为该接口提供独立专项测试文件 pytorch/test/nn/test_load_state_dict.py,并在 pytorch/test/test_nn.py、pytorch/test/nn/test_module_hooks.py 中覆盖 extra state 和 load hook 等关联语义。pytorch/test/nn/test_load_state_dict.py 覆盖范围如下。

官方测试文件与行号 测试用例 验证详情 NPU 适配结论
pytorch/test/nn/test_load_state_dict.py:31-48 test_load_state_dict_invalid 验证 numpy array 和 tuple 等非 Tensor 输入报错 输入类型校验,与设备无关,保持官方原样
pytorch/test/nn/test_load_state_dict.py:50-61 test_load_state_dict_type 验证非 Mapping 类型 state_dict 报 TypeError Python 类型校验,与设备无关,保持官方原样
pytorch/test/nn/test_load_state_dict.py:63-166 test_load_state_dict 覆盖正常加载、DDP module. 前缀消费、missing key、unexpected key、strict=False、shape mismatch、部分 state_dict 加载后旧值保持 需要迁移到 NPU module,并验证加载后目标 parameter/buffer 保持 NPU 且数值正确
pytorch/test/nn/test_load_state_dict.py:168-183 test_load_state_dict_BC 验证 BatchNorm 旧 metadata 缺失 num_batches_tracked 时自动补默认 long buffer 需要迁移 BatchNorm 到 NPU,验证补齐 buffer 保持在 NPU
pytorch/test/nn/test_load_state_dict.py:185-206 test_load_state_dict_child 构造多层嵌套 Sequential,验证子模块 pre-hook 接收的 state_dict 范围 需要迁移递归 module 到 NPU,验证 hook 和递归范围不受设备影响
pytorch/test/nn/test_load_state_dict.py:208-221 test_load_state_dict_ref_cycle 验证 LSTM load_state_dict 不产生 Tensor 引用环 Python 引用关系测试,与 NPU 设备无关,保持官方原样
pytorch/test/nn/test_load_state_dict.py:223-261 test_load_state_dict_custom 自定义 _save_to_state_dict / _load_from_state_dict,验证 serialized key 和子模块参数加载 需要迁移自定义 module 到 NPU,验证自定义 copy 可作用于 NPU 参数
pytorch/test/nn/test_load_state_dict.py:263-338 test_load_state_dict_assign_meta 验证 assign=True 加载到 meta module,参数/buffer 对象关系、顺序和输出一致性 以 meta tensor 语义为核心,强行迁移 NPU 会改变测试目标,保持官方原样
pytorch/test/nn/test_load_state_dict.py:340-386 test_load_state_dict_assign_with_optimizer 验证 assign=True 与 optimizer 构造时机和 optimizer state_dict 加载关系 重点是 optimizer 与 assign 语义,保持官方原样;NPU optimizer 可由优化器专项测试覆盖
pytorch/test/nn/test_load_state_dict.py:388-412 test_load_state_dict_assign_shape_stride 验证 assign=True 时同 shape 不同 stride 可加载,shape mismatch 报错 需要使用 NPU module 和 NPU state_dict tensor 覆盖 assign 路径
pytorch/test/nn/test_load_state_dict.py:414-424 test_load_state_dict_warn_assign 验证从 non-meta checkpoint copy 到 meta parameter 时 warning 文案 warning/meta 语义,与 NPU 设备无关,保持官方原样
pytorch/test/nn/test_load_state_dict.py:426-453 test_load_state_dict_with_unexpected_key 验证 unexpected key、strict=False 返回、合法 key 前缀加后缀的边界行为 需要迁移 module 和 unexpected key tensor 到 NPU,验证 key 解析与设备无关
pytorch/test/nn/test_load_state_dict.py:569-612 TestLoadStateDictSwap.test_swap_subclass 验证 Tensor subclass / wrapper subclass 与 module_load、assign、parameter/buffer 类型保持 Tensor subclass 协议测试,与 NPU 设备迁移目标不同,保持官方原样

官方测试已经覆盖 load_state_dict 的核心语义和主要边界行为,测试完备性充分。本次 NPU 适配不新增测试语义,而是在官方测试结构内将关键设备相关路径迁移到 NPU。

1.3 现有 patch 覆盖情况

原 torch-npu upstream patch:

  • pytorch-npu/test_upstream/test/nn/test_load_state_dict.py.patch

原 patch 仅在文件顶部导入 torch_npu 和 torch_npu.contrib.transfer_to_npu。但官方 test/nn/test_load_state_dict.py 本身没有 .cuda()、device="cuda" 或 torch.cuda 调用,因此原 patch 不会把关键测试路径迁移到 NPU,也不能证明 NPU module 或 NPU state_dict tensor 的 load_state_dict 行为。

【版本一致性与分支 patch 说明】

已检查 PyTorch v2.7.1、v2.9.0、v2.10.0、v2.11.0、v2.12.0 中 torch.nn.Module.load_state_dict 的实现和相关测试文件。

1. API 实现差异

torch.nn.Module.load_state_dict 位于 torch/nn/modules/module.py。五个版本间的差异如下。

API 版本差异 对 NPU 测试适配的影响
load_state_dict v2.7.1 -> v2.9.0 为 docstring 和内部 load() 返回类型标注变化;v2.10.0 -> v2.11.0 将 post-hook 返回值校验从 assert out is None 改为显式 if out is not None: raise AssertionError(...) 核心加载语义未变;NPU patch 覆盖的正常加载、missing/unexpected key、shape mismatch、BC、child hook、custom、assign shape/stride、unexpected key 后缀边界均不需要调整

2. 官方测试文件差异

测试文件 版本差异 对 patch 的影响
test/nn/test_load_state_dict.py v2.7.1 -> v2.9.0 仅调整 Tensor subclass helper 中 assert 格式;v2.9.0 -> v2.10.0 新增 test_scalar_param_1d_tensor_raises,并有少量 Python 写法调整;v2.10.0 -> v2.11.0 将 Tensor subclass helper 中的 assert 改为显式 raise AssertionError;v2.11.0 -> v2.12.0 无差异 关键设备相关用例均存在且语义一致。新增 test_scalar_param_1d_tensor_raises 只验证 scalar parameter 与 1D tensor 兼容/报错行为,设备无关,可保持官方原样

上述版本差异未改变 load_state_dict 的核心语义,也未破坏官方测试覆盖点。

likedislike
ascend-robotascend-robot成员
6月5日 添加了label:usage
zjucnzjucn
6月5日 修改了issue 的描述
zjucnzjucn
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
zjucnzjucn
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
zjucnzjucn
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
zjucnzjucn
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
zjucnzjucn
6月15日 修改了issue 的描述
zjucnzjucn
6月15日 修改了issue 的描述
zjucnzjucn
6月15日 修改了issue 的描述
ascend-robotascend-robot成员
6月15日 关闭了 issue
zjucnzjucn
6月15日 修改标题为 “[Usage]: API一致性说明:torch.nn.Module.load_state_dictapi的一致性检测”,原标题为“[Usage]: API一致性说明:torch.nn.Module.compile, torch.nn.Module.cpu, torch.nn.Module.float, torch.nn.Module.load_state_dict等一系列api的一致性检测 ”
ascend-robotascend-robot成员
6月15日 添加了label:resolved