已合并
test: add load planner api npu tests #35094
test: add load planner api npu tests #35094
已合并
zjucn创建于 5月8日
zjucn
zjucn
5月8日

【合入来源】

本 PR 覆盖以下 API:

  • torch.distributed.checkpoint.LoadPlan
  • torch.distributed.checkpoint.LoadPlanner
  • torch.distributed.checkpoint.LoadPlanner.set_up_planner
  • torch.distributed.checkpoint.LoadPlanner.create_local_plan
  • torch.distributed.checkpoint.LoadPlanner.create_global_plan
  • torch.distributed.checkpoint.LoadPlanner.finish_plan
  • torch.distributed.checkpoint.LoadPlanner.load_bytes
  • torch.distributed.checkpoint.LoadPlanner.resolve_tensor
  • torch.distributed.checkpoint.LoadPlanner.commit_tensor

说明:torch.distributed.checkpoint.LoadPlanner.resolve_bytes 是基类协议扩展点,默认实现为 NotImplementedError。当前 filesystem load 路径不依赖该方法,本 PR 聚焦验证实际加载路径中使用的 DefaultLoadPlanner.load_bytes,不对 resolve_bytes 新增专项测试。

【修改方案】

一、解决方案

针对 issue 中提出的 LoadPlan / LoadPlanner API 专项测试缺失问题,本 PR 新增 test_loadplan_api.py,以“API 级计划行为验证 + NPU state_dict 加载路径验证”的方式补齐相关测试。

本 PR 做如下修改:

  • 新增 test_loadplan_api.py,覆盖 LoadPlan / LoadPlanner 的计划生成、计划传递、bytes 加载、tensor 定位和 tensor 提交行为。
  • 使用真实 NPU tensor 构造目标 state_dict,验证 planner 在 NPU tensor、NPU tensor view、非连续 NPU tensor 场景下的行为。
  • 使用 FileSystemWriter / FileSystemReaderno_dist=True 路径,验证真实 checkpoint 读写流程可以加载 tensor 和 bytes 到 NPU state_dict
  • 通过自定义 PlanDataLoadPlanner 验证 planner_data / storage_data 在 local plan、global plan、finish plan 间的传递。
  • 通过自定义 MaterializeOnCpuLoadPlanner 验证 resolve_tensor 返回 CPU 临时 tensor 后,commit_tensor 可将数据回写到 NPU 目标 tensor。
  • 不修改 torch-npu 生产代码,不改变 API 签名、返回结构或运行时语义。

新增测试文件:

test/distributed/checkpoint/test_loadplan_api.py

二、用例完备性说明

本次新增 12 个测试用例,覆盖 LoadPlan / LoadPlanner 加载生命周期、错误边界、NPU 写入路径和跨版本兼容点。

1. test_default_load_planner_local_global_finish_plan

覆盖 API:

  • LoadPlan
  • LoadPlanner.set_up_planner
  • LoadPlanner.create_local_plan
  • LoadPlanner.create_global_plan
  • LoadPlanner.finish_plan

验证内容:

  • NPU tensor 生成 LoadItemType.TENSOR 类型 ReadItem
  • bytes 对象生成 LoadItemType.BYTE_IO 类型 ReadItem
  • tensor ReadItemdest_offsetsstorage_offsetslengths 正确。
  • DefaultLoadPlanner.create_global_plan 默认透传 local plan。
  • DefaultLoadPlanner.finish_plan 默认透传 global plan。

该用例证明默认 planner 可以基于 NPU 目标 state_dict 生成正确 LoadPlan,并保持默认 global / finish plan 语义。

2. test_default_load_planner_creates_multiple_tensor_read_items

覆盖 API:

  • LoadPlan
  • LoadPlanner.create_local_plan

验证内容:

  • 构造包含两个 tensor chunk 的 checkpoint metadata。
  • 单个 NPU 目标 tensor 生成两个 tensor ReadItem
  • 每个 ReadItemdest_indexstorage_indexstorage_offsetslengths 正确。

该用例覆盖 PyTorch 2.11.0 及 master 中 create_read_items_for_chunk_list 算法改写后的公共行为一致性。

3. test_default_load_planner_strict_and_partial_load

覆盖 API:

  • LoadPlanner.set_up_planner
  • LoadPlanner.create_local_plan

验证内容:

  • allow_partial_load=False 时,目标 state_dict 中存在 checkpoint 缺失 key 会抛出 Missing key in checkpoint
  • allow_partial_load=True 时,planner 只为 checkpoint 中存在的 key 生成读取计划。

该用例证明 set_up_planner 记录的目标 state_dict 会参与 strict / partial load 检查。

