File size: 12,025 Bytes
bcda938
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
"""Hybrid language model: Gated DeltaNet + gated attention (3:1), Block Attention Residuals.



Architecture follows Qwen3.5 (3 Gated DeltaNet layers per gated full-attention layer,

sigmoid output gate on attention, QK-norm) with Kimi's Block Attention Residuals replacing

the plain residual stream. arch="dense" swaps every DeltaNet layer for gated attention,

which gives a like-for-like Transformer baseline.

"""
import math
from dataclasses import dataclass

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint


@dataclass
class ModelConfig:
    vocab_size: int = 32768
    n_layer: int = 24
    d_model: int = 1024
    arch: str = "hybrid"      # "hybrid": every `attn_every`-th layer is attention, rest DeltaNet; "dense": all attention
    attn_every: int = 4
    n_head: int = 16          # attention query heads
    n_kv_head: int = 4
    head_dim: int = 128
    gdn_heads: int = 16
    gdn_head_dim: int = 128
    ffn_dim: int = 3584
    attnres_blocks: int = 8   # Block AttnRes: the 2*n_layer sublayers are split into this many blocks
    rope_theta: float = 10000.0
    max_seq_len: int = 2048
    softcap: float = 15.0
    grad_ckpt: bool = False   # recompute sublayers in backward to save VRAM


PRESETS = {
    # smoke tests only
    "tiny": dict(n_layer=4, d_model=256, n_head=4, n_kv_head=2, head_dim=64, gdn_heads=4, gdn_head_dim=64,
                 ffn_dim=768, attnres_blocks=4),
    # ~140M non-embedding: architecture bake-off size
    "small": dict(n_layer=12, d_model=768, n_head=12, n_kv_head=3, head_dim=128, gdn_heads=12, gdn_head_dim=128,
                  ffn_dim=2688, attnres_blocks=8),
    # ~500M non-embedding: Qwen3.5-0.8B layout with a 32k vocab
    "base": dict(n_layer=24, d_model=1024, n_head=16, n_kv_head=4, head_dim=128, gdn_heads=16, gdn_head_dim=128,
                 ffn_dim=3584, attnres_blocks=8),
}


def rms(x):
    return F.rms_norm(x, (x.size(-1),))


def inv_rms(x, eps=1e-6):
    return torch.rsqrt(x.float().pow(2).mean(-1) + eps)


class GatedAttention(nn.Module):
    """Causal GQA attention with QK-norm, RoPE and Qwen's elementwise sigmoid output gate."""

    def __init__(self, c: ModelConfig):
        super().__init__()
        self.nh, self.nkv, self.hd = c.n_head, c.n_kv_head, c.head_dim
        self.q_proj = nn.Linear(c.d_model, 2 * self.nh * self.hd, bias=False)  # query and gate
        self.kv_proj = nn.Linear(c.d_model, 2 * self.nkv * self.hd, bias=False)
        self.o_proj = nn.Linear(self.nh * self.hd, c.d_model, bias=False)
        inv_freq = 1.0 / (c.rope_theta ** (torch.arange(0, self.hd, 2).float() / self.hd))
        freqs = torch.outer(torch.arange(c.max_seq_len).float(), inv_freq)
        self.register_buffer("cos", freqs.cos()[None, :, None, :], persistent=False)
        self.register_buffer("sin", freqs.sin()[None, :, None, :], persistent=False)

    def rope(self, x):
        T = x.size(1)
        cos, sin = self.cos[:, :T].to(x.dtype), self.sin[:, :T].to(x.dtype)
        x1, x2 = x.chunk(2, -1)
        return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], -1)

    def forward(self, x):
        B, T, _ = x.shape
        q, gate = self.q_proj(x).split(self.nh * self.hd, -1)
        q = q.view(B, T, self.nh, self.hd)
        k, v = self.kv_proj(x).view(B, T, 2, self.nkv, self.hd).unbind(2)
        q, k = self.rope(rms(q)), self.rope(rms(k))
        if self.nkv != self.nh:
            k = k.repeat_interleave(self.nh // self.nkv, dim=2)
            v = v.repeat_interleave(self.nh // self.nkv, dim=2)
        y = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), is_causal=True)
        y = y.transpose(1, 2).reshape(B, T, self.nh * self.hd)
        return self.o_proj(y * torch.sigmoid(gate))


