File size: 5,338 Bytes
4397e12
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
"""Optimizers, following DeepSeek V4.1 §2.5:
  * Muon (Nesterov, decoupled WD) for backbone matrices; head-wise for query projections.
  * Momentum + Sinkhorn balancing (Algorithm 1) for token embedding, LM head and Engram tables.
  * AdamW for anything else (none by default: norms are parameter-free).
Update RMS for Muon is matched to AdamW via 0.2*sqrt(max(m, n)) (Moonlight), and the Sinkhorn
update uses gamma=0.18, so one learning rate works across groups.
"""
from __future__ import annotations

import math

import torch


@torch.no_grad()
def newton_schulz(G: torch.Tensor, steps: int = 5, eps: float = 1e-7) -> torch.Tensor:
    """Quintic NS orthogonalization; works on (..., m, n)."""
    a, b, c = 3.4445, -4.7750, 2.0315
    X = G.bfloat16()
    transpose = X.size(-2) > X.size(-1)
    if transpose:
        X = X.mT
    X = X / (X.norm(dim=(-2, -1), keepdim=True) + eps)
    for _ in range(steps):
        A = X @ X.mT
        B = b * A + c * A @ A
        X = a * X + B @ X
    if transpose:
        X = X.mT
    return X.to(G.dtype)


class Muon(torch.optim.Optimizer):
    def __init__(self, params, lr=3e-3, momentum=0.95, weight_decay=0.0, ns_steps=5):
        super().__init__(params, dict(lr=lr, momentum=momentum, weight_decay=weight_decay, ns_steps=ns_steps))

    @torch.no_grad()
    def step(self):
        for g in self.param_groups:
            for p in g["params"]:
                if p.grad is None:
                    continue
                st = self.state[p]
                if "buf" not in st:
                    st["buf"] = torch.zeros_like(p, dtype=torch.float32)
                buf = st["buf"]
                grad = p.grad.float()
                buf.lerp_(grad, 1 - g["momentum"])
                upd = grad.lerp(buf, g["momentum"])  # Nesterov
                heads = getattr(p, "_muon_heads", None)
                if heads:  # head-wise: orthogonalize each head's (head_dim, d_model) slice separately
                    u = upd.view(heads, -1, upd.size(-1))
                    u = newton_schulz(u, g["ns_steps"]).view_as(upd)
                    m, n = u.size(0) // heads, u.size(1)
                else:
                    u = newton_schulz(upd, g["ns_steps"])
                    m, n = u.shape
                scale = 0.2 * math.sqrt(max(m, n))
                if g["weight_decay"]:
                    p.mul_(1 - g["lr"] * g["weight_decay"])
                p.add_(u.to(p.dtype), alpha=-g["lr"] * scale)


class SinkhornMomentum(torch.optim.Optimizer):
    """V4.1 Algorithm 1: Nesterov momentum, mask near-zero rows, K (odd) alternating row/col L2
    normalizations ending on rows, scale by sqrt(n) for unit row RMS, lr correction gamma."""

    def __init__(self, params, lr=3e-3, momentum=0.95, K=5, gamma=0.18, tau=0.1, eps=1e-8):
        assert K % 2 == 1
        super().__init__(params, dict(lr=lr, momentum=momentum, K=K, gamma=gamma, tau=tau, eps=eps))

    @torch.no_grad()
    def step(self):
        for g in self.param_groups:
            for p in g["params"]:
                if p.grad is None:
                    continue
                st = self.state[p]
                if "buf" not in st:
                    st["buf"] = torch.zeros_like(p, dtype=torch.float32)
                buf = st["buf"]
                grad = p.grad.float()
                buf.lerp_(grad, 1 - g["momentum"])
                U = grad.lerp(buf, g["momentum"])
                rho = U.norm(dim=1)
                # Mean over rows that carry signal at all: huge hash tables are mostly untouched,
                # and a mean over all rows would let stale decayed momentum through the mask.
                nz = rho > 0
                rbar = rho[nz].mean() if nz.any() else rho.new_zeros(())
                U = U * (rho > g["tau"] * rbar)[:, None]
                for k in range(1, g["K"] + 1):
                    if k % 2 == 1:
                        U = U / (U.norm(dim=1, keepdim=True) + g["eps"])
                    else:
                        U = U / (U.norm(dim=0, keepdim=True) + g["eps"])
                U = U * math.sqrt(U.size(1))
                p.add_(U.to(p.dtype), alpha=-g["lr"] * g["gamma"])


def build_optimizers(model, lr=3e-3, weight_decay=0.0):
    sink, muon, adam = [], [], []
    sink_ids = {id(model.embed.weight), id(model.head.weight)}
    for m in model.engram_modules():
        sink_ids.add(id(m.table.weight))
    for p in model.parameters():
        if not p.requires_grad:
            continue
        if id(p) in sink_ids:
            sink.append(p)
        elif p.ndim == 2:
            muon.append(p)
        else:
            adam.append(p)
    opts = [Muon(muon, lr=lr, weight_decay=weight_decay), SinkhornMomentum(sink, lr=lr)]
    if adam:
        opts.append(torch.optim.AdamW(adam, lr=lr, betas=(0.9, 0.95), weight_decay=0.0))
    for o in opts:
        for g in o.param_groups:
            g["base_lr"] = g["lr"]
    return opts


def wsd_lr(step: int, total: int, warmup: int = 200, decay_frac: float = 0.2) -> float:
    """Warmup-stable-decay multiplier. Decay can also be triggered early by passing a smaller total."""
    if step < warmup:
        return (step + 1) / warmup
    decay_start = int(total * (1 - decay_frac))
    if step < decay_start:
        return 1.0
    return max(0.0, (total - step) / max(1, total - decay_start))