拓冰建站拓冰建站
首页 / 资讯中心 / 正文

CMuon优化器:动量正交化加速DiT训练的新方案

训练 Diffusion TransformerDiT时很多人会默认选择 AdamW 作为优化器毕竟它在分类、分割、生成模型里都有不错的表现。但真正把 DiT 从小规模跑到大规模之后会发现AdamW 收敛慢、loss 波动大对 batch size 和学习率都非常敏感。最近有一类基于“动量正交化”的优化器思路在语言模型和生成模型训练里都表现出很强的竞争力Muon 是其中的代表而 CMuon 进一步把正交化过程改成“分块”执行更适合 DiT 这类充满大型二维权重矩阵的模型。本文会围绕 CMuon 展开先讲清楚它背后的动量正交化原理再给出 PyTorch 实现并接一个简易 DiT 训练对比流程。适合有 PyTorch 基础、正在做扩散模型或 Transformer 生成模型训练的读者。学完以后你能理解 Muon/CMuon 为什么有效也能直接把它接入自己的训练循环。1. 为什么 Diffusion Transformer 训练又慢又不稳1.1 DiT 的基本结构Diffusion Transformer 并不是某种全新的网络范式而是把扩散模型中的去噪主干从 U-Net 换成了 Transformer。它的输入不是一张完整图像而是被切成的 patch 序列。一个典型 DiT 结构包含以下几部分Patch Embedding把图像切成固定大小的 patch然后通过卷积或线性层映射成 token。Timestep Embedding把扩散时间步 t 编码成向量注入到每个 Transformer Block。Class Embedding如果是条件生成还需要把类别标签映射成向量。Transformer Block包含 LayerNorm、多头自注意力、MLP、残差连接。输出层把 token 映射回 patch 级别的噪声预测。所以 DiT 本质上就是一个标准的视觉 Transformer只是输入输出都围绕扩散模型的噪声预测任务设计。这种结构的好处是模型表达能力更强对图像全局信息建模更充分但代价是参数量大、训练计算量大优化难度也更高。1.2 训练痛点我在实际训练 DiT 时最明显的四个痛点是收敛慢。扩散模型需要大量迭代才能学到有意义的噪声预测而 Transformer 参数规模又大AdamW 在部分任务上往往要跑很多步才能看到明显的 loss 下降。稳定性差。训练初期 loss 很容易剧烈震荡尤其是 batch size 较小、学习率偏高的时候甚至会出现长时间不下降的“平台期”。优化器内存开销大。AdamW 需要保存一阶动量和二阶动量这两个缓存和模型参数同规模。DiT 参数量一旦上来显存和内存压力非常明显。对超参数敏感。beta、epsilon、learning rate、weight decay 任何一个设置不合适都会导致收敛曲线差异巨大。这些痛点不是调一两个参数就能彻底解决的它们和优化器的更新规则本身强相关。1.3 优化器在训练中的角色优化器决定了每一步参数更新的方向和幅度。AdamW 的做法是对每个参数维护一阶矩估计和二阶矩估计然后做逐元素归一化。这个思路在非凸优化里非常稳健但也有副作用逐元素缩放忽略了参数之间的结构关系。二阶矩估计会让学习率被动态改变对噪声比较敏感。对于矩阵形式参数AdamW 并没有利用矩阵内部的几何结构。而动量正交化思路的出发点很简单既然 DiT 里大量参数是二维矩阵为什么不直接对整矩阵做“正交化”的更新这样做既能保留动量信息又能让更新方向更接近单位正交方向理论上可以加速收敛并提高稳定性。2. 核心概念动量正交化与 CMuon2.1 Momentum动量回顾动量方法是梯度下降的经典改进它把历史梯度累积起来减少更新方向的震荡。公式可以写成m_t μ * m_{t-1} g_t其中 μ 是动量系数g_t 是当前梯度。更新时用 m_t 而不是 g_t 作为方向。在 Muon/CMuon 这类方法中动量仍然存在只是它不再像 AdamW 那样被二阶矩归一化而是被正交化。所以我们可以把 CMuon 理解成“动量 正交化”的组合。2.2 为什么需要正交化先看一个事实在训练 Transformer 时权重矩阵的更新如果保持正交性通常有利于梯度传播也能避免权重值无限膨胀或退化。矩阵的正交化本质上是把一个矩阵投影到正交矩阵集合上。最朴素的做法是做 SVDX U S V^T 正交化结果 U V^T但 SVD 在每次迭代、每个大矩阵上做一次计算量非常大。于是 Newton-Schulz 迭代被引入它用多项式迭代逼近 SVD 的极分解结果每次迭代只涉及矩阵乘法和加减法计算效率远高于 SVD。2.3 Newton-Schulz 迭代Newton-Schulz 迭代的直觉是给定矩阵 X我们希望找到一个接近 X 的正交矩阵 P。通过反复执行如下形式的迭代A X X^T X a X b A X c A^2 X经过足够多次迭代后X 会逐渐收敛到正交矩阵。系数 a、b、c 取决于迭代次数常见 5 次迭代使用的系数为a 3.4445 b -4.7750 c 2.0315这些系数是通过最小化迭代误差得到的。每次迭代主要开销是两次矩阵乘法相比 SVD 要便宜很多。2.4 Muon 优化器的更新规则Muon 优化器对二维及以上的参数矩阵执行以下更新计算动量 m_t。对 m_t 做 Newton-Schulz 正交化。将正交化结果乘以学习率更新参数。对一维参数如 bias、LayerNorm 的 gamma/betaMuon 通常退化为类似 SGD 或 AdamW 的更新方式。Muon 的核心贡献是让参数更新方向保持正交性。这样做的好处是更新步长不再被二阶矩动态缩放训练更平滑。正交更新方向可以缓解梯度消失/爆炸。减少了对 beta、epsilon 等超参数的敏感度。2.5 CMuonChunked Momentum OrthogonalizationCMuon 的关键变化在“Chunked”一词上。它不再是直接对整个动量矩阵做正交化而是把矩阵拆成多个块对每个块分别做正交化最后拼接回去。为什么要分块降低计算开销。一个大矩阵的 Newton-Schulz 迭代涉及 m×m 和 m×n 的矩阵乘法。拆成块后每个块的计算量更小整体耗时更短。降低内存峰值。正交化过程中需要保存中间矩阵分块后同一时间只处理一个块内存占用更低。提升稳定性。过大的矩阵在做极分解时数值误差容易累积。分块相当于给每个子空间一个独立的规范化约束减轻极端数值波动。这里有一个重要的性质需要说明分块正交化得到的整体矩阵并不再是严格意义上的正交矩阵而是“块内正交、块间不保证正交”的近似。但实验表明这种近似在实际训练中不仅没有明显副作用反而在计算效率和稳定性上都有收益。这也是 Chunked Momentum Orthogonalization 想表达的核心用可接受的近似换取更快的训练和更稳定的收敛。3. 环境准备与项目结构3.1 运行环境本文的代码以 PyTorch 为基础不依赖额外的高阶库。你只需要Python 3.9 或更高版本PyTorch 2.0 或更高版本CUDA 可用设备可选CPU 也能跑示例重点不是某个具体版本而是把优化器实现和训练流程跑通。版本需要根据你的项目实际情况调整本文示例以常见环境为例重点演示配置思路。3.2 项目结构建议按下面的目录组织代码cmuon-demo/ ├── optimizer.py # Newton-Schulz 与 CMuon 实现 ├── dit_model.py # 简易 DiT 模型 ├── diffusion.py # 扩散过程与训练循环 └── main.py # 训练入口如果你只需要在自己的项目里接入 CMuon把optimizer.py复制过去即可。4. 手写 CMuon 优化器4.1 实现 Newton-Schulz 正交化先写最基础的正交化函数。它接收一个二维矩阵执行 5 次 Newton-Schulz 迭代后返回近似正交矩阵。# 文件路径optimizer.py import torch def zeropower_newtonschulz(G, steps5, eps1e-7): 使用 Newton-Schulz 迭代将矩阵 G 正交化。 G: [m, n] 的二维张量 steps: 迭代次数默认 5 m, n G.shape # 保证计算过程中行数不小于列数数值更稳定 transposed m n if transposed: G G.T # 归一化防止数值溢出 G G / (G.norm() eps) # 5 次迭代对应的经验系数 a, b, c (3.4445, -4.7750, 2.0315) X G for _ in range(steps): A X X.T B b * A c * (A A) X a * X B X if transposed: X X.T return X这段代码的关键点先做转置保证X X.T的维度更接近正方形有利于迭代收敛。归一化是一个非常重要的细节。如果不做归一化矩阵范数可能在迭代中膨胀导致数值不稳定。5 次迭代是经验值增加次数会更接近严格正交但计算量会线性增长。可以通过一个小测试验证正交化效果torch.manual_seed(42) G torch.randn(64, 32) X zeropower_newtonschulz(G, steps5) Gram X.T X error (Gram - torch.eye(32)).abs().max().item() print(f正交化误差最大偏差: {error:.6f})输出结果通常可以做到 1e-3 量级。如果你把 steps 提高到 10误差会更小但训练速度会下降。4.2 实现 Muon 优化器有了正交化函数Muon 的核心更新逻辑就很简单了。下面我们基于torch.optim.Optimizer实现一个适合训练矩阵参数的 Muon。# 文件路径optimizer.py追加 class Muon(torch.optim.Optimizer): Muon Optimizer 对二维及以上参数使用动量 正交化更新 对一维参数使用 SGD 风格更新。 def __init__(self, params, lr0.02, momentum0.95, ns_steps5, weight_decay0.0): defaults dict(lrlr, momentummomentum, ns_stepsns_steps, weight_decayweight_decay) super().__init__(params, defaults) def step(self, closureNone): loss None if closure is not None: loss closure() for group in self.param_groups: lr group[lr] momentum group[momentum] ns_steps group[ns_steps] weight_decay group[weight_decay] for p in group[params]: if p.grad is None: continue grad p.grad.data state self.state[p] if momentum_buffer not in state: # 动量的初始化保持全零 # 这样更新初期正交化作用于真实梯度方向训练更稳定。 state[momentum_buffer] torch.zeros_like(p.data) buf state[momentum_buffer] # 更新动量 buf.mul_(momentum).add_(grad) update buf if p.ndim 2: # 矩阵参数正交化 update zeropower_newtonschulz(update, stepsns_steps) # 一维参数直接使用动量可自行替换为 AdamW 风格 if weight_decay 0: p.data.mul_(1.0 - lr * weight_decay) p.data.add_(-lr * update) return loss这里值得注意的点是动量缓冲区初始化为零。这个设计并非随意零初始化的动量配合正交化可以让模型在初始阶段严格沿梯度方向更新避免早期更新方向被历史动量污染。4.3 实现 CMuon分块正交化CMuon 的核心改动在zeropower调用之前先把动量矩阵沿列方向切成若干块对每个块分别正交化再拼接。# 文件路径optimizer.py追加 def chunked_momentum_orthogonalization(G, steps5, chunk_size128, eps1e-7): 分块动量正交化 将矩阵 G 沿列方向拆成多个 chunk 每个 chunk 独立做 Newton-Schulz 正交化最后拼接。 m, n G.shape if n chunk_size: return zeropower_newtonschulz(G, stepssteps, epseps) chunks [] for start in range(0, n, chunk_size): end min(start chunk_size, n) chunk G[:, start:end] chunks.append(zeropower_newtonschulz(chunk, stepssteps, epseps)) return torch.cat(chunks, dim1) class CMuon(torch.optim.Optimizer): CMuon: Chunked Momentum Orthogonalization Optimizer 相比 Muon分块执行正交化降低计算和内存开销。 def __init__(self, params, lr0.02, momentum0.95, ns_steps5, chunk_size128, weight_decay0.0): defaults dict(lrlr, momentummomentum, ns_stepsns_steps, chunk_sizechunk_size, weight_decayweight_decay) super().__init__(params, defaults) def step(self, closureNone): loss None if closure is not None: loss closure() for group in self.param_groups: lr group[lr] momentum group[momentum] ns_steps group[ns_steps] chunk_size group[chunk_size] weight_decay group[weight_decay] for p in group[params]: if p.grad is None: continue grad p.grad.data state self.state[p] if momentum_buffer not in state: state[momentum_buffer] torch.zeros_like(p.data) buf state[momentum_buffer] buf.mul_(momentum).add_(grad) update buf if p.ndim 2: update chunked_momentum_orthogonalization( update, stepsns_steps, chunk_sizechunk_size ) if weight_decay 0: p.data.mul_(1.0 - lr * weight_decay) p.data.add_(-lr * update) return loss分块大小的选择需要权衡chunk_size 太大失去分块优势相当于退化成 Muon。chunk_size 太小每个块的行列比例失衡正交化误差增大更新方向容易被破坏。我建议从128或256开始尝试然后根据模型维度调整。4.4 验证分块正交化的效果我们可以做一个快速验证对比全量正交化和分块正交化的结果差异。torch.manual_seed(123) G torch.randn(256, 1024) # 全量正交化 full_out zeropower_newtonschulz(G, steps5) # 分块正交化 chunked_out chunked_momentum_orthogonalization(G, steps5, chunk_size128) print(full shape:, full_out.shape) print(chunked shape:, chunked_out.shape) # 检查分块结果的列正交性 # 块内正交、块间不强制正交 err (chunked_out.T chunked_out - torch.eye(1024)).abs().max().item() print(f分块后整体正交误差: {err:.4f})运行这段代码会看到分块结果的整体列正交误差比全量结果大但块内误差可控。这正是分块正交化的特点用块间正交性换取计算效率。5. 简易 DiT 训练实战下面我们把 CMuon 接入一个非常简化的 DiT 训练流程。这里不会追求模型性能重点是让你看到优化器替换的完整路径。5.1 构建简易 DiT 模型# 文件路径dit_model.py import math import torch import torch.nn as nn def timestep_embedding(t, dim, max_period10000): 时间步编码 t: [B] dim: 编码维度 half dim // 2 freqs torch.exp( -math.log(max_period) * torch.arange(half, devicet.device) / half ) args t[:, None].float() * freqs[None, :] return torch.cat([torch.cos(args), torch.sin(args)], dim-1) class PatchEmbed(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, embed_dim256): super().__init__() self.patch_size patch_size self.n_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # [B, D, H/p, W/p] x x.flatten(2).transpose(1, 2) # [B, N, D] return x class DiTBlock(nn.Module): def __init__(self, embed_dim256, n_heads8, mlp_ratio4.0): super().__init__() hidden_dim int(embed_dim * mlp_ratio) self.norm1 nn.LayerNorm(embed_dim) self.attn nn.MultiheadAttention(embed_dim, n_heads, batch_firstTrue) self.norm2 nn.LayerNorm(embed_dim) self.mlp nn.Sequential( nn.Linear(embed_dim, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, embed_dim), ) def forward(self, x): x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] x x self.mlp(self.norm2(x)) return x class SimpleDiT(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, embed_dim256, depth4, n_heads8, num_classes10): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) self.t_embed nn.Sequential( nn.Linear(embed_dim // 4, embed_dim), nn.SiLU(), nn.Linear(embed_dim, embed_dim), ) self.class_embed nn.Embedding(num_classes, embed_dim) self.pos_embed nn.Parameter(torch.zeros(1, self.patch_embed.n_patches, embed_dim)) self.blocks nn.ModuleList([ DiTBlock(embed_dim, n_heads) for _ in range(depth) ]) self.norm_out nn.LayerNorm(embed_dim) patch_area patch_size * patch_size * in_chans self.head nn.Linear(embed_dim, patch_area) def forward(self, x, t, yNone): x self.patch_embed(x) # [B, N, D] x x self.pos_embed t_emb self.t_embed(timestep_embedding(t, self.t_embed[0].in_features)) t_emb t_emb[:, None, :] if y is not None: c_emb self.class_embed(y)[:, None, :] t_emb t_emb c_emb x x t_emb for block in self.blocks: x block(x) x self.norm_out(x) x self.head(x) # 还原成图像形状 [B, C, H, W] B, N, D x.shape C 3 p int(math.sqrt(D // C)) h w int(math.sqrt(N)) * p x x.view(B, int(math.sqrt(N)), int(math.sqrt(N)), C, p, p) x x.permute(0, 3, 1, 4, 2, 5).reshape(B, C, h, w) return x5.2 扩散训练循环我们使用一个简化的 DDPM 风格训练流程直接对噪声做 MSE 回归。# 文件路径diffusion.py import torch import torch.nn.functional as F class SimpleNoiseSchedule: 预计算 DDPM 的 sqrt_alpha_cumprod 和 sqrt_one_minus_alpha_cumprod def __init__(self, timesteps1000, devicecpu): self.timesteps timesteps self.device device beta torch.linspace(1e-4, 0.02, timesteps, devicedevice) alpha 1.0 - beta alpha_cumprod torch.cumprod(alpha, dim0) self.sqrt_alpha_cumprod torch.sqrt(alpha_cumprod) self.sqrt_one_minus_alpha_cumprod torch.sqrt(1.0 - alpha_cumprod) def sample(self, x0, t): 根据时间步 t 生成带噪图像 x_noisy noise torch.randn_like(x0) sqrt_ac self.sqrt_alpha_cumprod[t][:, None, None, None] sqrt_one_ac self.sqrt_one_minus_alpha_cumprod[t][:, None, None, None] x_noisy sqrt_ac * x0 sqrt_one_ac * noise return x_noisy, noise def train_one_step(model, optimizer, x, y, schedule, device): 执行一个训练步骤返回 loss x x.to(device) y y.to(device) B x.shape[0] t torch.randint(0, schedule.timesteps, (B,), devicedevice) x_noisy, noise schedule.sample(x, t) pred model(x_noisy, t, y) loss F.mse_loss(pred, noise) optimizer.zero_grad() loss.backward() optimizer.step() return loss.item()扩散过程的核心逻辑是随机采样时间步 t按预计算的噪声系数把纯噪声叠加到输入图像上然后让模型预测这个噪声。模型输出和真实噪声之间的 MSE 就是训练损失。5.3 对比 AdamW 与 CMuon下面写一个对比入口分别在相同数据子集上使用 AdamW 和 CMuon 训练同样的模型。# 文件路径main.py import torch from dit_model import SimpleDiT from diffusion import SimpleNoiseSchedule, train_one_step from optimizer import CMuon def create_train_data(batch_size8, total_batches50, img_size32, devicecpu): 使用随机数据模拟训练集实际项目中替换成真实数据集即可。 torch.manual_seed(0) for _ in range(total_batches): x torch.randn(batch_size, 3, img_size, img_size, devicedevice) y torch.randint(0, 10, (batch_size,), devicedevice) yield x, y def run_experiment(optimizer_name, devicecpu): torch.manual_seed(42) model SimpleDiT(img_size32, patch_size4, in_chans3, embed_dim128, depth2, n_heads4).to(device) if optimizer_name adamw: optimizer torch.optim.AdamW(model.parameters(), lr1e-3) elif optimizer_name cmuon: optimizer CMuon(model.parameters(), lr0.02, momentum0.95, ns_steps5, chunk_size128) else: raise ValueError(f未知优化器: {optimizer_name}) schedule SimpleNoiseSchedule(timesteps1000, devicedevice) losses [] for step, (x, y) in enumerate(create_train_data(devicedevice)): loss train_one_step(model, optimizer, x, y, schedule, device) losses.append(loss) if step % 10 0: print(f[{optimizer_name}] step {step:3d}, loss {loss:.4f}) return losses if __name__ __main__: device cuda if torch.cuda.is_available() else cpu print(使用设备:, device) adamw_losses run_experiment(adamw, device) cmuon_losses run_experiment(cmuon, device)这里有一个关键点AdamW 和 CMuon 的学习率差异很大。AdamW 通常使用 1e-3 到 1e-4而 Muon/CMuon 推荐从 0.02 左右开始。原因是正交化后的更新方向已经被归一化单位步长更大所以需要较小的系数来保证稳定。5.4 结果说明在不同任务上两者的 loss 曲线会有差异但从这类动量正交化方法在生成模型训练中的常见表现看你通常会观察到CMuon 的 loss 在前几十步下降更快。CMuon 的 loss 曲线更平滑剧烈波动更少。相同步数下CMuon 生成的图像质量往往更接近收敛状态。需要强调的是上面的对比数据是随机模拟的不是一个真实基准。真实项目中请使用自己的数据集做对比并且每个优化器都做 grid search 式的学习率扫描否则对比没有意义。6. 常见问题与排查思路在把 Muon/CMuon 接入项目时下面几个问题出现的概率很高。问题现象常见原因解决思路训练 loss 出现 NaN学习率过大或正交化前矩阵范数异常膨胀降低学习率检查梯度是否包含 NaN给zeropower的归一化加 eps收敛速度明显慢于 AdamW学习率没有重新调整沿用 AdamW 的 lr 太小尝试 lr0.01 到 0.05 区间分块时检查 chunk_size 是否过小分块后 loss 波动变大chunk_size 过小正交化近似误差太大增大 chunk_size例如从 256 或 512 开始参数更新方向异常模型输出退化动量缓冲区初始值不是 0确保 momentum buffer 使用zeros_like初始化训练前期正常后期 loss 反弹weight decay 与 lr 不匹配参数被过度衰减降低 weight_decay或使用更小的 lr显存不足分块没有覆盖所有二维参数某些超大矩阵仍全量正交化确认所有 2D 参数都走 chunked 分支并在日志里打印参数形状排查时建议按“数据 → 模型 → 优化器 → 超参数”的顺序先用 CPU 跑几个 step确认数值没有异常。打印各层梯度范数观察是否有梯度消失或爆炸。关闭 weight decay验证 loss 是否恢复。单独替换优化器不改变其他训练配置做 A/B 测试。7. 最佳实践与工程建议7.1 学习率与超参数调试Muon/CMuon 的默认学习率参考区间是0.01 到 0.05。建议按以下流程调参先固定momentum0.95和ns_steps5。用 lr0.01 跑 100 步观察 loss 是否下降。如果 loss 下降太慢提高 lr 到 0.02、0.05。如果出现 loss 震荡降低 lr或增大 chunk_size。不要用 AdamW 的经验去套 MoM/CMuon二者的学习率量级完全不同。7.2 数值稳定性动量正交化的数值稳定性是整个方法的基础工程上需要注意在zeropower里对输入做归一化加上小 eps。对梯度做 clip建议grad_norm.clip_max_norm(1.0)防止异常梯度进入动量缓冲区。混合精度训练时建议在矩阵运算前把张量转为与模型一致的 dtype避免频繁类型转换。如果模型非常大使用 fsdp 或 deepspeed 时要确认 orthogonalization 计算发生在梯度同步之后。7.3 代码工程化生产项目中不建议直接复制训练循环里的优化器代码而是封装成独立模块用param_groups区分不同层的学习率。为 bias 和 norm 参数单独设置weight_decay0。在日志中记录每个参数的更新范数方便观测正交化是否生效。提供chunk_size的自动估算逻辑根据矩阵列数动态调整例如chunk_size min(128, n)。7.4 什么时候该用 CMuonCMuon 并不适合所有任务。我的经验是适合DiT、GAN 生成器、大规模 Transformer、需要长 step 训练的生成模型。可能不适合小规模模型、欠拟合场景、训练步数非常少的情况。此时正交化带来的额外计算可能超过收益。建议在正式切换前先用小规模模型做 500 到 1000 步的快速对比实验比较 AdamW 和 CMuon 的 loss 下降曲线再决定是否全面迁移。7.5 安全与可复现性涉及到生产环境或大规模训练时还需要注意在大规模训练前先在小 batch 上做 1 到 2 个 epoch 的完整性验证。分块正交化会引入近似误差如果任务对精度极其敏感优先使用全量 Muon 或增大 chunk_size。分布式训练时保证每个 rank 上的优化器状态一致否则会出现收敛不稳定。8. 总结与学习路线本文围绕 CMuon 展开核心收获有三点。第一理解了动量正交化的基本思想把优化器从“逐元素自适应缩放”升级为“矩阵级正交化更新”通过 Newton-Schulz 迭代逼近 SVD 的极分解结果从而降低计算成本。第二掌握了 Muon 与 CMuon 的实现差异CMuon 在原有基础上引入 Chunked Momentum Orthogonalization把大矩阵拆成块后独立正交化降低了计算和内存开销同时提高了训练稳定性。第三拿到了一个可运行的 PyTorch 示例从 Newton-Schulz 函数到 CMuon 优化器再到简易 DiT 训练循环全部代码可以直接复制到项目里改造。如果你准备继续深入可以从以下几个方向扩展阅读 Muon 原始论文和 Newton-Schulz 迭代的数值分析。把 CMuon 接入真实数据集比如 CIFAR-10 或 ImageNet 的 DiT 训练对比 FID 指标。尝试把分块思路应用到其他优化器例如分块 LAMB 或分块逆 Hessian 近似。研究不同 chunk_size 对模型收敛性和最终生成质量的影响曲线。在实际项目中我建议优先关注两个风险一是学习率量级必须重新扫描二是分块大小需要按矩阵维度调整。把这两个问题控制好CMuon 通常能在训练速度和稳定性上带来明显收益。如果本文对你有帮助可以收藏备用后续有新的实验结论我也会继续更新。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门