class DeltaNet(nn.Module):
    """Gated DeltaNet (linear attention with a delta-rule memory) from flash-linear-attention."""

    def __init__(self, c: ModelConfig):
        super().__init__()
        from fla.layers import GatedDeltaNet
        self.gdn = GatedDeltaNet(hidden_size=c.d_model, head_dim=c.gdn_head_dim, num_heads=c.gdn_heads,
                                 expand_v=1.0, mode="chunk", use_gate=True, use_short_conv=True, conv_size=4)

    def forward(self, x):
        return self.gdn(x)[0]


class SwiGLU(nn.Module):
    def __init__(self, c: ModelConfig):
        super().__init__()
        self.w_gate = nn.Linear(c.d_model, c.ffn_dim, bias=False)
        self.w_up = nn.Linear(c.d_model, c.ffn_dim, bias=False)
        self.w_down = nn.Linear(c.ffn_dim, c.d_model, bias=False)

    def forward(self, x):
        return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))


def attnres_mix(sources, source_inv_rms, query):
    """Block AttnRes: softmax over sources of <query, RMSNorm(source)>, then a weighted sum of sources."""
    logits = torch.stack([(s @ query.to(s.dtype)).float() * r for s, r in zip(sources, source_inv_rms)])
    weights = logits.softmax(0)
    h = weights[0, ..., None] * sources[0]
    for w, s in zip(weights[1:], sources[1:]):
        h = h + w[..., None] * s
    return h