4. test_default_load_planner_size_mismatch

覆盖 API:

  • LoadPlanner.set_up_planner
  • LoadPlanner.create_local_plan

验证内容:

  • 目标 NPU tensor shape 与 checkpoint metadata shape 不一致时抛出 Size mismatch

该用例证明目标 NPU tensor 的 shape 会参与 metadata shape 校验。

5. test_resolve_tensor_returns_npu_narrow_view

覆盖 API:

  • LoadPlanner.resolve_tensor
  • LoadPlanner.commit_tensor

验证内容:

  • resolve_tensor 根据 dest_offsetslengths 返回 NPU narrow view。
  • 向该 view 写入数据后,原始 NPU tensor 对应切片被正确更新。
  • 默认空 commit_tensor 不影响 view 写回结果。

该用例证明默认 planner 可以为 NPU tensor 返回可写入 view。

6. test_resolve_tensor_handles_non_contiguous_npu_target

覆盖 API:

  • LoadPlanner.resolve_tensor
  • LoadPlanner.commit_tensor

验证内容:

  • 目标 tensor 是非连续 NPU tensor。
  • resolve_tensor 返回的 view 仍位于 NPU。
  • 写入 view 后,非连续目标 tensor 对应区域正确更新。

该用例证明默认 planner 对非连续 NPU tensor 目标仍能正确定位写入位置。

7. test_load_bytes_updates_flattened_original_state_dict

覆盖 API:

  • LoadPlanner.load_bytes

验证内容:

  • DefaultLoadPlanner 默认开启 flatten_state_dict
  • bytes ReadItem 使用 flatten FQN。
  • load_bytes 反序列化结果后写回原始嵌套 state_dict

该用例证明 flatten 场景下 bytes 对象可以正确写回原始嵌套结构。

8. test_load_bytes_updates_unflattened_state_dict

覆盖 API:

  • LoadPlanner.load_bytes

验证内容:

  • DefaultLoadPlanner(flatten_state_dict=False, flatten_sharded_tensors=False) 关闭 flatten 行为。
  • bytes 对象直接写回顶层 state_dict

该用例证明 unflatten 场景下 bytes 对象可以直接写回目标 key。

9. test_load_state_dict_accepts_custom_plan_data

覆盖 API:

  • LoadPlan.storage_data
  • LoadPlan.planner_data
  • LoadPlanner.create_local_plan
  • LoadPlanner.create_global_plan
  • LoadPlanner.finish_plan

验证内容:

  • 自定义 planner 在 local plan 中写入 planner_data
  • 自定义 planner 在 global plan 中写入 storage_data 并改写 planner_data
  • finish_plan 能接收到 global plan 写入的数据。
  • 真实 load_state_dict 流程能完成 NPU tensor 加载。

该用例证明自定义 plan 数据可以经过真实 DCP load 生命周期传递。

10. test_filesystem_metadata_version_when_supported

覆盖 API:

  • FileSystemReader.read_metadata
  • Metadata.version 兼容路径

验证内容:

  • 读取真实 filesystem checkpoint metadata。
  • 当当前 PyTorch 版本支持 metadata.version 字段时,校验其等于官方 CURRENT_DCP_VERSION
  • 当当前 PyTorch 版本不支持该字段时,仅验证 metadata 可正常读取。

该用例覆盖 PyTorch 2.9.0 及之后版本新增 metadata version 字段的兼容行为。

11. test_custom_commit_tensor_materializes_cpu_tensor_to_npu

覆盖 API:

  • LoadPlanner.resolve_tensor
  • LoadPlanner.commit_tensor

验证内容:

  • 自定义 planner 的 resolve_tensor 返回 CPU 临时 tensor。
  • StorageReader 将 checkpoint 数据写入 CPU 临时 tensor。
  • 自定义 commit_tensor 将 CPU 临时 tensor copy 回 NPU 目标 tensor。
  • 加载后的 NPU tensor 与保存前数据一致。

该用例证明自定义 commit_tensor 可以支持临时 tensor materialize 后再写回 NPU 的扩展路径。

12. test_filesystem_load_tensor_and_bytes_to_npu_state_dict

覆盖 API:

  • LoadPlanner.load_bytes
  • LoadPlanner.resolve_tensor
  • LoadPlanner.commit_tensor
  • FileSystemReader
  • FileSystemWriter

验证内容:

  • 使用真实 FileSystemWriter 保存包含 NPU tensor 和 bytes 对象的 state_dict
  • 使用真实 FileSystemReaderDefaultLoadPlanner 加载到 NPU 目标 state_dict
  • 校验 tensor 数据和 bytes 对象均正确写回。

