☰
Muon Custom Sizing 实战:modded-nanogpt 将注意力与 MLP 参数合并进同一 reduce_scatter 调用,把 124M GPT 训练压进 150 秒
2026/10/3 2:16:37 网站建设 项目流程
  • 人工智能
  • 大模型
  • 预训练
  • 分布式训练
  • 模型优化
  • 深度学习

【免费下载链接】modded-nanogpt

NanoGPT (124M) in 90 seconds

项目地址:https://gitcode.com/GitHub_Trending/mo/modded-nanogpt
点击查看免费下载

本篇技术指南以 records/track_1_short/2025-09-23_MuonCustomSizing/README.md 为骨架,剖析 modded-nanogpt 在 124M 参数 GPT-2 训练(8 卡 H100)上的一项关键分布式优化:Muon Custom Sizing。它通过让注意力与 MLP 权重在存储层使用同一形状、在 forward 时按需 view 重塑,使两类参数可以被合并进同一条dist.reduce_scatter_tensor()调用,从而将世界纪录从 150.3843s 进一步压缩到149.6905s。读完本文,你将掌握 Muon 分布式 step 的三遍流水线、自定义参数分组算法、QKV/O 权重的 batched Newton-Schulz 正交化,以及驱动这一切的模型架构配合细节。

一、背景:Muon 优化器与分布式 step 的痛点

Muon(MomentUm Orthogonalized by Newton-Schulz)在 modded-nanogpt 中承担所有 2D 矩阵参数(attn 与 mlp 权重)的优化:它内部先跑标准 SGD-momentum,再用 Newton-Schulz 迭代把每个 2D 参数的更新替换为"最接近的正交矩阵"。其好处在于,Newton-Schulz 迭代可以稳定地在 bfloat16 下于 GPU 上执行。原文档与源码(Muon 定义)都明确提示:该优化器不应直接用于 embedding、最终输出层或任意 {0,1}-D 参数,这些应交给标准方法(如 AdamW/DistAdam);不过经验表明小规模 1D 参数交给 Muon 也很快——NS 近似相当于对梯度做幅度归一化,且这套超优化实现在小参数上的执行速度比当前 Adam 实现更快。

分布式场景下的核心矛盾是通信开销:每个 step 都需要对所有参数做一次 reduce-scatter(梯度取平均并按 rank 切分)与一次 all-gather(更新后的参数回传)。若按"形状分组",每个 param group 单独发起一次集合通信,则 12 层 GPT-2 的 attn(每层 1 个合并权重)与 mlp(每层 2 个矩阵)会产生大量小消息,通信延迟会显著拖慢训练。

二、核心思路:统一形状 + 合并 reduce_scatter

Muon Custom Sizing 的出发点非常直接,原文表述为:

The model stores all attn and mlp weights in the same shape, and then updates the view as needed on the forward pass. This enables attn and mlp weights to be contained within the samedist.reduce_scatter_tensor()call.

即:让模型把所有 attn 与 mlp 权重以相同 shape 存储,forward 时再按需 view,从而使两类参数能够进入同一条 reduce_scatter 调用。与之配套,模型架构被定制为满足:

(n_attn_layers + n_mlp_layers * 2) % 4 == 0

以保证 8 GPU 分片时零 padding(当前 12 层配置下:10 个 attn 层 + 11 个 mlp 层 ×2 = 32,恰好被 8 整除;被跳过的层见下文架构配合)。

2.1 参数分组调度(原文档 9 步调度)

原文档给出的调度共 9 步,前 4 步是 reduce-scatter(注意小参数前置):

  1. reduce scattersmear_gate(1 参数,7 个 padding 参数)
  2. reduce scatterattn_gate(10 参数,6 个 padding 参数)
  3. reduce scatterattn/mlpround 1(10 个 attn 参数,6 个 mlp 参数)
  4. reduce scatterattn/mlpround 2(16 个 mlp 参数)
  5. wait 步 1,计算第 1 组的 NS,并调度对应 all-gather
  6. wait 步 2,计算第 2 组的 NS,并调度对应 all-gather
  7. wait 步 3,计算第 3 组的 NS,并调度对应 all-gather(此时各 GPU 收到[2 ATTN, 2 ATTN, 2 ATTN, 2 ATTN, 2 ATTN, 2 MLP, 2 MLP, 2 MLP],收到 attn 类参数的 GPU 需在 NS 前先 reshape)
  8. wait 步 4,计算第 4 组的 NS,并调度对应 all-gather
  9. 等待每个 all-gather 完成,更新参数

文档同时记录了一个经验结论:"Empirically, leading with small params provides an additional 0.2s improvement."——把小的门控参数放在前面,额外带来约 0.2s 收益。这也是调度中smear_gate、attn_gate被排在最前的原因。

2.2 自定义参数分组算法

原文档给出generate_custom_param_groups的完整实现,其核心是按模块名打标签并排序,再按固定分片大小切组:

