已开启
分布式 Muon 优化器 #168
yao_yf创建于  5月27日
yao_yf成员
5月27日 创建

分布式 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 维参数报错测试)
likedislike
Yyao_yf成员
5月27日 关联了pull request:Add optimizer utilities module exports for distributed Muon