该用例证明默认 filesystem checkpoint 读写路径可以加载 tensor 和 bytes 到 NPU state_dict

覆盖关系总结

  • plan 生成与透传:test_default_load_planner_local_global_finish_plan
  • 多 chunk 读计划:test_default_load_planner_creates_multiple_tensor_read_items
  • strict / partial load:test_default_load_planner_strict_and_partial_load
  • shape mismatch:test_default_load_planner_size_mismatch
  • NPU tensor view 写入:test_resolve_tensor_returns_npu_narrow_view
  • 非连续 NPU tensor 写入:test_resolve_tensor_handles_non_contiguous_npu_target
  • bytes flatten 写回:test_load_bytes_updates_flattened_original_state_dict
  • bytes unflatten 写回:test_load_bytes_updates_unflattened_state_dict
  • 自定义 plan data 传递:test_load_state_dict_accepts_custom_plan_data
  • metadata version 兼容:test_filesystem_metadata_version_when_supported
  • 自定义 commit 回写:test_custom_commit_tensor_materializes_cpu_tensor_to_npu
  • 真实 filesystem load:test_filesystem_load_tensor_and_bytes_to_npu_state_dict

上述用例覆盖了 issue 中列出的全部 API,并将纯 Python 计划行为、NPU tensor 定位写入行为和真实 filesystem 加载路径分开验证,便于定位问题。

三、NPU 适配说明

本次不修改 torch-npu 生产代码,仅新增测试用例。

适配验证策略:

  • 使用 torch.zeros().to(device_type)torch.arange().to(device_type) 等方式构造真实 NPU state_dict
  • 使用 FileSystemWriter / FileSystemReaderno_dist=True 单进程路径,避免引入多进程 / HCCL 依赖,保持 API 测试稳定。
  • DefaultLoadPlanner.resolve_tensor 返回的 NPU view 进行直接写入验证。
  • 对非连续 NPU tensor 目标进行直接写入验证。
  • 对多 chunk checkpoint metadata 生成的多个 ReadItem 进行直接计划校验,覆盖 resharding 读计划行为。
  • 对 2.9.0 及之后版本的 checkpoint metadata version 字段进行兼容验证。
  • 通过自定义 MaterializeOnCpuLoadPlanner 覆盖 CPU 临时 tensor 经 commit_tensor 回写 NPU 的扩展路径。
  • 使用 NPU tensor 和 bytes 混合 state_dict 验证 FileSystemReader / FileSystemWriter 与默认 planner 的组合加载路径。

经检查,PyTorch 2.7.1、2.9.0、2.10.0、2.11.0 和 master 中 LoadPlan / LoadPlanner 加载侧公共 API 的方法签名保持稳定。版本间主要差异如下:

  • PyTorch 2.9.0 起 torch.distributed.checkpoint.Metadata 新增 version 字段,FileSystemWriter 会写入 CURRENT_DCP_VERSION
  • PyTorch 2.10.0 将部分内部 assert 改为显式 AssertionError,不影响正常加载路径。
  • PyTorch 2.11.0 起 planner_helpers.create_read_items_for_chunk_list 的多 chunk 匹配算法改为 sweep-line 实现,公共行为应保持一致。
  • master 主要新增 save plan 校验与缓存相关变化,不改变本 PR 覆盖的加载侧 API 签名。

为什么不需要修改 API:

  • LoadPlan / LoadPlanner 是 DCP load 阶段的通用计划与数据定位抽象,不绑定具体设备后端。
  • 本 PR 验证这些通用接口在目标 state_dict 位于 NPU 时仍能正确生成计划、定位 tensor/view、加载 bytes,并通过自定义 planner 完成 commit 回写。
  • NPU 适配目标是验证现有语义在 NPU 下可用,而不是改变 API 签名或返回结构。

【资料变更】

不涉及

【接口变更】

不涉及

【功能验证】

在 2.7.1、2.9.0、2.10.0、2.11.0、2.12.0以及 master 版本上执行新增测试,均通过。

执行命令:

python test_loadplan_api.py

结果示例:

----------------------------------------------------------------------
Ran 12 tests in 1.367s

OK

【CheckList】

PR提交人对以下CheckList自检项进行全量自检,自检通过或不涉及,均修改 [ ] 为 [x]

likedislike
Pull Request已成功合入, 合并人@ascend-robot
(感谢 zjucn 的贡献)
zjucnzjucn
5月8日 创建了 pull request,commit c10cd424
ascend-robot
ascend-robot成员
5月8日 评论:

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 /approve or /lgtm
  • Commenting /approve implies both code review (lgtm) and intent to merge (approve)