def generate_custom_param_groups(self, params): # implementation requires that a single GPU does not recieve both attn # and mlp params when a param group is split across GPUs module_ranks = { 'smear_gate': 1, # 1 param 'attn_gate': 2, # 10 params 'attn': 3, # 10 params 'mlp': 4, # 22 params } params = list(params) params.sort(key=lambda x: module_ranks.get(x.module)) idx = 0 group_sizes = [1,10,16,16] assert len(params)==sum(group_sizes) param_groups = [] for size in group_sizes: group_params = params[idx:idx+size] param_groups.append(dict(params=group_params)) idx += size return param_groups

assert len(params)==sum(group_sizes)硬性保证 43 个参数(1+10+10+22)被精确切成 4 组:[1, 10, 16, 16]。注释中的约束"一个 GPU 不应同时收到 attn 与 mlp 参数(当一个组被跨 GPU 切分时)"是正确性的关键——因为 attn 参数在 NS 前需要特殊 reshape(见 2.3),混装会导致 reshape 逻辑无法按组统一执行。

在完整训练脚本中,Muon.__init__通过开关custom_sizing=True(默认开启)选择分组策略:

def __init__(self, params, lr=0.02, weight_decay=0.01, momentum=0.95, custom_sizing=True): defaults = dict(lr=lr, weight_decay=weight_decay, momentum=momentum) if custom_sizing: param_groups = self.generate_custom_param_groups(params) else: param_groups = self.generate_standard_param_groups(params) super().__init__(param_groups, defaults)

generate_standard_param_groups则按 shape 去重分组(每个唯一 shape 一组),即 Custom Sizing 之前的旧策略,见 Muon 源码。

2.3 attn 权重的 batched NS 重塑

attn 合并权重qkvo_w的物理形状是(hdim, dim*4),但 Q/K/V/O 四部分需要独立做 Newton-Schulz 正交化,因此收到 attn 分片的 GPU 在 NS 前先做 reshape:

if getattr(params[module_idx],'module','none')=='attn': batch = 4 * original_shape[0] d1 = original_shape[1] d2 = original_shape[2] // 4 batched = batched_update_grads.view(batch, d1, d2) v_chunk = newton_schulz_triton(batched) v_chunk = v_chunk.view(original_shape) else: v_chunk = newton_schulz_triton(batched_update_grads)

即把[chunk_size, hdim, dim*4]的堆叠梯度 view 成[4*chunk_size, hdim, dim],让 Q、K、V、O 在 batch 维度上独立进入newton_schulz_triton,算完再 view 回原形状。newton_schulz_triton使用@torch.compile(dynamic=False, fullgraph=True)编译,执行 5 轮a*X + b*(X@X^T) + c*(X@X^T)@X形式的 NS 迭代(系数a,b,c = 3.4445, -4.7750, 2.0315,每次先按谱范数归一),并借助 Triton 对称矩阵乘 kernel 计算X @ X^T,从而支持 batch 矩阵的高效正交化。

三、Forward 侧的 shape 统一:模型如何配合

要让 attn 与 mlp 权重在存储层同形,模型必须做两处定制,原文档给出了 forward 代码:

3.1 注意力:合并 QKVO 权重 + forward 按需 view

self.qkvo_w = nn.Parameter(torch.empty(self.hdim, self.dim*4)) q, k, v = F.linear(x, self.qkvo_w.view(4,self.hdim, self.dim)[:3].flatten(end_dim=1).type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) y = F.linear(y, self.qkvo_w.view(4,self.hdim, self.dim)[3].type_as(y))

对应源码在 CausalSelfAttention:qkvo_w = nn.Parameter(torch.empty(self.hdim, self.dim*4))物理上是一个 2D 大矩阵,forward 时view(4, hdim, dim)拆成 Q/K/V/O 四个切片使用;初始化时前三片uniform_(-bound, bound)、输出片zero_()。qkvo_w通过self.qkvo_w.module='attn'打上模块标签,供分组算法识别。

3.2 MLP:c_fc 与 c_proj 同形

self.c_fc = nn.Parameter(torch.empty(dim, hdim)) self.c_proj = nn.Parameter(torch.empty(dim, hdim)) self.c_fc.module='mlp' self.c_proj.module='mlp'

见 MLP 定义。注释写明动机:"make both matrices have the same shape because optimizer sorts params by shape. 2 matrices × 12 layers = 24 total, which is divisible by 8 GPU world size"。c_fc/c_proj均为(dim, hdim),且都标记为mlp。

3.3 架构配合:计数必须满足整除条件

完整脚本中 12 层 Block 并非每层都有 attn 与 mlp:Block.__init__跳过layer_idx in [0, 7]的注意力(self.attn = ... if layer_idx not in [0, 7] else None),并跳过layer_idx != 0之外首层的 MLP(self.mlp = MLP(dim) if layer_idx != 0 else None),见 Block 定义。因此实际参与 Muon 的是10 个 attn 层 + 22 个 mlp 矩阵,与module_ranks注释(attn: 10 params, mlp: 22 params)一致,并满足(10 + 22) % 4 == 0,8 GPU 分片零 padding。可见 Custom Sizing 并非纯优化器改动,而是"优化器 + 架构"共同设计的结果。

四、step 的分布式流水线:三遍扫描实现

