资讯动态

CMuon优化器:用分块动量正交化加速Diffusion Transformer训练

发布时间:2026/8/30 8:08:06 来源:尧图企业网站定制
训练 Diffusion TransformerDiT时优化器经常是被忽略的一环。很多人会默认使用 AdamW把精力花在模型结构、采样器和数据集上。可在 DiT 的早期训练阶段损失曲线抖动、梯度范数异常、收敛变慢这些现象往往并不来自模型设计而是来自更新方向本身。本次要看的 CMuon 正是针对这类问题提出的优化器改进通过 Chunked Momentum Orthogonalization把“分块”和“动量正交化”组合起来目标是加速并稳定扩散 Transformer 训练。从标题看CMuon 的核心思路不复杂。它先保留动量法对历史梯度的平滑作用再对动量矩阵做正交化处理而且不是把整个大矩阵一次性分解而是按块处理。这种设计能在引入类似自然梯度的方向校正时避免过高的计算和存储开销。对 DiT 这种包含大量二维参数矩阵的架构来说这类方法比普通 AdamW 更贴近参数几何又比 Shampoo 类二阶优化器更容易落地。这篇文章不是概念介绍也不是把论文标题翻译一遍。我会先把 Chunked Momentum Orthogonalization 拆开讲清楚再给出一份可以落到 PyTorch 训练循环里的示意实现然后告诉你如何设计对比实验、观察显存与梯度波动以及遇到 loss 为 NaN、显存溢出、收敛变慢时怎么排查。适合正在复现 Diffusion Transformer、训练扩散模型或者想改造优化器的开发者。读完后你能快速判断这个优化器在你的任务上值不值得试。1. 核心能力速览能力项说明项目类型Diffusion Transformer 训练优化器 / 训练方法核心机制Chunked Momentum Orthogonalization分块动量正交化主要目标加速 DiT 训练收敛提高训练稳定性适用模型Diffusion TransformerDiT及含大量二维参数矩阵的模型基础优化器类 Muon 思路基于动量正交化非标准 AdamW硬件门槛需要按实际模型规模测试建议使用支持 CUDA 的 GPU存储开销分块正交化可降低全矩阵分解的存储压力具体取决于块大小启动方式作为优化器模块集成到 PyTorch 训练脚本无需独立服务是否支持 API不直接提供 HTTP API可嵌入训练框架使用是否支持批量任务可配合超参扫描脚本或实验管理工具批量运行适合读者复现扩散模型、研究训练稳定性、二次开发优化器的开发者上表信息主要来自标题本身与公开的优化器设计思路。任何显存占用、加速倍率、收敛步数等指标都需要以论文源码和你自己的训练环境实测结果为准。2. 适用场景与使用边界CMuon 最直接的适用场景是训练 Diffusion Transformer。DiT 将图像空间映射到隐空间后用 Transformer 块建模噪声预测过程结构上包含大量线性投影层和 MLP 矩阵。这些参数在训练过程中会经历剧烈变化尤其是时间步嵌入和条件注入部分很容易出现 loss 尖峰。CMuon 的动量正交化思路能够对参数更新方向做归一化理论上可以缓解这类不稳定。从标题中的“Chunked”来看它也很适合参数规模较大的模型。如果直接对整个全尺寸矩阵做正交化每一步都要计算矩阵分解或矩阵平方根逆计算量和内存都不可控。分块之后每个块的尺寸变小正交化成本显著下降。所以它特别适合那种“模型宽度很大、单层参数成矩阵”的训练场景。但也要说清楚使用边界。CMuon 不是万能的。它针对的是带有二维结构参数的模型对一维偏置、LayerNorm 参数、Embedding 表这类向量参数一般还是退化为普通动量更新。如果任务本身是一个小模型、小 batch 的简单分类任务换用 CMuon 可能看不到明显收益反而增加实现复杂度。另外如果训练过程已经使用 AdEMAMix、Schedule-free 这类更新规则再叠加 CMuon 时需要注意优化器状态重复累积的问题。使用边界还包括数据合规与生成合规。扩散模型训练数据必须来自合法授权渠道。如果模型后期用于生成人像、品牌元素或受版权保护的风格需要在发布和商用前确认授权边界不能因为换了优化器就忽略内容合规。这一点在后续最佳实践里还会强调。3. 前置知识Diffusion Transformer 训练难点要理解 CMuon 为什么选择“分块动量正交化”得先回到 DiT 训练的几个痛点。第一个痛点是模型结构导致的大矩阵更新。DiT 的核心模块是自注意力层和前馈网络。注意力层中的 QKV 投影是二维参数矩阵前馈网络的第一个线性层通常会把 hidden size 扩展到 4 倍以上。以 DiT-XL 这类规模为例单个矩阵可能包含数百万个参数。AdamW 会为每个参数保存一阶动量、二阶动量显存占用直接翻倍。更重要的是AdamW 的更新方向是按元素做归一化没有考虑矩阵整体方向的几何关系在某些损失地形陡峭的区域容易产生不稳定的更新。第二个痛点是扩散模型的训练目标会随噪声步变化。扩散模型每个 batch 都会采样不同的时间步对应不同强度的噪声。模型需要同时学会去小噪声和大噪声梯度分布在不同 time step 上差异很大。如果优化器对历史梯度的处理不够平滑就可能出现 loss 突然上升、梯度范数剧烈波动的情况。普通动量可以在一定程度上平滑梯度但不会改变更新方向上的几何分布。第三个痛点是训练刚开始时的冷启动不稳定。DiT 通常采用较大的学习率来加速收敛但大学习率配合 AdamW 容易在初始化阶段产生异常大的更新。Muon 这类优化器通过对动量矩阵做正交化把更新方向限制在更规范的子空间里本质上是对学习率步长做了一层“方向校正”。CMuon 进一步做分块就是希望在保留这种校正能力的同时让每一步的计算更轻。第四个痛点是长训练运行下的漂移问题。扩散模型训练到中后期loss 曲线虽然整体下降但局部会有一些尖峰。这些尖峰往往与某些参数矩阵的特征值分布变化有关。如果只靠 AdamW 的二阶矩估计很难及时捕捉这种分布变化。分块正交化相当于给每个块一个相对独立的坐标系能更灵敏地反映局部参数流形。4. 算法拆解Chunked Momentum Orthogonalization4.1 从 Muon 到分块CMuon 的命名包含三个关键词Chunked、Momentum、Orthogonalization。如果把它拆开看就是先把梯度做动量累积再对动量结果做正交化并且整个过程按块执行。Muon 这一类优化器的基本思路可以理解为“带正交化预处理的动量 SGD”。普通的 SGD 加上动量更新方向是历史梯度的加权和Muon 在得到动量矩阵后不直接拿这个矩阵去更新参数而是先做一个“极分解”或类似操作把矩阵分解成旋转/正交部分和尺度部分。更新时只保留正交部分再乘一个学习率系数。这样做的效果是让更新方向不依赖参数的绝对尺度只依赖方向信息类似一种轻量级的自然梯度。CMuon 的“Chunked”则是在 Muon 的基础上把整张参数矩阵按行或按列切成若干个小块。对每一个小块单独做正交化而不是对整张大矩阵做一次完整正交化。这种做法的动机很直接大矩阵的极分解或矩阵平方根逆计算成本高切块后每个块更小单次计算更快也更容易并行。分块后每个块都能保留局部方向信息代价是失去矩阵块之间的全局相关性。4.2 分块到底切什么分块方式会直接影响效果。常见的设计有两种一种是按矩阵的列方向切块。比如一个输出维度为 1024 的线性层参数是[1024, 4096]可以把 4096 的输入维度切成 4 个大小为[1024, 1024]的块。这样每个块都是方阵或接近方阵正交化更稳定。另一种是按行方向切块把多个输出头的投影分开处理。对于 Multi-Head Attention 的 QKV 投影这种方法比较自然因为每个注意力头本身就可以独立看待。分块大小是一个超参数块越大块内相关性保留越多但计算开销越高块越小计算越快但会损失跨特征维度的方向信息。在实际使用时需要结合模型宽度和显存容量来选。如果论文公开了推荐值优先使用论文配置。4.3 正交化怎么算正交化通常要计算给定矩阵的“正交因子”。最常见的方法包括 QR 分解、极分解和 Newton-Schulz 迭代。QR 分解实现简单torch.linalg.qr可以直接调用但 QR 分解并不是每次都能稳定地给出期望的旋转方向而且对非方阵的处理和显存占用需要额外注意。Newton-Schulz 是近似计算极分解的常用迭代方法。它通过多次迭代逼近矩阵的“正交部分”每一步只涉及矩阵乘法比较适合在 GPU 上跑。迭代次数越多结果越精确但计算量会线性增加。CMuon 如果走的是这条路线那么n_iters就是一个和计算速度直接相关的参数。需要注意正交化不是把矩阵变成单位阵而是保留矩阵的方向特征、消除尺度影响。这个区分很重要。如果把动量矩阵直接替换成单位矩阵那就等于丢掉了所有历史梯度信息正交化的目标是提取“方向”保留“长度”由学习率控制。4.4 与 AdamW 的差异对比维度AdamWCMuon分块动量正交化更新规则按参数元素归一化对动量矩阵做块级正交化方向感知没有矩阵方向感知有块内方向感知显存状态一阶 二阶动量动量 正交化临时缓冲计算成本低比 AdamW 高但低于完整二阶优化器典型问题大学习率下容易尖峰需要调好块大小和迭代次数这种差异在 DiT 训练中会体现得很明显。AdamW 每个参数独立缩放适合训练稳定、梯度尺度均匀的任务DiT 的梯度分布会随 time step 变化块级正交化相当于给更新方向做了一次“局部白化”能减少梯度尺度差异带来的抖动。5. 环境准备与训练实验前置条件5.1 硬件建议CMuon 本身是优化器不限制硬件但你要训练 Diffusion Transformer硬件门槛来自模型本身。建议先准备至少一张支持 CUDA 的 NVIDIA GPU。显存大小取决于你想跑 DiT-S、DiT-B 还是 DiT-L/XL。如果是第一次测试建议先用最小配置的 DiT 或自定义小 Transformer 结构跑通流程再逐步放大。显存占用可以在每次训练步骤里通过nvidia-smi观察。不要只凭经验估计。优化器多出来的显存主要集中在动量缓冲区和正交化中间结果上分块大小越大中间结果越大。5.2 软件依赖需要准备的软件环境包括Python 3.10 或 3.11具体版本以项目依赖为准。PyTorch 2.x建议使用支持 CUDA 的版本。CUDA 和配套驱动版本匹配很关键。可选的accelerate用于分布式训练和混合精度。可选的wandb或tensorboard用于记录 loss 和梯度指标。可选diffusers用于模型结构和数据 pipeline 参考。不需要额外安装独立服务。优化器以 Python 模块形式进入训练脚本。5.3 实验设计准备在跑正式实验前建议先固定以下内容随机种子。数据集切分。batch size。模型结构。训练步数或 epoch 数。这些内容不确定对比实验就没有意义。CMuon 和 AdamW 对比时两者必须使用完全相同的模型、数据顺序和种子只允许优化器超参数不同。6. 集成到训练流程示意实现下面给出一个教学级的 CMuon 思路示意实现用于理解核心逻辑不是论文官方代码。如果你要复现论文必须以开源仓库为准。import torch from torch.optim import Optimizer def orthogonalize_blocks(x, chunk_size128, n_iters5): 对二维张量按块做正交化示意实现。 这里为了方便展示对每个块使用 QR 分解提取正交因子。 实际论文如果使用 Newton-Schulz 迭代请用源码替换。 rows, cols x.shape # 这里沿列方向切块 blocks [] for start in range(0, cols, chunk_size): block x[:, start:start chunk_size] # block: [rows, block_cols] q, _ torch.linalg.qr(block) blocks.append(q) return torch.cat(blocks, dim1) class CMuonSchematic(Optimizer): 分块动量正交化优化器示意实现。 仅用于学习不保证与论文完全一致。 仅对二维参数做分块正交化一维参数保留普通动量更新。 def __init__(self, params, lr1e-3, momentum0.9, chunk_size128): defaults dict(lrlr, momentummomentum, chunk_sizechunk_size) super().__init__(params, defaults) torch.no_grad() def step(self, closureNone): loss None if closure is not None: with torch.enable_grad(): loss closure() for group in self.param_groups: lr group[lr] momentum group[momentum] chunk_size group[chunk_size] for p in group[params]: if p.grad is None: continue grad p.grad state self.state[p] if momentum_buffer not in state: state[momentum_buffer] torch.zeros_like(p) buf state[momentum_buffer] buf.mul_(momentum).add_(grad) if p.dim() 2: update orthogonalize_blocks(buf, chunk_sizechunk_size) else: update buf p.add_(-lr * update) return loss这段代码最重要的信息是分块入口和正交化步骤。orthogonalize_blocks中沿列方向切块把每块送入 QR 分解。在真实实现里你可能需要处理以下问题参数矩阵的行列数不是chunk_size的整数倍需要处理最后一个不完整块。QR 分解对某些块可能返回带符号的Q需要决定是否做符号统一。如果使用 Newton-Schulz 迭代迭代次数要加到状态字典里。梯度裁剪要在优化器 step 之前做。混合精度下正交化最好在 FP32 下完成避免低精度带来的方向偏差。训练循环里的接入方式很简单把原来的AdamW替换成上面的优化器即可from torch.utils.data import DataLoader model YourDiffusionTransformer() optimizer CMuonSchematic( model.parameters(), lr1e-4, momentum0.9, chunk_size512, ) for step, batch in enumerate(train_loader): noise torch.randn_like(batch[latent]) t torch.randint(0, num_timesteps, (batch[latent].size(0),), devicebatch[latent].device) loss model.loss(batch[latent], noise, t) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() if step % 100 0: print(fstep {step} loss {loss.item():.4f})这里没有额外启动 WebUI也没有 HTTP API。CMuon 不是一个服务而是训练环节里的一个组件。如果你的训练流程是离线批量任务可以把多组超参写在配置文件里循环启动。7. 功能测试与效果验证7.1 先做小规模冒烟测试不推荐直接把 CMuon 放到完整 DiT 训练上。第一次运行建议使用小模型、小分辨率和少量训练步数目标是确认优化器能正常更新、loss 能下降、显存不会溢出。冒烟测试配置示例模型一个 4 层、128 hidden size 的小 Diffusion Transformer。输入32x32 的低分辨率图或随机 latent。训练步数200 到 500 步。batch size8 或 16。chunk_size先设为 128。预期结果是 loss 在前 200 步内明显下降而不是直接变成 NaN 或始终不动。如果前 500 步 loss 完全不下降优先检查学习率是否过小、正交化方向是否正确。7.2 设计 AdamW 对比实验CMuon 是否有效要靠对比实验说话。标准做法是同一个模型结构。同一份数据集和 batch 顺序。固定所有随机种子。分别用 AdamW 和 CMuon 训练相同步数。记录训练 loss、验证 loss、FID 或生成样本质量。如果只想验证“稳定性”可以故意使用较大学习率跑短训练。比如 AdamW 在 lr3e-4 时 loss 出现尖峰CMuon 在同样学习率下 loss 曲线更平缓这就能说明正交化带来的稳定性收益。7.3 观察梯度范数除了 loss还要记录梯度范数。因为在 DiT 训练中loss 尖峰往往先表现为梯度范数异常。可以在每个 step 之后记录全参数梯度 L2 范数和每个 Transformer 块内的梯度范数。建议记录如下指标全局梯度 L2 范数。各层参数梯度范数。loss 的滑动平均和原始值。优化器更新步长的最大绝对值。参数梯度的 sign 变化频率。表格里可以用简单规则判断指标观察方式判断标准训练 lossTensorBoard / wandb整体下降无明显尖峰全局梯度范数每步记录尺度稳定不出现数量级跳变更新步长统计最大值不应出现单次更新把参数改变过多的现象有效 batch size固定保持一致7.4 判断成功与否不要只看某一个 step 的 loss。建议设置“有效步数比例”比如 1000 步里出现 loss 超过滑动平均 3 倍的步数占比。如果 AdamW 占比 5%CMuon 占比 1%可以认为稳定性有所提升。如果你主要目标是加速收敛那就关注达到同一验证指标所需的步数。比如 FID 从 30 降到 20AdamW 需要 100k 步CMuon 需要 80k 步这才叫加速。只说“loss 降得更快”还不够必须结合生成指标判断。8. 性能观察与资源占用优化器的性能观察重点有三个显存、计算时间、内存开销。显存方面建议用循环命令观察训练进程的显存占用# 每2秒采集一次显存状态 nvidia-smi --query-gpuindex,memory.used,utilization.gpu --formatcsv -l 2然后把同一时刻的显存占用与 batch size、分块大小对应起来。你可能会发现chunk_size越大正交化中间张量越大显存占用越高。如果显存不够优先减小chunk_size而不是一味降低 batch size因为小 batch 会影响梯度稳定性。计算时间方面要对比每个 step 的平均耗时。可以这样记录# 计时训练脚本 time python train.py --config config_cmuon.yaml但更准确的方式是在代码里记录每 100 个 step 的耗时import time start time.time() for step, batch in enumerate(train_loader): ... if step % 100 0: elapsed time.time() - start print(fstep {step}: {elapsed:.2f}s per 100 steps) start time.time()注意正交化的计算量会随块数变化。块越小遍历块带来的 Python 循环开销越大。在示意代码里块循环是 Pythonfor如果块数量很多会产生明显开销。真实实现里应该把所有块操作写成矩阵运算或使用torch.chunk加torch.linalg.qr的向量化形式。显存优化手段包括使用混合精度训练FP16 / BF16正交化过程回退到 FP32。开启 gradient checkpointing 减少激活显存。使用更小的chunk_size。如果使用了分布式训练需要检查优化器状态是否按分片保存。9. 常见问题与排查方法问题现象可能原因排查方式解决方案loss 直接变成 NaN学习率过大、正交化实现错误、混合精度损坏查看第一个出现 NaN 的 step 的梯度范数回退到 FP32降低学习率检查正交化函数是否输出非有限值loss 持续不下降学习率过低、更新方向被正交化破坏对比移除正交化后的 SGD 结果调大学习率或检查分块方式是否切错了维度显存溢出chunk_size 过大、batch size 过大观察 nvidia-smi 数据降低 chunk_size开启梯度检查点训练速度比 AdamW 慢很多块循环未向量化、正交化迭代次数过多统计每 100 步耗时减少迭代次数用批量矩阵运算替代循环分布式训练不收敛优化器状态未正确同步或分片检查 grad sync使用 PyTorch 官方 Optimizer 封装loss 出现周期性尖峰时间步采样分布导致梯度方差变化按 time step 分组记录 loss结合 min-SNR 加权或调整采样策略如果遇到 NaN可以先做最小化复现把模型改成单层线性层用同样的优化器训练随机数据。如果最小的线性层也出现 NaN问题就在优化器实现如果小模型正常、大模型异常问题更可能在数值精度或学习率上。也可以定期检查正交化输出def check_finite(tensor, name): if not torch.isfinite(tensor).all(): raise RuntimeError(f{name} contains NaN or Inf: {name})把这个检查插入到动量更新后、正交化后、参数更新后能快速定位问题出在哪一步。10. 最佳实践与使用建议第一次使用 CMuon 时不要把全部超参一次性调到位。建议按以下顺序验证先用最小的 DiT 结构跑通冒烟测试。固定学习率为 AdamW 的一半左右观察 loss 是否稳定下降。固定chunk_size为 256跑 1000 步对比。对比稳定后再扫描学习率和chunk_size的组合。学习率的选择需要单独说明。因为正交化会改变更新方向的实际尺度CMuon 的合适学习率不一定等于 AdamW 的学习率。最好每次只改一个变量并记录 loss 曲线和生成指标。分块大小建议以 128、256、512 为初始候选。如果你的模型层宽度是 1024分块大小 256 意味着每个块是 1024x256 或 256x256 的尺寸计算压力较小。如果宽度超大可以进一步缩小。在工程上还建议做几件事把所有超参写入 YAML 配置文件包括学习率、动量、块大小、迭代次数、seed、数据集路径。训练日志里记录优化器名称和分块配置方便后续复盘。批量扫描时使用同一套数据顺序避免数据加载顺序影响对比。多卡训练时先确认优化器状态能与 FSDP 或 DeepSpeed 兼容。定期保存 checkpoint至少保留最近两个防止实验中途崩掉后丢失全部结果。关于合规使用这里要单独强调。扩散模型的训练数据必须来自合法授权渠道生成内容如果涉及真人肖像、品牌标识或受版权保护的素材需要确认授权。优化器只是训练工具不能因为技术进步就忽略数据来源和生成内容的合法性。11. 总结与下一步CMuon 最值得尝试的点是它给 Diffusion Transformer 训练提供了一个比 AdamW 更“有方向感”的更新规则同时用分块把计算成本控制在可接受范围内。对你来说最先应该验证的是同样的学习率下loss 曲线是否更平稳、梯度范数是否更少出现尖峰。最容易踩的坑有两个一个是正交化实现写错比如把 QR 分解当成直接归一化导致更新信息丢失另一个是分块大小设置过大显存或计算时间直接爆掉。建议先用小模型把这两点跑明白再放大规模。后续可以继续扩展的方向包括把 CMuon 与时间步加权策略结合观察不同 noise level 下的梯度分布把它与学习率调度器配合测试余弦退火和 warmup 的搭配或者把它搬到视频生成、多模态 DiT 模型上看这类架构能否获得同样的稳定性和收敛收益。这篇文章的目的是帮你建立一个可执行的验证框架。到这里你可以直接去跑一个小实验把 CMuon 和 AdamW 放在同一份数据上对比 500 步。如果 loss 曲线、梯度范数和生成结果都有可复现的差异那这个优化器就值得继续深入。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价