已合并
Expanded DeviceMesh Class Implementation #94
lzy0920232创建于 1月9日
Expanded DeviceMesh Class Implementation #94
已合并
lzy0920232创建于 1月9日
从refs/pull/94/head合入到master
lzy0920232
1月9日

相关的Issue

https://gitcode.com/mindspore/hyper-parallel/issues/188

原因(目的、解决的问题等)

在 Hyper-Parallel 框架中,分布式训练需要一个统一的设备拓扑抽象来管理多维并行策略。DeviceMesh 模块旨在:

  • 提供设备拓扑抽象:将物理设备组织成多维逻辑网格,支持数据并行(DP)、张量并行(TP)、流水线并行(PP)等多种并行策略
  • 统一通信组管理:为不同并行维度创建和缓存通信组,简化集合通信操作
  • 支持子网格操作:允许从原始网格中提取子网格,满足混合并行场景需求
  • 支持网格扁平化:提供 flatten 方法将多维网格转换为一维,便于全局通信操作

描述(做了什么,变更了什么)

核心实现

新增 DeviceMesh 类 (hyper_parallel/core/device_mesh.py)

  • 实现 DeviceMesh 类,提供设备网格的拓扑抽象
  • 实现 init_device_mesh 工厂函数,支持自动 rank_list 生成和缓存机制
  • 支持多种操作:
    • 子网格获取:通过 __getitem__ 方法按维度名称获取子网格
    • 通信组获取:通过 get_group 方法获取指定维度的通信组
    • 本地排名计算:通过 get_local_rank 方法获取当前进程在指定维度的本地排名
    • 网格扁平化:通过 flatten 方法将多维网格转换为一维网格
  • 实现完整的输入验证和错误处理
  • 支持通信组缓存,避免重复创建

主要功能

方法 功能 应用场景
__getitem__ 获取子网格 获取特定并行维度的设备组
get_group 获取通信组 梯度同步、张量通信
get_local_rank 获取本地排名 确定本地数据/参数分片
flatten 扁平化网格 全局通信、Checkpoint 保存
get_device_num_along_axis 获取维度设备数 计算本地张量形状
get_rank_list_along_axis 获取维度排名列表 创建自定义通信组
get_global_shape 计算全局张量形状 从分布式张量恢复全局形状

测试用例

单元测试 (tests/mindspore/ut/test_device_mesh.py)

新增 5 个单元测试用例:

测试用例 测试内容
test_init_device_mesh_basic 基础功能:自动 rank_list 生成、缓存机制
test_device_mesh_getitem_valid 子网格获取:单维度、多维度指定
test_device_mesh_get_group_valid 通信组获取:按名称、按索引
test_device_mesh_get_local_rank 本地排名计算:不同 rank 位置
test_device_mesh_flatten 网格扁平化:属性验证、通信组创建

变更文件列表

新增文件(2个):

文件 行数 说明
hyper_parallel/core/device_mesh.py 712 行 DeviceMesh 核心实现
tests/mindspore/ut/test_device_mesh.py 120 行 单元测试

代码统计:

  • 新增代码:约 832 行
  • 新增测试:5 个 UT

测试用例(新增、改动、可能影响的功能)

新增测试用例

单元测试(UT):

  • ✅ test_init_device_mesh_basic - init_device_mesh 基础功能测试
  • ✅ test_device_mesh_getitem_valid - 子网格获取测试
  • ✅ test_device_mesh_get_group_valid - 通信组获取测试
  • ✅ test_device_mesh_get_local_rank - 本地排名计算测试
  • ✅ test_device_mesh_flatten - 网格扁平化测试

测试执行方式

# 运行单元测试
cd /path/to/hyper-parallel
pytest tests/mindspore/ut/test_device_mesh.py -v

可能影响的功能

  • Layout 系统:DeviceMesh 作为 Layout 的底层支撑,Layout 依赖 DeviceMesh 进行设备管理
  • 分布式算子:分布式算子通过 DeviceMesh 获取通信组进行集合通信
  • 现有功能:纯新增功能,不影响现有模块

测试覆盖范围

  • ✅ 正常场景:覆盖所有主要方法的正常使用
  • ✅ 缓存机制:验证 DeviceMesh 和通信组的缓存复用
  • ✅ 多维网格:2D 和 3D 网格场景
  • ✅ Mock 测试:使用 Mock 隔离平台依赖

测试状态

  • ✅ 所有单元测试通过
  • ✅ 代码质量检查通过(无 Linter 错误)
likedislike
Pull Request已成功合入, 合并人@
(感谢 lzy0920232 的贡献)
AAtomGit-Bot
1月9日 指派了 CODEOWNER kisnwang 审查
AAtomGit-Bot
1月9日 添加了   mindspore-cla/yes 标签
AAtomGit-Bot
1月9日 添加了   pr-check-pass 标签
yangzhenzhang成员1月9日进行代码检视1
yangzhenzhang1月9日评论:

DeviceMesh相关的东西,可以单独弄个 device_mesh.py

likedislike
yangzhenzhang成员1月9日进行代码检视1
yangzhenzhang1月9日评论:

里面具体的类型也体现一下,比如tuple(int), tuple(str)之类的

likedislike
changzherui
changzherui成员1月9日进行代码检视1
changzherui
changzherui1月9日评论:

当前一定要指定alias_name 和 rank_list 吗?
rank_list有没有默认值,alias_name 可不可以后续指定

likedislike
changzherui
changzherui成员1月9日进行代码检视1
changzherui
changzherui1月9日评论:

这个方法的返回值是啥

likedislike
changzherui
changzherui成员1月9日进行代码检视1
changzherui
changzherui1月9日评论:

这个方法在什么场景会使用?
会返回一个新的DeviceMesh对象?alias_name是拼接后的?

likedislike
changzherui
changzherui成员1月9日进行代码检视1
changzherui
changzherui1月9日评论:

所有的类型校验后续可以抽出来个功能方法,可以参考mindspore\python\mindspore_checkparam.py

likedislike
yangzhenzhang成员1月9日进行代码检视1
tests/mindspore/ut/test_device_mesh.py
@@ -0,0 +31,4 @@
31+ # Mock communication group creation (return Mock object)
32+ mock_group = Mock()
33+ platform_mock.create_group.return_value = mock_group
34+ yield platform_mock
yangzhenzhang1月9日评论:

格式--逗号后面加空格

likedislike
yangzhenzhang成员1月9日进行代码检视1
tests/mindspore/ut/test_device_mesh.py
@@ -0,0 +107,4 @@
107+ assert device_mesh.rank_list == (0, 2, 1, 3) # Flattened from custom layout
108+ assert device_mesh.ndim == 2
109+ 
110+ def test_device_mesh_direct_construction_with_list(self, mock_platform):
yangzhenzhang1月9日评论:

这边的门禁好像对用例注释有规范,一般按这种格式:
'''
Feature:
Description:
Expectation:
'''

likedislike
liuchongming74
liuchongming74成员
1月9日 评论:

/ai-review

likedislike
Llzy0920232
6月5日 修改了pull request 的描述