已关闭
[Usage]: API一致性说明:torch.nn.Module.load_state_dictapi的一致性检测 #2266
zjucn创建于 6月5日关闭于 6月15日
6月5日 添加了label:usage
6月5日 修改了issue 的描述
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
6月12日 关联了pull request:test: add nn module load_state_dict npu test patch
6月15日 修改了issue 的描述
6月15日 修改了issue 的描述
6月15日 修改了issue 的描述
6月15日 关闭了 issue
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的一致性检测 ”
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的一致性检测 ”
6月15日 添加了label:resolved
在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。
环境信息
环境信息
使用场景及问题
【涉及 API】
torch.nn.Module.load_state_dict1.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覆盖范围如下。pytorch/test/nn/test_load_state_dict.py:31-48test_load_state_dict_invalidpytorch/test/nn/test_load_state_dict.py:50-61test_load_state_dict_typepytorch/test/nn/test_load_state_dict.py:63-166test_load_state_dictmodule.前缀消费、missing key、unexpected key、strict=False、shape mismatch、部分 state_dict 加载后旧值保持pytorch/test/nn/test_load_state_dict.py:168-183test_load_state_dict_BCnum_batches_tracked时自动补默认 long bufferpytorch/test/nn/test_load_state_dict.py:185-206test_load_state_dict_childpytorch/test/nn/test_load_state_dict.py:208-221test_load_state_dict_ref_cyclepytorch/test/nn/test_load_state_dict.py:223-261test_load_state_dict_custom_save_to_state_dict/_load_from_state_dict,验证 serialized key 和子模块参数加载pytorch/test/nn/test_load_state_dict.py:263-338test_load_state_dict_assign_metaassign=True加载到 meta module,参数/buffer 对象关系、顺序和输出一致性pytorch/test/nn/test_load_state_dict.py:340-386test_load_state_dict_assign_with_optimizerassign=True与 optimizer 构造时机和 optimizer state_dict 加载关系pytorch/test/nn/test_load_state_dict.py:388-412test_load_state_dict_assign_shape_strideassign=True时同 shape 不同 stride 可加载,shape mismatch 报错pytorch/test/nn/test_load_state_dict.py:414-424test_load_state_dict_warn_assignpytorch/test/nn/test_load_state_dict.py:426-453test_load_state_dict_with_unexpected_keystrict=False返回、合法 key 前缀加后缀的边界行为pytorch/test/nn/test_load_state_dict.py:569-612TestLoadStateDictSwap.test_swap_subclassmodule_load、assign、parameter/buffer 类型保持官方测试已经覆盖
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。五个版本间的差异如下。load_state_dictv2.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(...)2. 官方测试文件差异
test/nn/test_load_state_dict.pyv2.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的核心语义,也未破坏官方测试覆盖点。