# 正确用法:先分离 1 维参数
muon_params = [p for p in model.parameters() iflen(p.shape) >= 2]
adamw_params = [p for p in model.parameters() iflen(p.shape) < 2]
no_comm, comm_groups = group_parameters_by_sharding(muon_params) # OK
group_parameters_by_sharding(adamw_params) # ValueError!
算法伪代码:
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)
分布式 Muon 优化器 — 参数分组设计文档(v2)
1. 背景与动机
Muon 是一种矩阵级优化器,其核心操作(Newton-Schulz 正交化迭代)需要对完整矩阵进行计算,而非像 AdamW 那样逐元素独立更新。Newton-Schulz 迭代本质上是对矩阵做正交化,只能作用于 2 维及以上的张量——1 维参数(如 bias、layer norm scale)不是矩阵,无法参与 Muon 运算。
在分布式训练中,模型参数以 DTensor 形式分布在多个设备上,不同参数的切分方式不同:
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):Replicateno_comm_paramsShard的 dim 均不在{n-2, n-1}中no_comm_paramsShard的 dim ∈{n-2, n-1}comm_params示例:
[128][8, 16]Replicate(), Replicate()[8, 16]Shard(0), Replicate()[8, 16]Replicate(), Shard(1)[4, 8, 16]Shard(0), Replicate()[4, 8, 16]Replicate(), Shard(2)[4, 8, 16]Replicate(), Shard(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
[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
[Replicate(), Replicate(), Shard(1)]:dp 和 cp 都是 replicate3. 数据结构
4. 函数接口
4.1
_validate_param_ndimdef _validate_param_ndim(dtensor: DTensor) -> None输入:一个 DTensor 参数
输出:无(校验通过则静默返回)
异常:
ValueError— 当参数维度 < 2 时逻辑:检查
len(dtensor.shape) < 2,若为真则抛出包含形状信息的ValueError4.2
extract_shard_infodef extract_shard_info(dtensor: DTensor) -> ShardInfo输入:一个 DTensor 参数(必须 >= 2 维)
输出:ShardInfo 对象
异常:
ValueError— 当参数维度 < 2 时逻辑:
_validate_param_ndim校验维度dtensor.shape获取tensor_ndimdtensor.placements获取每个 mesh 维度的 placementis_replicate()→ 记录到replicate_mesh_dimsis_shard()→ 将placement.dim加入shard_dims4.3
calculate_replicate_groupdef calculate_replicate_group( dtensor: DTensor, shard_info: Optional[ShardInfo] = None, ) -> List[int]输入:DTensor 参数,可选预计算的 ShardInfo
输出:排序后的 rank 列表
逻辑:
[device_mesh.rank]4.4
group_parameters_by_shardingdef 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(内含维度校验)_is_no_comm_param)no_comm_params_placements_key计算分组键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. 边界情况处理
ValueError,提示使用 AdamW 等其他优化器shard_dims为空 → 归入no_comm_params[current_rank],表示无通信对象{ndim-2, ndim-1},逻辑一致is_shard()返回 True,dim属性可用,自动兼容7. 文件清单
hyper_parallel/core/optimizer/param_grouping.pyhyper_parallel/core/optimizer/__init__.pytests/ut/core/optimizer/test_param_grouping.py