"""modeling_bind2_0.py -- bind2_0 architecture: forced-bottleneck delta-overwrite recurrence (2026-07-12). d=384, L=12, 3:1 GatedDeltaNet:chunked-attn (9 GDN + 3 chunked-attn), ~24M (matched to bind1 23.9M). - GDN layer (fla GatedDeltaNet): delta-rule state S_t = a_t S_{t-1}(I - b_t k k^T) + b_t v k^T, o_t = S_t q_t. Per-token state UPDATE carried along the sequence = the "apply an operation to update state" mechanism (the verb-dynamics bind1 lacks; bind1 only labels roles). - CHUNKED attention (forced bottleneck, P0): strict intra-chunk, NON-overlapping mask M[t,s]=1 iff floor(t/C)==floor(s/C) and s<=t. Cross-chunk info can flow ONLY through S. - Unified readout: LM head on the residual stream after all blocks (GDN o_t already added to residual) -> every position, no hard routing. """ import math import torch, torch.nn as nn, torch.nn.functional as F from fla.layers import GatedDeltaNet # --- self-contained generic primitives (bind2_0 is architecturally independent of bind1/mono; # these are standard transformer building blocks, copied here so modeling_bind2_0.py stands alone) --- def build_rope(T, D, device, base=10000.0): inv = 1.0 / (base ** (torch.arange(0, D, 2, device=device).float() / D)) t = torch.arange(T, device=device).float() f = torch.outer(t, inv) emb = torch.cat([f, f], dim=-1) return emb.cos(), emb.sin() def rotate_half(x): x1, x2 = x.chunk(2, dim=-1) return torch.cat((-x2, x1), dim=-1) def apply_rope(x, cos, sin): cos = cos[None, None]; sin = sin[None, None] return x * cos + rotate_half(x) * sin class SwiGLU(nn.Module): def __init__(self, d, h): super().__init__() self.w1 = nn.Linear(d, h, bias=False); self.w3 = nn.Linear(d, h, bias=False); self.w2 = nn.Linear(h, d, bias=False) def forward(self, x): return self.w2(F.silu(self.w1(x)) * self.w3(x)) def _hdim(d): return ((int(8 / 3 * d) + 63) // 64) * 64 class ChunkedAttn(nn.Module): """Forced bottleneck: causal attention restricted to within non-overlapping chunks of size C.""" def __init__(self, d, nh, chunk): super().__init__() self.nh = nh; self.hd = d // nh; self.chunk = chunk self.qkv = nn.Linear(d, 3 * d, bias=False); self.o = nn.Linear(d, d, bias=False) def forward(self, x, cos, sin): B, T, D = x.shape qkv = self.qkv(x).view(B, T, 3, self.nh, self.hd).permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] q = apply_rope(q, cos, sin); k = apply_rope(k, cos, sin) idx = torch.arange(T, device=x.device) same = (idx[:, None] // self.chunk) == (idx[None, :] // self.chunk) causal = idx[:, None] >= idx[None, :] keep = same & causal mask = torch.zeros(T, T, device=x.device, dtype=q.dtype).masked_fill(~keep, float("-inf")) o = F.scaled_dot_product_attention(q, k, v, attn_mask=mask) return self.o(o.transpose(1, 2).reshape(B, T, D)) class GDNBlock(nn.Module): def __init__(self, d, idx, mlp_hidden, gdn_heads=4, gdn_hd=72): super().__init__() self.n1 = nn.RMSNorm(d) self.gdn = GatedDeltaNet(hidden_size=d, num_heads=gdn_heads, head_dim=gdn_hd, layer_idx=idx) self.n2 = nn.RMSNorm(d); self.mlp = SwiGLU(d, mlp_hidden) def forward(self, x): m = self.gdn(self.n1(x))[0] # fla returns (output, attn, cache) x = x + m return x + self.mlp(self.n2(x)) class AttnBlock(nn.Module): def __init__(self, d, nh, chunk, mlp_hidden): super().__init__() self.n1 = nn.RMSNorm(d); self.attn = ChunkedAttn(d, nh, chunk) self.n2 = nn.RMSNorm(d); self.mlp = SwiGLU(d, mlp_hidden) def forward(self, x, cos, sin): x = x + self.attn(self.n1(x), cos, sin) return x + self.mlp(self.n2(x)) class Bind2_0LM(nn.Module): def __init__(self, vocab, d=384, depth=12, nh=6, chunk=32, mlp_hidden=576, gdn_heads=4, gdn_hd=72): super().__init__() self.emb = nn.Embedding(vocab, d) self.kinds = ["attn" if (i + 1) % 4 == 0 else "gdn" for i in range(depth)] # 3:1 GDN:attn self.blocks = nn.ModuleList([ GDNBlock(d, i, mlp_hidden, gdn_heads, gdn_hd) if k == "gdn" else AttnBlock(d, nh, chunk, mlp_hidden) for i, k in enumerate(self.kinds)]) self.nf = nn.RMSNorm(d); self.head = nn.Linear(d, vocab, bias=False) self.head.weight = self.emb.weight self.d = d; self.nh = nh; self.chunk = chunk nl = depth for m in self.modules(): if isinstance(m, nn.Linear): nn.init.normal_(m.weight, std=0.02) nn.init.normal_(self.emb.weight, std=0.02) def forward(self, ids): cos, sin = build_rope(ids.shape[1], self.d // self.nh, ids.device) h = self.emb(ids) for blk, k in zip(self.blocks, self.kinds): h = blk(h) if k == "gdn" else blk(h, cos, sin) return self.head(self.nf(h))