CLA Signature Pass

zjucn, thanks for your pull request. All authors of the commits have signed the CLA. 👍

likedislike
ascend-robotascend-robot成员
5月8日 添加了label:ascend-cla/yes
ascend-robot
ascend-robot成员
5月8日 评论:

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

Protected Branch Version Release
master
v2.7.1
v2.9.0
v2.10.0
v2.11.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月8日 评论:

Ascend docs pipeline is running...

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

✅ 跳过 docs ci 检查,没有需要检查的文档文件

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

compile

likedislike
ascend-robot
ascend-robot成员
5月8日 评论:
流水线 PR-pipeline_pytorch#20171 运行中
阶段 任务名 状态 详情
编译构建 Build_X86 🕚 >>>
Build_ARM 🕚 >>>
Build_LibTorch_x86 🕚 >>>
Build_LibTorch_ARM 🕚 >>>
Build_X86_torchair 🕚 >>>
Build_ARM_torchair 🕚 >>>
patch_test 🕚 >>>
恶意代码检查 Antipoison 🕚 >>>
编码安全与规范检查 CodeCheck >>>
check_error 🕚 >>>
CodeCheck_lintrunner 🕚 >>>
开源片段检查 SCA 🕚 >>>
开发者测试 UT_X86_Part_01 🕚 >>>
UT_X86_Part_02 🕚 >>>
UT_ARM_A3_Part_01 🕚 >>>
UT_ARM_A3_Part_02 🕚 >>>
UT_DIST_X86_Part_01 🕚 >>>
UT_DIST_X86_Part_02 🕚 >>>
UT_DIST_X86_Part_03 🕚 >>>
UT_DIST_X86_Part_04 🕚 >>>
UT_inductor_Part_01 🕚 >>>
UT_inductor_Part_02 🕚 >>>
UT_inductor_Part_03 🕚 >>>
UT_inductor_Part_04 🕚 >>>
UT_ARM_A2_Part_01 🕚 >>>
UT_ARM_A2_Part_02 🕚 >>>
UT_ARM_A2_Part_03 🕚 >>>
流水线 PR-pipeline_pytorch 🕚 >>>
likedislike
ascend-robotascend-robot成员
5月8日 添加了label:ci-pipeline-running
ascend-robot
ascend-robot成员
5月8日 评论:

Ascend docs pipeline is running...

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

✅ 跳过 docs ci 检查,没有需要检查的文档文件

likedislike
ascend-robotascend-robot成员
5月8日 删除了label:docs-ci-pipeline-running
ascend-robotascend-robot成员
5月8日 添加了label:docs-ci-pipeline-success
ascend-robotascend-robot成员
5月8日 删除了label:ci-pipeline-running
ascend-robot
ascend-robot成员
5月8日 评论:
流水线 PR-pipeline_pytorch#20171 已完成
阶段 任务名 状态 详情
编译构建 Build_X86 >>>
Build_ARM >>>
Build_LibTorch_x86 >>>
Build_LibTorch_ARM >>>
Build_X86_torchair 🛑 >>>
Build_ARM_torchair 🛑 >>>
patch_test 🛑 >>>
恶意代码检查 Antipoison >>>
编码安全与规范检查 CodeCheck >>>
check_error >>>
CodeCheck_lintrunner >>>
开源片段检查 SCA >>>
开发者测试 UT_X86_Part_01 🛑 >>>
UT_X86_Part_02 🛑 >>>
UT_ARM_A3_Part_01 🛑 >>>
UT_ARM_A3_Part_02 🛑 >>>
UT_DIST_X86_Part_01 >>>
UT_DIST_X86_Part_02 >>>
UT_DIST_X86_Part_03 >>>
UT_DIST_X86_Part_04 >>>
UT_inductor_Part_01 🛑 >>>
UT_inductor_Part_02 🛑 >>>
UT_inductor_Part_03 🛑 >>>
UT_inductor_Part_04 🛑 >>>
UT_ARM_A2_Part_01 >>>
UT_ARM_A2_Part_02 >>>
UT_ARM_A2_Part_03 >>>
流水线 PR-pipeline_pytorch >>>
likedislike
ascend-robotascend-robot成员
5月8日 添加了label:ci-pipeline-passed
zjucnzjucn
5月9日 update merge request[project id: 7404318, iid: 35094, commit_id: 8ae2c0117d2f69a6fda7c927c862f5e4afc868a2] virtual merging success
zjucnzjucn
5月9日 强制推送  1 个提交:c5b038b6-test: add load planner api npu tests
zjucnzjucn
5月9日 update merge request[project id: 7404318, iid: 35094, commit_id: 291d346a872694c79e96be3f72f69957b9d35282] virtual merging success
ascend-robotascend-robot成员
5月9日 删除了label:ci-pipeline-passed
ascend-robot
ascend-robot成员
5月9日 评论:

