跳转至

CMuon: Accelerating and Stabilizing Diffusion Transformer Training via Chunked Momentum Orthogonalization

会议: ECCV2026
Paper: https://eccv.ecva.net/virtual/2026/poster/5960
PDF: https://media.eventhosts.cc/Conferences/ECCV2026/pdfs/14727.pdf
领域: 训练效率 / 扩散模型
关键词: 分块动量正交化、子空间干扰、扩散Transformer、功能分块、学习率缩放

一句话总结

CMuon 将 DiT 中融合存储的 QKV、FFN gate/up 和 AdaLN 投影按功能分块后独立正交化动量,使 675M DiT-XL 在 ImageNet-1K 256×256 上用 200 epochs 达到 FID 1.18,优于 AdamW 训练 400 epochs 的 1.21,且收益不局限于训练早期。

研究背景与动机

Muon 的吸引力在于它不是逐元素调整梯度,而是对二维权重的动量矩阵做正交化,改变更新的奇异值谱,使原本相对弱的方向不至于一直被强方向压制。但在扩散 Transformer 上,早期损失下降更快并不自动意味着最终生成质量更好:本文的 DiT-B 实验中,Muon 与 AdamW 训练到 400 epochs 时 FID 都是 2.78。需要解决的不是“能不能让训练开始得更快”,而是为何这种优势到了后期会消失。

作者把问题追到一个容易被视为纯工程细节的操作:将不同功能的投影拼成一个大矩阵以提高计算效率。Q、K、V 可以由一次线性运算得到,AdaLN 的缩放、平移与门控也可以融合计算;然而,当优化器以整个张量为正交化单位时,这种存储边界同时变成了统计耦合边界。不同投影虽然共享输入维度,却未必具有一致的梯度主方向,用统一预条件器可能让一个分支的统计改变另一个分支的更新。

核心 idea:保留高效的融合前向计算,但在优化器中恢复参数的功能边界,先分块、再独立正交化,并显式控制分块后的更新尺度,让实现上的拼接不再强制不同功能共享优化几何。

方法详解

整体框架

CMuon 改的是优化步骤,不是生成模型的前向架构,也没有增加表示对齐损失。每轮先对 flow-matching 目标反向传播并形成 Nesterov 型动量更新,再对指定二维权重依次执行功能分块、独立正交化和尺度校准,最后把块拼回原形状更新权重;其余参数由 AdamW 处理。

这里的“独立”仅指正交化统计不跨功能块混合,不意味着整个网络的学习彼此独立。各块仍在同一模型、同一损失下联合训练,反向传播仍会传递网络内的相互影响。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    input["Flow Matching<br/>反向传播"] --> momentum["二维投影的<br/>Nesterov 动量"]
    input --> adam["其他参数<br/>AdamW"]
    momentum --> split["功能分块"]
    split --> orth["独立正交化"]
    orth --> scale["尺度校准"]
    scale --> update["拼接并更新权重"]
    adam --> update

关键设计

1. 功能分块:把计算融合与优化分组分开

任意切成相同大小的小块不是本文的关键,块必须对应原来具有不同语义的投影。以隐藏维度为 1152 的 DiT-XL 为例,QKV 权重为 [3456, 1152],沿第 0 维拆成 3 个 [1152, 1152] 矩阵;FFN gate/up 的 [6144, 1152] 拆成 2 个 [3072, 1152];AdaLN 的 [6912, 1152] 拆成 6 个 [1152, 1152],分别对应注意力和 FFN 分支的调制与门控。三类矩阵均沿较长维度切分,参数数量和前向计算含义不变。

这一步之所以重要,是因为融合张量的行区间并不是可任意互换的特征组。Q、K、V 分别影响注意力匹配和内容聚合;FFN 的 gate 与 up 也承担不同计算角色。CMuon 允许这些角色在正交化时使用自己的统计,而不是让一个梯度较强的功能块通过共享预条件器影响所有块。实现时需要依据具体模型的参数布局识别边界,不能只凭矩阵“看起来很长”就机械切分。

