mindspore-ai / mindspore-ai/hyper-parallel

分布式 Muon 优化器

Open
#283 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

分布式 Muon 优化器 — 参数分组设计文档(v2)


1. 背景与动机

Muon 是一种矩阵级优化器,其核心操作(Newton-Schulz 正交化迭代)需要对完整矩阵进行计算,而非像 AdamW 那样逐元素独立更新。Newton-Schulz 迭代本质上是对矩阵做正交化,只能作用于 2 维及以上的张量——1 维参数(如 bias、layer norm scale)不是矩阵,无法参与 Muon 运算。

在分布式训练中,模型参数以 DTensor 形式分布在多个设备上,不同参数的切分方式不同:

  • 切分不涉及最后两个维度的参数:每个设备上已经拥有完整的矩阵切片,可以直接本地计算 Newton-Schulz 迭代,无需通信
  • 切分涉及最后两个维度的参数:每个设备只持有矩阵的一部分,需要先 all-gather 收集完整矩阵,再进行正交化计算
  • 1 维参数直接报错拒绝,必须由其他优化器(如 AdamW)处理

2. 核心概念

2.1 维度校验规则

1 维参数不兼容 Muon,传入即报错。 调用方应先过滤掉 1 维参数,将其交给 AdamW 等逐元素优化器。

# 正确用法:先分离 1 维参数
muon_params = [p for p in model.parameters() if len(p.shape) >= 2]
adamw_params = [p for p in model.parameters() if len(p.shape) < 2]

no_comm, comm_groups = group_parameters_by_sharding(muon_params)  # OK
group_parameters_by_sharding(adamw_params)  # ValueError!
2.2 最后两个维度判定规则

对于一个形状为 [d0, d1, ..., d_{n-2}, d_{n-1}] 的 n 维张量(n >= 2):

条件 分组
所有 mesh 维度均为 Replicate no_comm_params
所有 Shard 的 dim 均不在 {n-2, n-1} no_comm_params
任一 Shard 的 dim ∈ {n-2, n-1} comm_params

示例

张量形状 切分方式 分组 原因
[128] 任意 报错 1 维参数,不能参与 Muon
[8, 16] Replicate(), Replicate() no_comm 完全复制
[8, 16] Shard(0), Replicate() comm dim 0 是最后两维之一
[8, 16] Replicate(), Shard(1) comm dim 1 是最后两维之一
[4, 8, 16] Shard(0), Replicate() no_comm dim 0 不在最后两维 {1, 2}
[4, 8, 16] Replicate(), Shard(2) comm dim 2 是最后维
[4, 8, 16] Replicate(), Shard(1) comm dim 1 是倒数第二维
2.3 通信域(Replicate Group)

通信域是一组 rank 的列表,这些 rank 上持有完全相同的数据副本。在 all-gather 之前,同组的 rank 之间需要交换数据以重建完整矩阵。

计算方法:对于每个 Replicate 的 mesh 维度,当前 rank 可以沿该维度"看到"一组 peer rank。当有多个 replicate 维度时,通过枚举所有 replicate 维度的坐标值笛卡尔积,得到所有可达的 rank。

示例:2×4 mesh(dp=2, tp=4),当前 rank=2

mesh 布局:        rank 编号:
dp=0: [0,1,2,3]   coord(0,0)=0, coord(0,1)=1, coord(0,2)=2, coord(0,3)=3
dp=1: [4,5,6,7]   coord(1,0)=4, coord(1,1)=5, coord(1,2)=6, coord(1,3)=7
  • [Replicate(), Shard(1)]:dp 维度 replicate,rank 2 的坐标是 (0,2),沿 dp 维度的 peer 是 rank 2 和 rank 6 → replicate_group = [2, 6]
  • [Shard(0), Shard(1)]:无 replicate 维度 → replicate_group = [2](仅自己)

3D mesh 示例:2×2×2 mesh(dp=2, cp=2, tp=2),当前 rank=0

coord(0,0,0)=0  coord(0,1,0)=2
coord(1,0,0)=4  coord(1,1,0)=6
coord(0,0,1)=1  coord(0,1,1)=3
coord(1,0,1)=5  coord(1,1,1)=7
  • [Replicate(), Replicate(), Shard(1)]:dp 和 cp 都是 replicate
    • 沿 dp(dim0) 的坐标范围: {0, 1}
    • 沿 cp(dim1) 的坐标范围: {0, 1}
    • 笛卡尔积: (0,0)→0, (0,1)→2, (1,0)→4, (1,1)→6
    • replicate_group = [0, 2, 4, 6]

3. 数据结构

┌─────────────────────────────────────────────────────────┐
│                      ShardInfo                          │
├─────────────────────────────────────────────────────────┤
│ tensor_ndim: int          # 张量维度数(>= 2)           │
│ placements: Sequence[Placement]  # 每个 mesh 维度的放置   │
│ device_mesh: DeviceMesh   # 所属设备网格                  │
│ shard_dims: set           # 被切分的张量维度集合           │
│ replicate_mesh_dims: list # Replicate 的 mesh 维度索引    │
└─────────────────────────────────────────────────────────┘

┌─────────────────────────────────────────────────────────┐
│                   CommParamGroup                        │
├─────────────────────────────────────────────────────────┤
│ params: List[DTensor]     # 组内参数列表                  │
│ shard_info: ShardInfo     # 共享的切分信息                │
│ replicate_group: List[int] # 通信域 rank 列表(已排序)   │
└─────────────────────────────────────────────────────────┘

4. 函数接口

4.1 _validate_param_ndim
def _validate_param_ndim(dtensor: DTensor) -> None

输入:一个 DTensor 参数

输出:无(校验通过则静默返回)

异常ValueError — 当参数维度 < 2 时

