musabc commited on
Commit
073320f
·
verified ·
1 Parent(s): b5dc2be

upload muon.py

Browse files
Files changed (1) hide show
  1. muon.py +121 -0
muon.py ADDED
@@ -0,0 +1,121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Muon Optimizer — Keller Jordan, NanoGPT speedrun.
3
+
4
+ Newton-Schulz iterasyonu ile orthogonalize edilmiş momentum.
5
+ 2D ağırlıklara (Linear weight) uygulanır. 1D parametreler (norm weight,
6
+ bias, embedding) AdamW'ye verilir.
7
+
8
+ Referans:
9
+ https://github.com/KellerJordan/modded-nanogpt
10
+ https://kellerjordan.github.io/posts/muon/
11
+
12
+ Kullanim:
13
+ # Param ayri:
14
+ muon_params = [p for p in model.parameters() if p.dim() >= 2 and p.requires_grad]
15
+ other_params = [p for p in model.parameters() if p.dim() < 2 and p.requires_grad]
16
+
17
+ # Embedding ve lm_head'i muon'dan ayir (yaygin best practice)
18
+ embed_params = [model.wte.weight] # tied ise lm_head dahil
19
+ muon_params = [p for p in muon_params if not any(p is e for e in embed_params)]
20
+ other_params = other_params + embed_params
21
+
22
+ optimizer_muon = Muon(muon_params, lr=2e-2, momentum=0.95)
23
+ optimizer_adam = torch.optim.AdamW(other_params, lr=3e-4, ...)
24
+ """
25
+
26
+ import torch
27
+
28
+
29
+ @torch.no_grad()
30
+ def newton_schulz(G: torch.Tensor, steps: int = 5) -> torch.Tensor:
31
+ """G matrisini orthogonalize et (yaklasik USV^T -> UV^T).
32
+
33
+ Newton-Schulz quintic iteration. bf16'da kararli, hizli.
34
+ """
35
+ assert G.ndim == 2
36
+ a, b, c = (3.4445, -4.7750, 2.0315)
37
+ X = G.to(torch.bfloat16)
38
+ # Boyut yonune gore transpose (her iki yonde de calissin)
39
+ if X.size(0) > X.size(1):
40
+ X = X.T
41
+ # Spektral normu yaklasik 1'e cek
42
+ X = X / (X.norm() + 1e-7)
43
+ for _ in range(steps):
44
+ A = X @ X.T
45
+ B = b * A + c * (A @ A)
46
+ X = a * X + B @ X
47
+ if G.size(0) > G.size(1):
48
+ X = X.T
49
+ return X.to(G.dtype)
50
+
51
+
52
+ class Muon(torch.optim.Optimizer):
53
+ """Muon: Momentum + orthogonalize edilmiş update.
54
+
55
+ Sadece 2D parametreler için. 1D'leri AdamW ile ayrı eğit.
56
+
57
+ Args:
58
+ params: 2D parametreler iterable
59
+ lr: 0.02 (AdamW'nin ~50x'i, çünkü update'ler ortonormal)
60
+ momentum: 0.95
61
+ nesterov: True (genelde daha iyi)
62
+ ns_steps: Newton-Schulz iter sayisi (5 default)
63
+ """
64
+ def __init__(self, params, lr=0.02, momentum=0.95, nesterov=True, ns_steps=5):
65
+ defaults = dict(lr=lr, momentum=momentum, nesterov=nesterov, ns_steps=ns_steps)
66
+ super().__init__(params, defaults)
67
+
68
+ @torch.no_grad()
69
+ def step(self, closure=None):
70
+ loss = None
71
+ if closure is not None:
72
+ with torch.enable_grad():
73
+ loss = closure()
74
+
75
+ for group in self.param_groups:
76
+ lr = group["lr"]
77
+ momentum = group["momentum"]
78
+ nesterov = group["nesterov"]
79
+ ns_steps = group["ns_steps"]
80
+
81
+ for p in group["params"]:
82
+ if p.grad is None:
83
+ continue
84
+ if p.ndim < 2:
85
+ raise ValueError(
86
+ f"Muon sadece >=2D param destekler, {p.ndim}D bulundu. "
87
+ "1D paramları AdamW'ye ver.")
88
+
89
+ g = p.grad
90
+ state = self.state[p]
91
+ if "momentum_buffer" not in state:
92
+ state["momentum_buffer"] = torch.zeros_like(g)
93
+
94
+ buf = state["momentum_buffer"]
95
+ buf.mul_(momentum).add_(g)
96
+
97
+ # Nesterov momentum
98
+ if nesterov:
99
+ g = g.add(buf, alpha=momentum)
100
+ else:
101
+ g = buf
102
+
103
+ # Reshape if needed (e.g., conv weight)
104
+ original_shape = g.shape
105
+ if g.ndim > 2:
106
+ g = g.view(g.size(0), -1)
107
+
108
+ # Newton-Schulz orthogonalization
109
+ g_orth = newton_schulz(g, steps=ns_steps)
110
+
111
+ # Scale: sqrt(max(out, in) / min(out, in)) — ~spectral norm
112
+ # Ya da basitce sqrt(d_out / d_in) gibi.
113
+ # Modded-nanogpt: scale = max(1, p.shape[0]/p.shape[1]) ** 0.5
114
+ scale = max(1.0, g_orth.size(0) / g_orth.size(1)) ** 0.5
115
+
116
+ # Geri reshape
117
+ g_orth = g_orth.view(original_shape)
118
+
119
+ p.add_(g_orth, alpha=-lr * scale)
120
+
121
+ return loss