class LM(nn.Module):
    def __init__(self, c: ModelConfig):
        super().__init__()
        self.config = c
        self.embed = nn.Embedding(c.vocab_size, c.d_model)
        self.sublayers = nn.ModuleList()
        for i in range(c.n_layer):
            is_attn = c.arch == "dense" or (i % c.attn_every == c.attn_every - 1)
            self.sublayers.append(GatedAttention(c) if is_attn else DeltaNet(c))
            self.sublayers.append(SwiGLU(c))
        n_sub = len(self.sublayers)
        self.block_size = math.ceil(n_sub / c.attnres_blocks)
        # one pseudo-query per sublayer plus one for the final read-out; zero init = uniform average at start
        self.attnres_queries = nn.Parameter(torch.zeros(n_sub + 1, c.d_model))
        self.lm_head = nn.Linear(c.d_model, c.vocab_size, bias=False)
        self.init_weights()

    @torch.no_grad()
    def init_weights(self):
        c = self.config
        nn.init.normal_(self.embed.weight, std=1.0)
        nn.init.normal_(self.lm_head.weight, std=0.001)
        s = 3 ** 0.5 * c.d_model ** -0.5
        for m in self.sublayers:
            if isinstance(m, GatedAttention):
                nn.init.uniform_(m.q_proj.weight, -s, s)
                nn.init.uniform_(m.kv_proj.weight, -s, s)
                nn.init.zeros_(m.o_proj.weight)
            elif isinstance(m, SwiGLU):
                nn.init.uniform_(m.w_gate.weight, -s, s)
                nn.init.uniform_(m.w_up.weight, -s, s)
                nn.init.zeros_(m.w_down.weight)
            else:  # DeltaNet keeps FLA's init for its gates/decay; match the rest to the attention layers
                g = m.gdn
                for lin in (g.q_proj, g.k_proj, g.v_proj, g.g_proj):
                    nn.init.uniform_(lin.weight, -s, s)
                nn.init.zeros_(g.o_proj.weight)

    def num_params(self, non_embedding=True):
        n = sum(p.numel() for p in self.parameters())
        return n - self.embed.weight.numel() - self.lm_head.weight.numel() if non_embedding else n

    def _run(self, sub, x):
        if self.config.grad_ckpt and self.training:
            return checkpoint(sub, x, use_reentrant=False)
        return sub(x)

    def hidden(self, idx):
        """Final normalized hidden states [B, T, d_model]."""
        x0 = rms(self.embed(idx))
        blocks, blocks_inv = [x0], [inv_rms(x0)]   # completed block sums (block 0 = token embedding)
        partial = None                             # running sum of sublayer outputs in the current block
        for j, sub in enumerate(self.sublayers):
            srcs = blocks if partial is None else blocks + [partial]
            inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
            h = attnres_mix(srcs, inv, self.attnres_queries[j])
            out = self._run(sub, rms(h)).float()
            partial = out if partial is None else partial + out
            if (j + 1) % self.block_size == 0:
                blocks.append(partial)
                blocks_inv.append(inv_rms(partial))
                partial = None
        srcs = blocks if partial is None else blocks + [partial]
        inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
        return rms(attnres_mix(srcs, inv, self.attnres_queries[-1]))

    def forward(self, idx, targets=None):
        h = self.hidden(idx)
        if targets is None:
            return self._logits(h)
        # Loss in chunks, recomputing each chunk's logits in backward, so the full
        # (tokens x vocab) fp32 logit tensor never has to sit in VRAM.
        # Targets of -1 (prompts, padding during fine-tuning) are ignored.
        h, targets = h.flatten(0, 1), targets.flatten()
        loss = 0.0
        for i in range(0, h.size(0), self.LOSS_CHUNK):
            loss = loss + checkpoint(self._chunk_loss, h[i:i + self.LOSS_CHUNK], targets[i:i + self.LOSS_CHUNK],
                                     use_reentrant=False)
        return loss / (targets >= 0).sum().clamp(min=1)

    LOSS_CHUNK = 4096

    def _logits(self, h):
        logits = self.lm_head(h).float()
        if self.config.softcap:
            logits = self.config.softcap * torch.tanh(logits / self.config.softcap)
        return logits

    def _chunk_loss(self, h, targets):
        return F.cross_entropy(self._logits(h), targets, reduction="sum", ignore_index=-1)

    def token_logprobs(self, idx, targets):
        """Log-probability of each target token, [B, T] (0 where the target is -1). Chunked like the loss."""
        h, t = self.hidden(idx).flatten(0, 1), targets.flatten()
        parts = [checkpoint(self._chunk_logprobs, h[i:i + self.LOSS_CHUNK], t[i:i + self.LOSS_CHUNK], use_reentrant=False)
                 for i in range(0, h.size(0), self.LOSS_CHUNK)]
        return torch.cat(parts).view(targets.shape)

    def _chunk_logprobs(self, h, targets):
        lp = torch.log_softmax(self._logits(h), -1).gather(1, targets.clamp(min=0)[:, None])[:, 0]
        return lp * (targets >= 0)

    def param_groups(self):
        """Split parameters: hidden matrices go to Muon, everything else to AdamW."""
        muon, other = [], []
        for name, p in self.sublayers.named_parameters():
            (muon if p.ndim == 2 and min(p.shape) >= 64 else other).append(p)
        other.append(self.attnres_queries)
        return dict(muon=muon, embed=[self.embed.weight], head=[self.lm_head.weight], other=other)


def resize_vocab(state_dict, new_rows):
    """Grow the embedding and output matrices for added tokens; new rows start at the mean of the old ones."""
    for key in ("embed.weight", "lm_head.weight"):
        w = state_dict[key]
        if w.size(0) < new_rows:
            extra = w.float().mean(0, keepdim=True).expand(new_rows - w.size(0), -1).to(w.dtype)
            state_dict[key] = torch.cat([w, extra])
    return state_dict


@torch.no_grad()
def generate(model, idx, max_new_tokens, temperature=0.8, top_k=50):
    """Simple sampling that re-runs the full context each step (no cache; fine for short samples)."""
    for _ in range(max_new_tokens):
        ctx = idx[:, -model.config.max_seq_len:]
        with torch.autocast("cuda", dtype=torch.bfloat16):
            logits = model(ctx)[:, -1, :].float()
        if temperature <= 0:
            nxt = logits.argmax(-1, keepdim=True)
        else:
            logits = logits / temperature
            if top_k:
                v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
                logits[logits < v[:, [-1]]] = -float("inf")
            nxt = torch.multinomial(F.softmax(logits, -1), 1)
        idx = torch.cat([idx, nxt], 1)
    return idx