Notification

This pull request source branch has changed, so removes the following label(s): ci-pipeline-passed.

likedislike
AtlasAccountAtlasAccount成员
5月9日 添加了label:ci-pipeline-failed
ascend-robot
ascend-robot成员
5月9日 评论:

Ascend docs pipeline is running...

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

✅ 跳过 docs ci 检查,没有需要检查的文档文件

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

compile

likedislike
ascend-robotascend-robot成员
5月9日 删除了label:ci-pipeline-failed
ascend-robot
ascend-robot成员
5月9日 评论:
流水线 PR-pipeline_pytorch#20295 运行中
阶段 任务名 状态 详情
编译构建 Build_X86 🕚 >>>
Build_ARM 🕚 >>>
Build_LibTorch_x86 🕚 >>>
Build_LibTorch_ARM 🕚 >>>
Build_X86_torchair 🕚 >>>
Build_ARM_torchair 🕚 >>>
patch_test 🕚 >>>
恶意代码检查 Antipoison 🕚 >>>
编码安全与规范检查 CodeCheck 🕚 >>>
check_error 🕚 >>>
CodeCheck_lintrunner 🕚 >>>
开源片段检查 SCA 🕚 >>>
开发者测试 UT_X86_Part_01 🕚 >>>
UT_X86_Part_02 🕚 >>>
UT_ARM_A3_Part_01 🕚 >>>
UT_ARM_A3_Part_02 🕚 >>>
UT_DIST_X86_Part_01 🕚 >>>
UT_DIST_X86_Part_02 🕚 >>>
UT_DIST_X86_Part_03 🕚 >>>
UT_DIST_X86_Part_04 🕚 >>>
UT_inductor_Part_01 🕚 >>>
UT_inductor_Part_02 🕚 >>>
UT_inductor_Part_03 🕚 >>>
UT_inductor_Part_04 🕚 >>>
UT_ARM_A2_Part_01 🕚 >>>
UT_ARM_A2_Part_02 🕚 >>>
UT_ARM_A2_Part_03 🕚 >>>
流水线 PR-pipeline_pytorch 🕚 >>>
likedislike
ascend-robotascend-robot成员
5月9日 添加了label:ci-pipeline-running
ascend-robotascend-robot成员
5月9日 删除了label:ci-pipeline-running
ascend-robot
ascend-robot成员
5月9日 评论:
流水线 PR-pipeline_pytorch#20295 已完成
阶段 任务名 状态 详情
编译构建 Build_X86 >>>
Build_ARM >>>
Build_LibTorch_x86 >>>
Build_LibTorch_ARM >>>
Build_X86_torchair 🛑 >>>
Build_ARM_torchair 🛑 >>>
patch_test 🛑 >>>
恶意代码检查 Antipoison >>>
编码安全与规范检查 CodeCheck >>>
check_error >>>
CodeCheck_lintrunner >>>
开源片段检查 SCA >>>
开发者测试 UT_X86_Part_01 🛑 >>>
UT_X86_Part_02 🛑 >>>
UT_ARM_A3_Part_01 🛑 >>>
UT_ARM_A3_Part_02 🛑 >>>
UT_DIST_X86_Part_01 >>>
UT_DIST_X86_Part_02 >>>
UT_DIST_X86_Part_03 >>>
UT_DIST_X86_Part_04 >>>
UT_inductor_Part_01 🛑 >>>
UT_inductor_Part_02 🛑 >>>
UT_inductor_Part_03 🛑 >>>
UT_inductor_Part_04 🛑 >>>
UT_ARM_A2_Part_01 >>>
UT_ARM_A2_Part_02 >>>
UT_ARM_A2_Part_03 >>>
流水线 PR-pipeline_pytorch >>>
likedislike
ascend-robotascend-robot成员
5月9日 添加了label:ci-pipeline-passed
sunyu-xuan成员
5月11日 评论:

/lgtm

likedislike
zjucnzjucn
5月11日 修改了pull request 的描述
liwei386成员
5月11日 评论:

/approve

likedislike
ascend-robotascend-robot成员
5月11日 添加了label:approvedlgtm
ascend-robot
ascend-robot成员
5月11日 评论:

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.

likedislike
ascend-robotascend-robot成员
5月11日 合入了pull request
zjucnzjucn
5月11日 修改了pull request 的描述