逻辑:检查 len(dtensor.shape) < 2,若为真则抛出包含形状信息的 ValueError

4.2 extract_shard_info
def extract_shard_info(dtensor: DTensor) -> ShardInfo

输入:一个 DTensor 参数(必须 >= 2 维)

输出:ShardInfo 对象

异常ValueError — 当参数维度 < 2 时

逻辑

  1. 调用 _validate_param_ndim 校验维度
  2. dtensor.shape 获取 tensor_ndim
  3. dtensor.placements 获取每个 mesh 维度的 placement
  4. 遍历 placements:
    • is_replicate() → 记录到 replicate_mesh_dims
    • is_shard() → 将 placement.dim 加入 shard_dims
4.3 calculate_replicate_group
def calculate_replicate_group(
    dtensor: DTensor,
    shard_info: Optional[ShardInfo] = None,
) -> List[int]

输入:DTensor 参数,可选预计算的 ShardInfo

输出:排序后的 rank 列表

逻辑

  1. 若无 replicate mesh 维度 → 返回 [device_mesh.rank]
  2. 计算当前 rank 在 mesh 中的多维坐标
  3. 对每个 replicate mesh 维度,获取其坐标取值范围
  4. 枚举所有 replicate 维度坐标的笛卡尔积
  5. 将每组坐标映射回 rank,收集到集合中
  6. 排序后返回
算法伪代码:
  coord = rank_to_coordinate(current_rank)
  dim_ranges = [range(mesh_shape[d]) for d in replicate_mesh_dims]
  result = {}
  for combo in cartesian_product(dim_ranges):
      new_coord = coord.copy()
      for dim_idx, val in zip(replicate_mesh_dims, combo):
          new_coord[dim_idx] = val
      result.add(coordinate_to_rank(new_coord))
  return sorted(result)
4.4 group_parameters_by_sharding
def group_parameters_by_sharding(
    params: List[DTensor],
) -> Tuple[List[DTensor], List[CommParamGroup]]

输入:模型所有 DTensor 参数列表(每个参数必须 >= 2 维)

输出(no_comm_params, comm_params_same_shard)

异常ValueError — 当任一参数维度 < 2 时

逻辑

  1. 遍历每个参数,调用 extract_shard_info(内含维度校验)
  2. 判断是否为 no_comm 参数(_is_no_comm_param
  3. 若是 → 加入 no_comm_params
  4. 若否 → 用 _placements_key 计算分组键
    • 键相同 → 加入已有 CommParamGroup
    • 键不同 → 创建新 CommParamGroup,计算 replicate_group
分组键生成规则:
  Shard(dim)   → ("Shard", dim)
  Replicate()  → ("Replicate",)
  Partial(op)  → ("Partial", op)

5. 使用示例

from hyper_parallel.core.optimizer.param_grouping import (
    extract_shard_info,
    calculate_replicate_group,
    group_parameters_by_sharding,
)

# 第一步:分离 1 维参数(bias、norm scale 等),交给 AdamW
all_params = list(model.parameters())
muon_params = [p for p in all_params if len(p.shape) >= 2]
adamw_params = [p for p in all_params if len(p.shape) < 2]

# 第二步:对 Muon 参数按切分方式分组
no_comm_params, comm_groups = group_parameters_by_sharding(muon_params)

# no_comm_params: 可直接本地做 Newton-Schulz 迭代
for param in no_comm_params:
    update = newton_schulz(param.to_local())

# comm_groups: 需要先 all-gather 再计算
for group in comm_groups:
    for param in group.params:
        full_param = all_gather(param, group=group.replicate_group)
        update = newton_schulz(full_param)
        # 再 reduce-scatter 回去

单独使用各函数:

# 提取单个参数的切分信息
info = extract_shard_info(some_dtensor)
print(f"shard_dims={info.shard_dims}, replicate_mesh_dims={info.replicate_mesh_dims}")

# 计算通信域
replicate_ranks = calculate_replicate_group(some_dtensor)
# 或复用已计算的 ShardInfo 避免重复计算
replicate_ranks = calculate_replicate_group(some_dtensor, shard_info=info)

6. 边界情况处理

场景 处理方式
1 维参数 抛出 ValueError,提示使用 AdamW 等其他优化器
完全复制参数 shard_dims 为空 → 归入 no_comm_params
多个维度同时切分 只要任一 shard dim 在最后两维中,就是 comm param
多个 replicate 维度 笛卡尔积枚举所有坐标组合
无 replicate 维度 返回 [current_rank],表示无通信对象
高维张量(>2 维) 最后两维 = {ndim-2, ndim-1},逻辑一致
StridedShard 继承自 Shard,is_shard() 返回 True,dim 属性可用,自动兼容
Partial 在 placements_key 中正确编码,不影响分组判定

7. 文件清单

文件 用途
hyper_parallel/core/optimizer/param_grouping.py 核心实现(3 个公开函数 + 1 个内部校验函数 + 2 个 dataclass)
hyper_parallel/core/optimizer/__init__.py 模块导出
tests/ut/core/optimizer/test_param_grouping.py 26 个单元测试(含 1 维参数报错测试)

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 168
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/168

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with hyper_parallel/core/optimizer/param_grouping.py and its interfaces, then inspect hyper_parallel/core/optimizer/init.py for exports. Run tests/ut/core/optimizer/test_param_grouping.py first to understand the 26 specified cases, including one-dimensional parameters and replicate groups. Done means the grouping, validation, rank-list calculation, and documented edge cases are covered by passing tests.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
distributed-systems, machine-learning
Issue type
Feature
Difficulty
4/5
Estimated time
3-5 days
Activity status
Active
Clarity
Clearly specified
Newbie friendliness
68/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.