File size: 8,746 Bytes
3b1575e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
# model_dna_peer_v21.py - v21 ternary controller + SHARED-POOL PEER experts.
# Changes vs model_dna_peer.py (v20), per docs/2026-07-25_v21:
#   P1  ONE shared expert pool used by 3 PEER layers at stride 4 (Memory Layers at
#       Scale, arXiv:2412.09764) instead of 16 private pools -> same params, 16x
#       fewer gathers. Other 13 blocks keep a cheap dense SwiGLU FFN.
#   P2  F.embedding_bag(mode='sum', per_sample_weights=...) for the down-projection
#       (numerically identical to the einsum contraction, ~3x faster).
#   P3  expert tables held in bf16 (nn.Embedding is autocast-EXEMPT, so the v20 code
#       silently gathered fp32 and then cast the whole [N,h,k,d] tensor).
#   P4  sparse=True on the pool -> sparse grads, no dense 507M grad / Adam state.
#   P7  h*k raised from 8 to 32 (heads=4, topk=8) toward PEER's recommended
#       granularity, plus query BatchNorm for expert usage (PEER arXiv:2407.04153).
import sys; sys.path.insert(0, '/root/dna')
import torch, torch.nn as nn, torch.nn.functional as F
from model_dna import DnaChat, diag_scan, counts

def tern(w):                              # BitNet b1.58 ternary STE (weight-only)
    s = w.abs().mean().clamp_min(1e-5)
    wt = (w / s).round().clamp_(-1, 1) * s
    return w + (wt - w).detach()

class BitLinear(nn.Module):
    def __init__(self, i, o):
        super().__init__(); self.weight = nn.Parameter(torch.empty(o, i)); nn.init.normal_(self.weight, std=0.02)
    def forward(self, x): return F.linear(x, tern(self.weight))

class ExpertPool(nn.Module):
    """ONE shared bank of E = nk^2 rank-1 experts, reused by every PEER layer (P1),
    with sparse gradients (P4).

    dtype=None keeps fp32 master weights (DEFAULT). P3 (bf16 tables) halves gather
    bandwidth but MEASURED bf16 SparseAdam updates carry quantization noise
    (mean|delta| 1.8e-4 vs 1.5e-3 clean), so it is opt-in: pass dtype=torch.bfloat16."""
    def __init__(self, nk, d, sparse=True, dtype=None):
        super().__init__()
        self.nk, self.E, self.d = nk, nk * nk, d
        self.up = nn.Embedding(self.E, d, sparse=sparse)
        self.down = nn.Embedding(self.E, d, sparse=sparse)
        nn.init.normal_(self.up.weight, std=0.02)
        nn.init.normal_(self.down.weight, std=0.02 / (d ** 0.5))
        if dtype is not None:
            self.up.to(dtype); self.down.to(dtype)

class PeerFFN(nn.Module):
    """Product-key retrieval over a SHARED expert pool.
    out = sum_k softmax(score)_k * Wdown_k * silu(Wup_k . x);  active FLOPs ~ h*k*d."""
    def __init__(self, pool, d, dk=128, topk=8, heads=4, qbn=True):
        super().__init__()
        self.pool, self.d, self.dk, self.topk, self.h = pool, d, dk, topk, heads
        self.nk = pool.nk
        self.sparse = pool.up.sparse   # F.embedding/_bag need this EXPLICITLY
        self.q = nn.Linear(d, heads * 2 * dk, bias=False)
        self.Ka = nn.Parameter(torch.randn(heads, self.nk, dk) * 0.02)
        self.Kb = nn.Parameter(torch.randn(heads, self.nk, dk) * 0.02)
        # query BatchNorm (PEER 3.3 / PKM): keeps expert usage near 100%
        self.qbn = nn.BatchNorm1d(heads * 2 * dk) if qbn else None

    def route(self, x):
        """Product-key routing. Pure dense ops -> safe for torch.compile."""
        N = x.size(0); k = self.topk; nk = self.nk; h = self.h
        q = self.q(x)
        if self.qbn is not None: q = self.qbn(q)
        q = q.view(N, h, 2, self.dk)
        qa, qb = q[:, :, 0], q[:, :, 1]                    # [N,h,dk]
        sa = torch.einsum('nhd,hkd->nhk', qa, self.Ka)     # [N,h,nk]
        sb = torch.einsum('nhd,hkd->nhk', qb, self.Kb)
        va, ia = sa.topk(k, -1); vb, ib = sb.topk(k, -1)   # [N,h,k]
        cs = (va.unsqueeze(-1) + vb.unsqueeze(-2)).reshape(N, h, -1)
        ci = (ia.unsqueeze(-1) * nk + ib.unsqueeze(-2)).reshape(N, h, -1)
        tv, ti = cs.topk(k, -1)                            # [N,h,k]
        eid = torch.gather(ci, -1, ti).reshape(N * h, k)    # flatten heads into bags
        gate = torch.softmax(tv, -1).reshape(N * h, k)
        return eid, gate

    def gather(self, x, eid, gate):
        """Sparse expert gather. Kept OUT of the compiled graph: inductor cannot
        lower aten.embedding with sparse=True (LoweringException)."""
        N = x.size(0); h = self.h
        W = self.pool.up.weight
        u = F.embedding(eid, W, sparse=self.sparse)                 # [N*h,k,d]
        xr = x.unsqueeze(1).expand(N, h, self.d).reshape(N * h, self.d).to(W.dtype)
        hid = F.silu(torch.bmm(u, xr.unsqueeze(-1)).squeeze(-1))    # [N*h,k]
        # P2: weighted-sum gather == einsum('nk,nkd->nd'), ~3x faster, bit-equal
        out = F.embedding_bag(eid, self.pool.down.weight, mode='sum',
                              per_sample_weights=(gate.to(W.dtype) * hid),
                              sparse=self.sparse)
        return out.reshape(N, h, self.d).sum(1).to(x.dtype) / h

    def forward(self, x):                        # x: [N, d]
        eid, gate = self.route(x)
        return self.gather(x, eid, gate)