2. 独立正交化:解除共享预条件器造成的方向干扰

理解 Muon 可以先看精确极分解的理想形式:若动量矩阵的奇异值分解为 \(M=U\Sigma V^\top\),则 \(\mathrm{Orth}(M)=UV^\top\)。它主要保留方向结构,重新塑造奇异值尺度;实际使用 Newton–Schulz 迭代近似这一操作,而非每步显式执行 SVD。对矩形矩阵,更准确的说法是半正交,并非所有行列都能同时正交。

作者用纵向堆叠的梯度块说明耦合来源。下式是对正文第 3.3 节的规范化转写,假设 Gram 矩阵可逆;实际优化时应把输入理解为带动量的更新矩阵,而不是绕过动量直接处理原始梯度:

\[ U_i^{\mathrm{fused}}=G_i\left(\sum_{j=1}^{N}G_j^\top G_j\right)^{-1/2}, \qquad U_i^{\mathrm{chunk}}=G_i(G_i^\top G_i)^{-1/2}. \]

差别就在逆平方根里的统计范围。融合形式让每个块共享所有块的 Gram 矩阵之和;当主方向不对齐时,一个块会改变另一个块的缩放方向。分块形式不再混合这部分统计,随后仍把更新按原次序拼接。这个推导解释了可能的干扰机制,但不是“不论数据和架构,分块都优于融合”的收敛定理。

具体更新沿用 Algorithm 1:动量缓冲区先按 \(M\leftarrow\mu M+G\) 累积,再构造 \(G+\mu M\) 作为正交化输入,然后才切块。不能先对每个原始梯度块正交化再累积动量,因为正交化是非线性操作,顺序交换会得到不同算法。缓存未收录所引用的附录 Algorithm 3,因此不能据此补写 Newton–Schulz 的具体迭代系数和步数。

3. 尺度校准:区分方向收益与步长收益

分块会改变矩阵形状,也会改变半正交矩阵的范数。如果直接沿用融合大矩阵的缩放,实验中更快的收敛可能部分来自更大的步长。CMuon 默认采用 Moonlight 缩放,按每块形状计算 \(\alpha_c=0.2\sqrt{\max(d_{\mathrm{out}},d_{\mathrm{in}})}\),并据此应用权重衰减和参数更新。这使 AdamW 的基础学习率更容易沿用,但不意味着所有缩放规则无需调参。

对纵向堆叠的 \(N\) 个块,若每块满足 \(d_{\mathrm{out}}\ge d_{\mathrm{in}}\),在精确半正交假设下,融合与分块更新满足下列关系;这里 \(G\) 表示完整堆叠矩阵,\(\alpha=0.2\sqrt{Nd_{\mathrm{out}}}\)

\[ \|\alpha\,\mathrm{Orth}(G)\|_F =\left\|\alpha_c \begin{bmatrix}\mathrm{Orth}(G_1)\\ \vdots\\ \mathrm{Orth}(G_N)\end{bmatrix}\right\|_F =0.2\sqrt{Nd_{\mathrm{out}}d_{\mathrm{in}}}. \]

因此默认分块不是单纯扩大总更新幅度,而是重新分配块间的更新范数和方向。论文另提供可选的 rescale 开关,再乘 \(\sqrt{N_{\mathrm{chunk}}}\) 来加快早期训练;开启后就不再具有上面的全局等范数解释。作者对未分块 Muon 的对应层也施加相同 rescale,以区分“步长变大”和“统计解耦”两种作用。

一个完整示例

以 DiT-XL 的 QKV 层为例:反向传播产生 [3456, 1152] 梯度,优化器先在原形状上更新动量,再把 Nesterov 输入沿行拆为 Q、K、V 三块。每块独立近似极分解,然后使用每块的 Moonlight 系数 \(0.2\sqrt{1152}\) 缩放,拼回 [3456, 1152] 后更新同一份融合权重。

