MCPcopy Create free account
hub / github.com/AlmondGod/tinyworlds / _create_muon_split

Function _create_muon_split

utils/optimizer_utils.py:28–69  ·  view source on GitHub ↗
(model, args)

Source from the content-addressed store, hash-verified

26
27
28def _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
72def _split_decay_params(model):

Callers 1

create_optimizerFunction · 0.85

Calls 1

MuonClass · 0.90

Tested by

no test coverage detected