Scaling Muon for Diffusion Transformers
Chenghao Li, Xiao Han, Xinxin Huang, Wei Liu, Boyang Li, Bing Xiao, Heran Zhang, Juanma Perez Rua, Ke Xu, Kangning Liu, Linjun Kuang, Na Li, Tan Wang, Tian Xie, Wei Peng, Yang Pei, Yifan Xu, Yuanhao Zhai, Yuwei Lin, Zhe Wang, Zihao He, Daniel Li, Junbiao Tang, Ziyang Jiang, Dake Chen
cs.LG, cs.AI, cs.CV
2026-08-21
On 1.3B-15B DiTs, Muon beats AdamW by 12.9-19.1% best FD-DINO. Periodic Row-wise Muon matches quality, cuts optimizer time 47-54%, and reaches its best 34-65% faster.
Muon treats each two-dimensional weight as a matrix. It runs five Newton–Schulz (NS5) steps on the momentum so the update is balanced across singular directions. On LLM pretraining that has been enough to match AdamW with fewer FLOPs. Small diffusion runs hinted at the same optimization win, with a catch: loss, sample quality, and wall-clock time can rank optimizers differently. Nobody had measured this on Diffusion Transformers (DiTs) at ten-billion scale.
At that scale an optimizer is only as good as its GPU-hours. NS5 is a handful of GEMMs on the full matrix. Under FSDP, parameters and momentum are sharded by rows, so a spectral update needs an all-gather of the whole momentum, then a replicated NS5 on every rank. The extra work can cancel the step-efficiency edge. The paper asks two questions: does Muon still beat AdamW from 1.3B to 15B DiTs, and can the systems tax be cut without giving the quality back.
They train text-to-image MMDiTs from scratch on GPIC-Full (100M Flickr/Wikimedia image-text pairs, captioned by Qwen3-VL) at 512x512. Four widths: about 1.3B, 4B, 9B, 15B, same dual-stream backbone with a frozen FLUX.1-schnell VAE and frozen CLIP-L, CLIP-G, and T5-XXL text encoders. Muon touches only the 2-D matrices inside Transformer blocks, about 99% of trainable weights. Biases, norms, and everything outside the blocks stay on AdamW. All three optimizers run 60k steps, global batch 4,096, 256 H100s, PyTorch FSDP2.
Periodic Row-wise Muon splits the two geometries.
The design bet is local stability: when the momentum stays away from rank degeneracy, small changes in momentum produce only small changes in the polar direction, so a full spectral map is not needed every step. K sets the refresh rate. Gamma scales the RowNorm branch relative to NS5, because the two maps live on different feasible sets. Setting gamma=1 (reuse the Muon learning rate) has no justification. They swept K in {2,3,4} and gamma in {0.10,0.15,0.25,0.35} once on the 1.3B model for 30k steps on 16 H100s, scoring mean FD-DINO at 20k/25k/30k. Going from K=2 to K=3 cuts NS5 by 33% and raises late FD-DINO by 2.8%; K=4 saves another 25% of refreshes but costs 8.2% FD-DINO. Default is K=3, gamma=0.15, copied to every larger model with no retune.
The distributed path matches the two branches. Off-refresh, RowNorm runs on shards: wide matrices have local rows and need no optimizer collective; tall matrices all-reduce a handful of per-column norm scalars instead of gathering the full matrix. On refresh, matrices are bucketed so the next all-gather overlaps NS5 on the current bucket. For K=3 the logical optimizer payload drops to about one third of vanilla Muon.
Muon's step-wise edge holds at every scale. Validation loss stays below AdamW for the whole 60k run, and best observed FD-DINO (Frechet distance in DINOv2 space, lower is better) improves 12.9-19.1%. Replot the same curves against wall-clock and AdamW hits intermediate losses first. The per-step tax eats the algorithmic lead.
Periodic Row-wise Muon keeps the quality and converts it into time. Final-checkpoint FD-DINO beats AdamW by 11.1-17.8%. Versus vanilla Muon, best FD-DINO stays within 0.5% at 1.3B and 4B, and is about 4.5% better at 9B and 2.7% better at 15B. Time to each method's own best FD-DINO drops 33.7%, 36.1%, 64.8%, and 57.4% at the four scales.
| Scale | AdamW final FD-DINO | Muon | Periodic |
| 1.3B | 53.93 | 46.08 | 45.65 |
| 4B | 51.39 | 41.62 | 42.25 |
| 9B | 41.26 | 33.97 | 36.70 |
| 15B | 40.23 | 33.51 | 33.99 |
Those are 60k finish-line numbers, not the best point seen in training. At 9B, Periodic lifts GenEval2 AM from 51.77 to 57.33 and GM from 11.35 to 15.97. At 15B, Coverage and Density beat Muon, while GenEval2 falls from 57.93 to 52.26. The metrics do not move together.
Against vanilla Muon, optimizer time falls 46.9-54.3%, end-to-end step time 15.7-24.3%, logical communication volume 66.7% at every scale. On 15B, optimizer time goes from 2.884 to 1.317 (normalized to mean 1.3B AdamW step = 1) and step time from 6.527 to 5.002.
Ablations pin down both pieces. On 1.3B, RowNorm every step lands at best FD-DINO 51.322, essentially AdamW's 51.56. Periodic with gamma=1 gets 44.294; gamma=0.15 gets 41.913, next to vanilla Muon's 41.733. On 15B systems, a naive periodic baseline that still materializes full momentum every step sits at step time 5.685. Sharded RowNorm cuts communication from 28.62 to 9.54 GiB/rank/step; bucketing and overlap take step time to 5.002.
If a team already wants Muon on a large DiT, this is the variant that turns the quality edge into fewer GPU-hours. K and gamma selected at 1.3B transferred to 15B, so the defaults are at least sticky inside this MMDiT family.
It is not a new optimizer family. Lowering spectral-refresh frequency, row-wise normalization, and sharded Muon runtimes all exist. The paper's job is to run periodic RowNorm plus shard-local stats through real 1.3B-15B DiT training and report both sample quality and systems cost. Step-wise loss oversells vanilla Muon. Wall-clock alone undersells it. Looking at both is what makes Periodic the rational default.
The authors list the obvious holes: one MMDiT family, one GPIC dataset, one 512 resolution, one 32-node H100 plus FSDP2 setup. Video DiTs, higher resolution, and other parallelism are untested. K and gamma are global constants. Refresh steps still all-gather and materialize full momentum, so the speedup will move with topology.
A few numbers need a second look. Best observed FD-DINO cherry-picks the training trajectory; Muon and Periodic peak at different steps, and some checkpoints swing hard (Periodic 4B hits 62.89 at 45k; 15B hits 52.25 at 50k). At the 60k finish line, 9B Periodic FD-DINO is worse than Muon (36.70 vs 33.97), and 15B GenEval2 is worse. The headline mixes best-of-run and final checkpoint. Learning rates are also asymmetric: Muon's main LR is roughly 10-17x AdamW's, with a separate 1e-4 AdamW group. No code or kernel is released; every timing number is from this 256-GPU cluster.