已合并
test(distributed/checkpoint): Add StorageWriter storage_meta API tests #35076
test(distributed/checkpoint): Add StorageWriter storage_meta API tests #35076
已合并
Jwerr创建于 5月8日
Jwerr
Jwerr
5月8日

【合入来源】

关联 Issue:#1895
PyTorch 社区缺少对 torch.distributed.checkpoint.StorageWriter.storage_meta(实际类路径:torch.distributed.checkpoint.storage.StorageWriter.storage_meta)的直接 API 级测试,且未覆盖 NPU。本 PR 在 test/distributed/checkpoint/test_StorageWriter_api.py 补齐单卡 + 多卡 NPU 用例。

【修改方案】

一、API功能说明

torch.distributed.checkpoint.StorageWriter.storage_meta是 torch.distributed.checkpoint.StorageWriter 基类下的可选钩子,父类默认返回 None,子类可按需覆写并返回 torch.distributed.checkpoint.metadata.StorageMeta。
torch.distributed.checkpoint.metadata.StorageMeta 用于描述 checkpoint 存储元信息,包含 checkpoint_id、save_id、load_id、modules 等字段。其中 checkpoint_id 表示存储位置 ID,save_id 表示本次保存 UUID,modules 默认为空列表。
在保存流程中,writer.storage_meta() 返回的 torch.distributed.checkpoint.metadata.StorageMeta 会通过 torch.distributed.checkpoint.save 透传给 torch.distributed.checkpoint.planner.SavePlanner.set_up_planner(state_dict, storage_meta=...),用于保存规划阶段感知 writer 侧元数据。

二、测试文件 test_StorageWriter_api.py 验证内容

1. 基类默认行为验证
   定义 _MinimalStorageWriter 作为 torch.distributed.checkpoint.storage.StorageWriter 最小实现类,仅实现抽象方法,不覆写 torch.distributed.checkpoint.storage.StorageWriter.storage_meta,验证:
* torch.distributed.checkpoint.storage.StorageWriter.storage_meta 可调用;
* 父类默认 storage_meta() 返回 None。
2. FileSystemWriter 覆写行为验证
   通过 torch.distributed.checkpoint.FileSystemWriter 验证真实 writer 的 storage_meta 返回值:
* 返回值类型为 torch.distributed.checkpoint.metadata.StorageMeta;
* checkpoint_id 与初始化 checkpoint 目录一致;
* save_id 为合法 UUID;
* modules 默认为空列表;
* reset(checkpoint_id=new_dir) 后 checkpoint_id 更新,save_id 变化。
3. 保存流程验证
   通过 torch.distributed.checkpoint.save 验证 storage_meta 在实际保存流程中的可用性:
* 使用 NPU Tensor 执行 torch.distributed.checkpoint.save(..., no_dist=True) 后,writer.storage_meta() 仍返回合法 torch.distributed.checkpoint.metadata.StorageMeta;
* 空 state_dict 保存后,writer.storage_meta() 仍保持有效;
* save -> reset -> save 后 save_id 发生变化,验证 writer 生命周期更新有效。
4. SavePlanner 透传验证
   定义 _CapturingPlanner 继承 torch.distributed.checkpoint.default_planner.DefaultSavePlanner,在 set_up_planner 中捕获 storage_meta,验证:
* torch.distributed.checkpoint.save 会将 writer.storage_meta() 返回值透传给 torch.distributed.checkpoint.planner.SavePlanner.set_up_planner;
* planner 捕获到的 checkpoint_id、save_id 与 writer.storage_meta() 保持一致。
5. 多 NPU 分布式验证
   TestStorageMetaDistributed 继承 torch.testing._internal.distributed._shard.sharded_tensor.ShardedTensorTestBase,通过 torch_npu.testing.common_distributed.with_comms 初始化 HCCL 进程组,并通过 torch_npu.testing.common_distributed.skipIfUnsupportMultiNPU(2) 限定 2 卡 NPU 环境运行,验证:
* rank 0 创建共享 checkpoint 目录,并通过 torch.distributed.broadcast_object_list 广播给其他 rank;
* 分布式 torch.distributed.checkpoint.save 后,各 rank 的 checkpoint_id 通过 torch.distributed.all_gather_object 收集并保持一致;
* 分布式 save -> reset -> save 后 save_id 发生变化;
* 分布式保存后 writer.storage_meta() 仍返回 torch.distributed.checkpoint.metadata.StorageMeta,且 save_id 为合法 UUID。

三、NPU适配

torch.distributed.checkpoint.storage.StorageWriter.storage_meta 和 torch.distributed.checkpoint.metadata.StorageMeta 均属于 Python 层 checkpoint 元数据抽象,不涉及 NPU 算子、NPU kernel、设备内存管理或 HCCL 通信协议本身,因此 API 本身无需针对 NPU 修改。
本 PR 主要在测试用例层面验证 NPU 场景可用性:
* 单卡用例继承 torch_npu.testing.testcase.TestCase,使用 torch.Tensor.npu() 构造 NPU Tensor 后进入 torch.distributed.checkpoint.save 流程;
* 多卡用例通过 torch_npu.testing.common_distributed.with_comms 初始化 HCCL 进程组;
* 通过 torch_npu.testing.common_distributed.skipIfUnsupportMultiNPU(2) 限定多 NPU 环境,避免设备数量不足时误失败;
* 通过 torch.distributed.broadcast_object_list 保证各 rank 使用同一 checkpoint 目录;
* 通过 torch.distributed.all_gather_object 验证 checkpoint_id 跨 rank 一致;
* 测试结束后通过 destroy_pg 释放分布式进程组资源。

【资料变更】

不涉及

【接口变更】

不涉及

【功能验证】

在2.7.1 2.9.0 2.10.0 2.11.0 2.12.0以及master版本上执行该用例,均通过,本分支对应2.10.0日志(完整版在评论区)如下:

/root/work/test_StorageWriter_api.py:303: FutureWarning: `save_state_dict` is deprecated and will be removed in future versions.Please use `save` instead.
  save_state_dict(
.
----------------------------------------------------------------------
Ran 8 tests in 156.320s

OK (skipped=6)

【CheckList】

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

likedislike
Pull Request已成功合入, 合并人@ascend-robot
(感谢 Jwerr 的贡献)
JwerrJwerr
5月8日 创建了 pull request,commit 4cd30713
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

Jwerr, 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
此处折叠了220条消息 查看更多
sunyu-xuan成员
5月12日 评论:

/lgtm

likedislike
liwei386成员
5月12日 评论:

/approve

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

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月12日 合入了pull request