原文档只给出调度纲要,完整实现位于 Muon.step,它把 9 步调度落实为三段:

  • 第一遍(发起 reduce-scatter):对每个 param group 计算padded_num_params(向上取整到world_size的倍数),把每个参数的.gradstack 成一个大张量,多余的 padding 用torch.zeros_like(params[0].grad)补齐,然后异步发起dist.reduce_scatter_tensor(grad_chunk, stacked_grads, op=dist.ReduceOp.AVG, async_op=True)。所有组的 reduce-scatter 一次性全部发出,最大化通信重叠。
  • 第二遍(等待 → 本地 NS → 发起 all-gather):逐个组reduce_future.wait(),先对本地分片做 momentum 更新(momentum_buffer.lerp_(grad, 1-momentum)与update_grad = grad.lerp(momentum_buffer, momentum)),同时把参数复制进updated_param_chunk并施加权重衰减;随后把update_gradsstack 成 batch,按 2.3 的逻辑(attn 先 reshape)统一调用newton_schulz_triton,把 NS 结果以alpha=-eff_lr_val写回缓冲;最后异步发起dist.all_gather_into_tensor(stacked_params, updated_param_chunk, async_op=True)。
  • 第三遍(收尾):等待所有 all-gather 完成,torch.unbind后把结果逐个p.copy_(..., non_blocking=True)写回原参数。

其中有效学习率与权重衰减按组一次性算好以向量化:eff_lr_val = lr * max(1, hdim/dim)^0.5 * lr_mul,eff_weight_decay_val = lr * wd * wd_mul。训练脚本中 Muon 的实际超参是lr=0.05, momentum=0.95, weight_decay=0.0(见 优化器初始化),且 momentum 在头 300 步从 0.85 线性升温到 0.95;配套的DistAdam负责 scalar/head/embed 参数(lr=0.008, betas=(0.8, 0.95), eps=1e-8)。这也印证了 docstring 中"小参数前置额外 0.2s"的经验:小分组先完成 reduce-scatter,其 NS 与 all-gather 可以与后续大分组的通信并行。

五、实测收益与运行环境

原文档末尾给出对比数据:

  • 复跑此前纪录(rerunning prior record):150.3843s,三次样本[150.393, 150.347, 150.413]
  • 新运行时(new runtime):149.6905s,四次样本[149.686, 149.678, 149.775, 149.623]

两者同为 8 卡环境,差距约0.7s,其中小参数前置贡献约 0.2s。仓库内完整运行日志(b067b4ac-…txt 末尾)可进一步佐证:1680 步、train_time:149775ms、step_avg:89.15ms、最终val_loss:3.2792,运行环境为PyTorch 2.9.0.dev20250726+cu126、Triton 3.4.0、8× NVIDIA H100 80GB(驱动 570.148.08)。模型配置见脚本Hyperparameters:train_batch_size=2048*24*8、train_max_seq_len=128*16、num_iterations=1640、cooldown_frac=0.5,模型为vocab_size=50257, num_layers=12, num_heads=6, head_dim=128, model_dim=768的 GPT-2 规模,torch.compile(model, dynamic=False, fullgraph=True)编译后训练。训练脚本完整可复现,原始记录文件一并保留在 2025-09-23_MuonCustomSizing 目录(README 之外的 4 个.txt均为同配置的多次运行记录)。

六、经验总结与适用前提

  1. 通信次数决定分布式优化器的延迟下限。Muon Custom Sizing 的核心价值不是减少数据量,而是把多次小消息合并为更少的集合通信调用,减少启动与同步开销。
  2. 参数形状统一是一种"存储换视图"的设计:权重在内存中以合并形状存放、forward 按需 view,代价是每次使用时的 view 开销,收益是优化器侧可以整组批处理。文档与源码明确此设计"enables attn and mlp weights to be contained within the same reduce_scatter call"。
  3. 分组边界必须保证同组同构:assert强制 43 个参数精确入组、组内单一模块类型,这是 batched NS 与"attn 先 reshape"逻辑能成立的前提;跨 GPU 混装 attn/mlp 会破坏该不变式。
  4. 适用前提:该实现针对 8 GPU(world_size=8,脚本断言8 % world_size == 0)与 12 层 124M 配置调优,参数数量整除性((n_attn_layers + n_mlp_layers*2) % 4 == 0)是零 padding 的关键;换层数、换卡数时需重新推导分组方案。这是 modded-nanogpt 冲刺 90 秒 WR 系列中的一环,后续记录(如 2025-09-27_BF16CE、2025-10-24_NorMuon)沿用了同一框架并继续演进,可对照阅读以观察该技术的后续变化。
  • 人工智能
  • 大模型
  • 预训练
  • 分布式训练
  • 模型优化
  • 深度学习

【免费下载链接】modded-nanogpt

NanoGPT (124M) in 90 seconds

项目地址:https://gitcode.com/GitHub_Trending/mo/modded-nanogpt
点击查看免费下载

相关推荐

上一篇:lego v4 到 v5 迁移指南:CLI 命令、目录结构与 Go 库 API 全面升级对照
下一篇:Herdr 三种键盘模式实战指南:terminal、prefix、navigate 全覆盖

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询