CMuon优化器:用分块动量正交化解决DiT训练收敛慢与不稳定

发布时间:2026/8/30 6:12:51
CMuon优化器:用分块动量正交化解决DiT训练收敛慢与不稳定 先把结论放在前面CMuon 是一个面向 Diffusion TransformerDiT训练的优化器改进方案核心思路是分块动量正交化Chunked Momentum Orthogonalization。它要解决的是 DiT 训练里最常见的两个痛点收敛慢、训练过程不稳定。如果你正在训练扩散模型或者准备从 UNet 架构切到 DiT 架构又不想只依赖把学习率调到很小来换取稳定那 CMuon 值得你花时间看一下。这类优化器改进不像换模型结构那么直观需要同时理解动量的作用、正交化的目的、分块带来的计算与稳定性收益。这篇文章我会按自己实测时会用的顺序来拆先聊 DiT 训练为什么难再解释 CMuon 的核心机制接着给一份可落地的复现思路、参数含义、验证方式、排查链路最后说清楚哪些场景适合用它哪些场景别急着换。1. 先搞清楚它到底在解决什么问题CMuon 不是一个新的模型也不是一个损失函数它属于优化器层面的改进。优化器在训练里的地位往往被低估。很多人遇到 DiT 训练不收敛或者 loss 震荡第一反应是改网络结构、调总步数、换学习率表很少有人去怀疑优化器本身的方向更新方式有问题。1.1 为什么 DiT 训练比普通 Transformer 更敏感DiT 网络和文本 Transformer 有相似之处都是多层 Transformer block。但 DiT 的输入不是离散 token而是连续的时间步嵌入、类别条件、带噪图像块。这个差异非常关键。连续输入导致梯度分布不像文本模型那样集中在 embedding 空间而是散落在整个特征空间里。加上扩散训练要不断采样噪声步模型需要在同一个参数空间里同时拟合不同噪声强度的数据。同一个 batch 里可能有的样本是低噪声、接近原图有的样本是高噪声、接近纯噪声。两者的梯度方向差异很大。这种情况下普通 Adam 虽然能工作但对学习率、warmup、EMA 都非常敏感。稍微把学习率调大一点短时间内 loss 可能降得更快但很快就会遇到 spike调小了收敛又太慢。1.2 从 Muon 到 CMuon 的逻辑变化CMuon 的前身是 Muon 这类“动量正交化”思路。Muon 的基本想法是不要把原始梯度直接用于更新而是先对动量做一次正交化让更新方向更均匀、更去相关。到 CMuon 这里变化在于“分块”。大矩阵不做全局正交化而是按块切分每块独立完成正交化更新。这样做看起来只是工程上的取舍实际上它同时影响计算效率和训练稳定性。我在本地复现时最大的感受是CMuon 的代码改动量并不大但你必须理解分块这个动作为什么能避免很多奇怪问题。不然看到 loss 先降后涨你根本不知道是优化器的问题还是分块大小调得不对。2. 分块动量正交化的核心概念怎么理解才不吃力动量正交化这个词听起来有点重但如果拆开看每一部分都是常见操作。动量维护一个梯度历史滑动平均让更新方向不完全跟随当前梯度减少震荡。这个大家都很熟。正交化把动量矩阵的内部方向做一次投影让不同特征方向的尺度更均衡。最简单的理解是训练过程中有些方向梯度特别大有些方向特别小。如果直接按梯度更新大方向会主导小方向被掩盖。正交化之后各个方向之间的相关性降低更新方向更像“等权融合”后的结果。分块不把整个权重矩阵当做一个大矩阵做正交化而是按行、列或某个固定维度切成多块分别处理。2.1 全局正交化为什么在 DiT 上不够友好全局正交化本身不是新东西。问题是 DiT 的权重矩阵通常很大尤其是 attention 里的 QKV 投影、MLP 层动辄上千维甚至几千维。对一个大矩阵做完整正交化计算成本会明显上升。更麻烦的是全局正交化对奇异值比较小的方向很敏感。一个大矩阵里如果有部分方向能量很弱正交化后这些方向会被放大放大之后如果正好赶上连续几步的梯度噪声就很容易在训练中期出现 loss spike。我遇到过不止一次前面 20k 步都正常某一步开始梯度方向异常紧接着 loss 直接翻倍。分块之后正交化只在局部块内进行。每个块的方向范围变小强方向不会影响弱方向弱方向也不会跨块影响其他区域。整体训练会稳定很多。2.2 分块大小怎么理解分块大小是 CMuon 里最需要手动关心的参数。如果分块太小比如每块只有 16 维正交化的范围不够很多相关方向还是没有被处理效果会退化甚至和普通动量差不多。如果分块太大又回到了全局正交化的老路计算和稳定性问题都会回来。可以先把分块大小设成 128 或 256跑一个小规模实验。观察两个东西一个是 wall-clock 时间另一个是训练后期的 loss 曲线是否平滑。分块大小不是越大越好也不是越小越好而是要和你的矩阵维度、batch size、任务复杂度匹配。注意不要一上来就同时调分块大小和学习率。先固定学习率只调分块大小找到稳定区间后再动学习率表这样才分得清是哪个参数引起的波动。3. 复现前先把环境和输入条件理清楚CMuon 的代码量不大但训练 DiT 的工程成本不小。如果你已经有 DiT 训练代码可以直接把优化器替换成 CMuon。如果还没有建议不要从零写完整扩散训练流程而是先跑通一个小规模的 DiT 训练脚本越小越好。3.1 硬件与软件条件训练 DiT 类模型最低建议准备一块 16GB 显存以上的 GPU。不是 16GB 以下不能跑而是图像类扩散模型默认输入分辨率越高显存消耗越快。如果你只有 8GB 显存必须把 batch size、图像分辨率、patch size 都降下来否则连梯度累积都没法稳定跑。依赖方面核心是 PyTorch。建议先确认 PyTorch 版本、CUDA 版本和你训练框架的匹配情况。不要以为 pip install 成功就结束。我见过太多 case代码能启动但 AMP 在某个算子上报错最后发现是 PyTorch 版本太老。如果你的训练脚本里用了torch.compile或者 Deepspeed/FSDP就要额外注意优化器是否支持。CMuon 这种自定义优化器在单卡上跑通不代表多卡零改动。先把单卡跑通这是最重要的前置条件。3.2 用最小配置验证不要一上来就训练完整 ImageNet 或者大分辨率生成模型。建议先跑一个 64x64 或者 128x128 的 toy 配置batch size 小一点总步数 5k 到 10k观察优化器能否正常启动。每一步耗时是否可接受。动量 buffer 是否正确保存和加载。loss 是否在几十步内开始下降。有没有 NaN 或者梯度溢出。这套最小验证基本一两个小时内就能出结果。如果你在 toy 配置上都频繁出现异常那问题通常不在 CMuon 本身而在你现有训练代码和优化器之间的兼容性。4. 单步更新流程伪代码和参数边界CMuon 的更新流程大致可以分成四步计算梯度、更新动量、分块正交化、按学习率应用到参数。下面这个伪代码是教学简化版重点是把结构展示清楚不是某个具体实现。# CMuon 简化伪代码实际项目中需要结合 AMP 和梯度裁剪 def cmuon_update(param, optimizer_state, grad, lr, momentum0.9, chunk_size128): # 1. 更新动量 if momentum not in optimizer_state: optimizer_state[momentum] torch.zeros_like(param) optimizer_state[momentum].mul_(momentum).add_(grad) # 2. 分块正交化 update chunked_orthogonalize(optimizer_state[momentum], chunk_size) # 3. 应用学习率 param.data.add_(update, alpha-lr)这里有几个点需要细说。第一动量系数通常取 0.9 到 0.95 之间。过小分块正交化效果不明显过大更新滞后训练早期收敛会变慢。第二正交化之前要不要做梯度裁剪我的建议是保留原有梯度裁剪逻辑不要因为换了优化器就关掉裁剪。CMuon 可以减轻梯度噪声但不能完全替代梯度裁剪在长训练里的保护作用。第三分块正交化函数要自己实现常见做法是 Newton-Schulz 迭代。过程不复杂但需要注意不要破坏 PyTorch 的 autograd 图。优化器更新时已经不需要梯度了正交化过程应该发生在 detach 后的张量上。4.1 分块正交化的一种实现思路对每个块做正交化最简单稳定的办法是迭代法。假设块矩阵为 X目标是让 X 的方向近似正交化可以反复执行Y X Y 0.5 * Y * (3I - Y^T Y)迭代次数越多结果越接近正交。实际使用时迭代 2 到 5 次通常就够。迭代次数太多计算成本上升收益不明显迭代次数太少正交化效果变弱稳定性提升有限。分块时要注意 reshape 方式。二维权重可以直接按行或列切块高维参数需要先 reshape 成二维再按固定 chunk_size 切。不同层的参数大小不同分块逻辑必须兼容所有参数不能假设每层都能被 chunk_size 整除。常见做法是在分块时丢弃尾部或者对最后一块单独处理。4.2 关键参数的建议学习率CMuon 的学习率不宜直接照抄 Adam 的常见值。如果你之前用 AdamW 时学习率是 1e-4换成 CMuon 后建议先按 1e-3 到 5e-3 量级试但必须用小步数实验确认。不要一上来就用 1e-2那需要非常规范的训练流程和足够小的梯度噪声。weight decay很多 Muon 类优化器实现会选择把 weight decay 和动量正交化分开。建议先用一个很小的 weight decay比如 0.01 或 0.1甚至先不开启等训练曲线稳定后再加。这样可以减少变量。梯度裁剪max grad norm 从 1.0 开始如果 loss 依然不稳定降到 0.5。不要一开始就极致压缩到 0.1那样会严重拖慢收敛。注意这里给的是通用排查顺序。不同基线、不同数据、不同模型规模下最佳参数一定不同。以你的实际效果为准。5. 怎么验证“加速”和“稳定”这两个核心收益很多优化器类项目容易让人陷入“看起来很有道理跑一跑发现没什么用”的尴尬。原因不是原理有问题而是验证方式太粗糙。只看最终指标、只跑一个配置很难判断优化器到底在哪个环节带来了收益。5.1 加速效果怎么量化加速最直接的标准不是“每秒 step 数涨了多少”而是“达到相同质量指标需要的时间和步数”。建议固定模型结构、数据、总 batch size、学习率表和随机种子在一组配置里同时跑 AdamW 和 CMuon记录两个东西相同训练步数下的 loss 值。达到目标 loss 或目标 FID 时两者分别用了多少步、多少分钟。wall-clock 时间也很重要。如果 CMuon 每一步因为分块正交化多了不少计算即使步数减少总时间可能也没优势。实际测下来分块正交化的额外开销通常远小于整次 forward-backward 的开销但“通常”不等于“一定”。在小模型、小 batch 场景下额外开销占比可能很高这时你会看到 CMuon 收敛步数更少但每秒步数更慢。所以验证加速必须同时看步数和墙钟时间。5.2 稳定性怎么判断稳定性不能只看最终 loss。建议额外记录这几个指标loss 是否出现突然 spike。是否有任何一步梯度或者输出变成 NaN。loss 曲线中后部是否明显比 AdamW 平滑。使用相同 EMA 参数时EMA 版本是否比原始 checkpoint 更稳定。我一般会把“训练中 loss 超过前后若干步均值 30% 的事件次数”统计出来。CMuon 如果有效这个次数应该明显少于 AdamW。这里 30% 只是一个参考阈值你可以按自己的任务调。如果只有最终 FID 提升但训练曲线频繁剧烈波动那只能说明它在这个配置下碰巧收敛到了好位置稳定性的结论还不够可靠。再多换几个 seed 跑确认不是随机运气。5.3 建议对比表对比维度AdamWCMuon学习率敏感度较高需要精细 warmup相对更低但仍要调loss 平滑度依赖具体 lr通常更平滑达到目标 loss 的步数基线需要实测每步额外开销低分块正交化有少量开销对大矩阵的适应性好需要调 chunk_size这张表不是结论是一个记录框架。数字要由你的实验填写。6. 常见失败场景按这个顺序排查换优化器之后出问题的概率并不低。如果你在 DiT 训练里用 CMuon 遇到异常不要直接回退到 AdamW先按下面这个顺序排查。6.1 训练发散或 NaN第一步看日志里 NaN 首次出现的位置。是在前向里、反向里还是更新后如果更新后立刻 NaN说明更新量过大或分块正交化步骤数值不稳定。排查顺序关闭 AMP用全精度跑 500 步看是否复现。降低学习率一个数量级看 NaN 是否消失。把 chunk_size 调小一些比如从 256 降到 128看是否改善。检查正交化实现里有没有在非 detach 张量上操作导致梯度图异常。确认输入和 label 没有异常值比如空洞 mask 造成的无穷大。如果关闭 AMP 后正常问题大概率在混合精度和正交化算子的数值稳定性上。6.2 loss 不降或收敛过慢loss 不降时很多人会直接加大学习率这是错误做法。CMuon 的问题往往不在学习率而分块大小和动量系数不匹配。排查顺序确认动量 buffer 是否正确初始化有没有在 load checkpoint 时丢失。检查 weight decay 是否过大压住了更新方向。把 chunk_size 调大一点增强正交化范围。把学习率先固定单独扫描 momentum用 0.85、0.9、0.95 几档对比。如果仍然不降回看数据集和 loss 是否稳定很多时候是数据问题而不是优化器问题。6.3 速度反而变慢如果你发现单个 step 时间明显变长优先看分块正交化的实现复杂度。Newton-Schulz 迭代如果实现不好或者每次都对非常大的矩阵做重复迭代会造成额外开销。可以先测一下正交化函数单独跑一次耗时。如果耗时超过整个 step 的 10%说明实现需要优化。可以从这三个方向优化减少迭代次数。用更小的 chunk_size。只对特定层启用分块正交化比如 attention 里的 QKV 层不对 embedding 和 norm 层用。注意速度慢不一定是 CMuon 的问题。先看有没有显存换入换出、数据加载瓶颈、CPU 预处理抢占 GPU 利用率。这些都会让单步耗时变高。7. 什么场景适合用 CMuon什么场景要谨慎优化器不是越复杂越好。你的任务、模型规模和现有基础设施都会影响 CMuon 的实际收益。7.1 适合的判断清单如果你的情况符合下面几条CMuon 可以优先试模型主体是 Diffusion Transformer或者类似视觉 Transformer 的连续输入回归任务。训练曲线经常出现 loss spike靠降低学习率已经压不住了。你有足够长的训练计划总步数在 100k 以上优化器带来的步数收益会累积。权重矩阵维度较大比如 attention hidden size 在 512 以上。矩阵太小分块正交化的收益不明显。7.2 要谨慎的情况如果你的模型比较小比如 hidden size 只有 128参数量很小CMuon 很可能不会比 AdamW 有明显优势反而增加实现和调试成本。如果你在跑一个已经有成熟超参的 baseline目标是严格对照某个已有实现建议先不要替换优化器。优化器一变lr、warmup、weight decay 都要重新搜索否则结论无法归因。另外如果你需要用到非常规权重访问方式比如 LoRA、部分参数冻结、选择性 offload也要确认 CMuon 的实现是否能和这些机制兼容。很多自定义优化器只支持常规 dense 参数更新遇到 LoRA 的 low-rank 参数可能需要额外适配。7.3 两个隐藏坑第一个坑是 checkpoint 兼容性。CMuon 需要保存动量 buffer 和分块相关信息。如果训练中途要 load 一个 AdamW 的 checkpoint动量 buffer 对不上只能重新开始累积动量。前期可能会有几 k 步的不稳定期。第二个坑是超参搜索成本。CMuon 的学习率、动量、chunk_size、迭代次数形成一个新的参数面比 AdamW 要多调整一个 chunk_size。如果你的实验预算只够跑少量配置可能还没调到合适区间就放弃了。建议先用 1k 步小实验粗筛再用粗筛结果进入正式训练。8. 如果我从头试一次会这样安排如果我现在要在一个新项目里试用 CMuon会按下面这个节奏走而不是直接改完优化器就开始长训练。第一步先拿一个很小的 DiT 训练脚本分辨率降到 64 或 128总步数 2k 到 5k验证优化器代码本身没有 bug。重点看动量 buffer、梯度裁剪、AMP 这几块能不能同时工作。第二步固定 batch size 和模型结构跑一组 AdamW 和 CMuon 的对比。学习率各选 2 到 3 个点总实验控制在 6 到 8 个以内。记录 loss 下降曲线和单步耗时先不看 FID只看训练阶段的差距。第三步如果 CMuon 在训练曲线上明显更平滑或者达到相同 loss 的步数更少再进入正式实验。此时再调 chunk_size 和 momentum找一组在稳定性和收敛速度之间最均衡的参数。第四步长训练时把 EMA、checkpoint 保存、日志记录都配好。分块正交化相关状态最好在 checkpoint 里单独保存避免和模型参数混在一起。踩过几次坑之后我最大的感受是CMuon 这类优化器确实能解决一部分 DiT 训练问题但它不是免调参的开关。真正决定成败的仍是数据质量、模型实现、梯度流和训练配置之间的匹配程度。先把最小实验跑稳再谈加速和稳定是最不会出错的做法。