class DenseFFN(nn.Module):
    """Cheap ternary SwiGLU FFN for the 13 non-PEER blocks."""
    def __init__(self, d, ff):
        super().__init__(); self.up = BitLinear(d, 2 * ff); self.down = BitLinear(ff, d)
    def forward(self, x):
        a, b = self.up(x).chunk(2, -1); return self.down(F.silu(b) * a)

class FastBlockTP(nn.Module):
    def __init__(self, d, ff, pool=None, dk=128, topk=8, heads=4):
        super().__init__()
        self.n1 = nn.LayerNorm(d); self.proj = BitLinear(d, 4 * d); self.o = BitLinear(d, d)
        self.decay = nn.Parameter(torch.full((d,), 2.0))
        self.n2 = nn.LayerNorm(d)
        self.is_peer = pool is not None
        self.ffn = PeerFFN(pool, d, dk, topk, heads) if self.is_peer else DenseFFN(d, ff)
    def recur(self, x, C=16):
        k, v, r, g = self.proj(self.n1(x)).chunk(4, -1)
        g = torch.sigmoid(g + self.decay); u = (1 - g) * torch.tanh(k)
        s = diag_scan(g, u, C)
        return x + self.o(torch.sigmoid(r) * s * torch.sigmoid(v))

    def forward(self, x, C=16):
        x = self.recur(x, C)
        B, T, d = x.shape
        h = self.n2(x).reshape(B * T, d)
        return x + self.ffn(h).reshape(B, T, d)

class DnaPeerV21(DnaChat):
    """3 shared-pool PEER layers at stride 4; dense ternary FFN elsewhere."""
    def __init__(self, vocab=32000, d=512, layers=16, ff=1536, a=2048, b=2048,
                 nk=704, dk=128, topk=8, pheads=4, peer_layers=(4, 8, 12), chunk=16,
                 pool_dtype=None):
        super().__init__(vocab, d, layers, ff, a, b, chunk)
        self.nk, self.dk, self.topk, self.pheads = nk, dk, topk, pheads
        self.peer_layers = tuple(peer_layers)
        self.pool = ExpertPool(nk, d, dtype=pool_dtype)
        self.blocks = nn.ModuleList([
            FastBlockTP(d, ff, self.pool if i in self.peer_layers else None, dk, topk, pheads)
            for i in range(layers)])
    def config(self):
        c = super().config()
        c.update(nk=self.nk, dk=self.dk, topk=self.topk, pheads=self.pheads,
                 peer_layers=list(self.peer_layers)); return c

def peer_counts(m):
    """(total controller params, active params per token)."""
    ctrl = sum(p.numel() for p in m.parameters())
    active = 0
    for blk in m.blocks:
        active += blk.proj.weight.numel() + blk.o.weight.numel()
        if blk.is_peer:
            active += blk.ffn.q.weight.numel() + m.pheads * m.topk * 2 * m.d
        else:
            active += blk.ffn.up.weight.numel() + blk.ffn.down.weight.numel()
    active += m.embed.weight.numel()
    active += m.ra.weight.numel() + m.rb.weight.numel()
    return ctrl, active

def compile_(m):
    """Compile everything EXCEPT the sparse expert gathers.

    torch.compile was this project's single biggest lever (v8 T2: 3.1x), but inductor
    raises LoweringException on aten.embedding with sparse=True. So we compile the
    recurrence of every block, the dense FFNs, and the PEER routers, and leave only
    PeerFFN.gather (the sparse part) eager."""
    for blk in m.blocks:
        blk.recur = torch.compile(blk.recur)
        if blk.is_peer: blk.ffn.route = torch.compile(blk.ffn.route)
        else: blk.ffn = torch.compile(blk.ffn)
    return m


def param_groups(m):
    """Split params: Muon (2D recurrence), sparse pool, dense AdamW (P4/P6)."""
    muon, sparse, dense = [], [], []
    pool_ids = {id(m.pool.up.weight), id(m.pool.down.weight)}
    muon_ids = set()
    for blk in m.blocks:
        muon += [blk.proj.weight, blk.o.weight]
    muon_ids = {id(p) for p in muon}
    for p in m.parameters():
        if id(p) in pool_ids: sparse.append(p)
        elif id(p) in muon_ids: pass
        else: dense.append(p)
    return muon, sparse, dense