| 26 | |
| 27 | |
| 28 | def _create_muon_split(model, args): |
| 29 | from models.muon import Muon |
| 30 | |
| 31 | muon_params = [] |
| 32 | adamw_decay = [] |
| 33 | adamw_no_decay = [] |
| 34 | |
| 35 | for name, param in model.named_parameters(): |
| 36 | if not param.requires_grad: |
| 37 | continue |
| 38 | # Muon only makes sense for 2D weight matrices (not embeddings, not biases) |
| 39 | if param.ndim == 2 and "embed" not in name: |
| 40 | muon_params.append(param) |
| 41 | elif param.ndim == 1 or name.endswith(".bias") or "norm" in name: |
| 42 | adamw_no_decay.append(param) |
| 43 | else: |
| 44 | adamw_decay.append(param) |
| 45 | |
| 46 | lr = args.learning_rate |
| 47 | momentum = getattr(args, "muon_momentum", 0.95) |
| 48 | backend_steps = getattr(args, "muon_backend_steps", 5) |
| 49 | |
| 50 | optimizers = [] |
| 51 | |
| 52 | if muon_params: |
| 53 | muon_opt = Muon( |
| 54 | muon_params, lr=lr, momentum=momentum, |
| 55 | backend_steps=backend_steps, weight_decay=0.01, |
| 56 | ) |
| 57 | optimizers.append(muon_opt) |
| 58 | |
| 59 | # AdamW for the rest |
| 60 | adamw_groups = [] |
| 61 | if adamw_decay: |
| 62 | adamw_groups.append({"params": adamw_decay, "weight_decay": 0.01}) |
| 63 | if adamw_no_decay: |
| 64 | adamw_groups.append({"params": adamw_no_decay, "weight_decay": 0}) |
| 65 | if adamw_groups: |
| 66 | adamw_opt = optim.AdamW(adamw_groups, lr=lr, betas=(0.9, 0.999), eps=1e-8, fused=True) |
| 67 | optimizers.append(adamw_opt) |
| 68 | |
| 69 | return optimizers |
| 70 | |
| 71 | |
| 72 | def _split_decay_params(model): |