未分块时对应系数是 \(0.2\sqrt{3456}\)。分块后虽然单块系数变小,却有 3 块一起贡献范数,所以默认整体范数不变;若开启 rescale,则每块系数额外乘 \(\sqrt{3}\)。这个例子只展开原文的张量形状与缩放规则,不代表新增的数值实验,也不要求把前向 QKV 线性层拆成三个模块。

损失函数 / 训练策略

训练仍使用 Flow Matching:数据端位于时间 0,噪声端位于时间 1,在两者间做线性插值,并预测从数据到噪声的速度。缓存公式排版存在损坏,下面依据第 3.1 节清晰的文字定义转写,不补充原文未给定的时间采样分布:

\[ x_t=(1-t)x+tz,\qquad \mathcal L(\theta)=\mathbb E_{x,z,c,t} \left[\|F_\theta(x_t,t,c)-(z-x)\|_2^2\right]. \]

推理从时间 1 的噪声向时间 0 积分,本文类条件实验中的 \(c\) 对应类别条件。优化器变化不修改这一生成目标,也不引入 REPA 等辅助表示对齐技术。

实验采用 batch size 1024、bf16、恒定学习率、梯度裁剪最大范数 1.0,以及衰减为 0.9999 的 EMA 模型评估。AdamW 设置为 \((\beta_1,\beta_2)=(0.9,0.95)\)、weight decay 0;仅注意力、FFN、AdaLN 的二维权重交给 Muon/CMuon,一维参数、嵌入和最终层仍由 AdamW 更新。DiT-XL 学习率分析测试了 \(\{1,2,3\}\times10^{-4}\),默认附近的 \(2\times10^{-4}\) 表现较好;缓存没有提供完整附录级复现配置。

实验关键数据

主实验

下表选取原文 Table 2,任务为 ImageNet-1K 256×256 类条件生成,统一使用 30 NFE、EMA 和 FID-50K;FID 越低越好。NFE 是采样时的函数评估次数,不是训练步数;“未报告”不能视为实验失败。

模型 / VAE 优化器 FID@80ep FID@200ep FID@400ep
DiT-B 130M / VA-VAE AdamW 5.87 3.49 2.78
DiT-B 130M / VA-VAE Muon 5.50 3.31 2.78
DiT-B 130M / VA-VAE CMuon 5.14 3.03 2.57
DiT-XL 675M / VA-VAE AdamW 1.66 1.30 1.21
DiT-XL 675M / VA-VAE Muon 1.65 1.29 未报告
DiT-XL 675M / VA-VAE CMuon 1.46 1.18 未报告

在 XL 的同一 200-epoch 预算下,CMuon 比 Muon 低 0.11 FID,比 AdamW 低 0.12。其 200 epochs 的 1.18 还优于 AdamW 400 epochs 的 1.21,因此支持“以一半训练轮数达到更好质量”;正文“超过 2 倍加速”的表述不能直接替换成已经实测的墙钟时间或 GPU 小时节省。

消融实验

下表来自 Table 3,固定 130M DiT-B,只改变哪些投影启用分块;None 就是普通 Muon。

启用分块的位置 FID@80ep FID@200ep 相比 None 的后期表现
None 5.50 3.31 基线
FFN 5.43 3.28 降低 0.03
QKV 5.32 3.35 反而升高 0.04
AdaLN 6.00 3.23 降低 0.08,但早期更差
FFN + QKV + AdaLN 5.14 3.02 降低 0.29

Table 3 的完整配置为 3.02,Table 2 同类配置为 3.03,缓存未解释差别,此处按各表原值保留。更重要的是不能照搬正文“任意单块都改善后期”的概括:QKV 单独分块的 3.35 确实差于 3.31,证据支持联合配置,而不是每块都有单调收益。

