mindspore-ai / mindspore-ai/hyper-parallel
分布式 Muon 优化器
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 时
逻辑:
- 调用
_validate_param_ndim校验维度 - 从
dtensor.shape获取tensor_ndim - 从
dtensor.placements获取每个 mesh 维度的 placement - 遍历 placements:
is_replicate()→ 记录到replicate_mesh_dimsis_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 列表
逻辑:
- 若无 replicate mesh 维度 → 返回
[device_mesh.rank] - 计算当前 rank 在 mesh 中的多维坐标
- 对每个 replicate mesh 维度,获取其坐标取值范围
- 枚举所有 replicate 维度坐标的笛卡尔积
- 将每组坐标映射回 rank,收集到集合中
- 排序后返回
算法伪代码:
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 时
逻辑:
- 遍历每个参数,调用
extract_shard_info(内含维度校验) - 判断是否为 no_comm 参数(
_is_no_comm_param) - 若是 → 加入
no_comm_params - 若否 → 用
_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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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