关键发现

  • Table 5 将分块与 rescale 分开控制:Muon 的 FID@40ep/@80ep 为 5.67/1.65,Muon + rescale 为 4.26/1.55,CMuon 为 5.32/1.50,CMuon + rescale 为 3.78/1.46。这支持两者互补,但 80 epochs 本身不足以证明任意长训练下的最终收益。
  • Table 4 中 Vanilla、KellerJordan、Moonlight 缩放的 Muon FID 分别为 8.73、7.40、3.31,CMuon 为 6.94、6.00、3.03。作者沿用 AdamW 学习率且未为前两种规则重调,因此不能把这些数字理解为缩放策略的充分调参排名。
  • Table 6 中 AdamW 在 \(3\times10^{-4}\)、200 epochs 下达到 1.26,而 CMuon 在 \(2\times10^{-4}\)、140 epochs 下为 1.27;这是接近而非超过。该表 AdamW 在 \(2\times10^{-4}\)、200 epochs 为 1.29,与 Table 2 的 1.30 也存在未解释的小差异。
  • SD-VAE 的 DiT-B 对照在 200 epochs 时为 AdamW 3.60、Muon 3.30、CMuon 3.06,说明收益不只出现在默认 VA-VAE,但仍是同一个数据集和分辨率范围内的验证。

亮点与洞察

  • 张量布局并非优化中性的实现细节。 对矩阵级优化器而言,合并投影就会改变预条件器看到的统计集合;应同时设计高效计算布局和有意义的优化分组。
  • 控制更新范数,使机制解释更有辨识度。 默认等范数缩放与独立的 rescale 消融有助于区分方向修正和步长增大,避免把全部加速都归因于“解除耦合”。
  • 看完整质量曲线,而不只看早期 loss。 Muon 在 B 模型 400 epochs 时回到 AdamW 的 2.78,CMuon 则到 2.57,这比只展示早期收敛更贴合生成模型训练的实际目标。

局限与展望

  • 验证范围有限。 当前证据集中于 ImageNet-1K 256×256、130M/675M DiT;没有直接验证大规模文本到图像、视频或其他架构,也没有证明语义分块在所有矩阵优化器中都有效。
  • 效率证据主要是训练进度。 作者称额外开销可忽略,但所给缓存没有完整的墙钟耗时、显存、吞吐或分布式通信表;多次小矩阵运算是否保持真实硬件效率,需要独立测量。
  • 消融与统计不够完备。 单块和全块对照不能分离所有交互效应,缺少两两组合、随机等尺寸分块、种子方差与置信区间。进一步比较语义分块和任意分块,才能更严格地定位收益来源。
  • 材料与报告边界需要保留。 缓存止于参考文献,没有附录 A/B 或 Algorithm 3,且存在 3.02/3.03、1.29/1.30 的跨表差异;不能自行补出 Newton–Schulz 配置、理论保证或“修正后”的唯一数字。

相关工作与启发

  • 与 AdamW 对比: AdamW 使用逐元素自适应统计,CMuon 对选定二维权重做矩阵级动量正交化,同时仍依赖 AdamW 处理其他参数。本文是混合优化方案,不能简单称为完全替代 AdamW。
  • 与 Muon / Moonlight 对比: Muon 提供极因子更新,Moonlight 提供按形状的缩放;CMuon 的主要新增点是正交化前的功能分块及相应尺度处理,不是另创一个全新的正交化算子。
  • 与 REPA / VA-VAE 对比: REPA 从表示对齐改善训练,VA-VAE 改变潜空间质量,而 CMuon 改变优化几何。本文无 REPA 的结果说明分块可独立起效,但与 REPA 叠加是否仍有收益尚未由这些实验建立。

评分

  • 新颖性:4/5。改动简单,但准确指出融合张量边界对矩阵级优化的影响,并给出机制解释。
  • 实验充分度:4/5。覆盖模型规模、VAE、分块位置、缩放和学习率,但缺少跨任务、重复种子与真实硬件成本测量。
  • 写作质量:3/5。问题与算法主线清楚,但跨表数字及单块消融的文字概括存在不一致,缓存公式和附录也不完整。
  • 价值:4/5。对使用融合投影的 DiT 训练具有直接实践意义,迁移前仍需核对参数布局和实测吞吐。