# ========================================================================= # splat_attention.py — standalone Splat Attention (aleph-addressed, # softmax-free attention through a shared blackboard) # ========================================================================= # From the AlephLM-0 / aleph-splat research line (AbstractPhil). Single # file, no dependencies beyond torch. Experimental — the measured record, # including failures, is summarized below so you know what you're holding. # # THE MECHANISM # Every head is a tiny frozen "aleph" codebook: K unit anchor # directions read through a closed-form SIGNED address # u_k = cos(x, a_k)/tau, w_k = sinh(u_k) / sum_j cosh(u_j) # (a reconstructive read — no argmax, no top-k, no softmax selection; # weights are signed, so an anchor can contribute negatively). # Attention is a write/read through the codebook cells: # write: cells = sum_j wg_j (x) v_j (per head) # read: out_i = wg_i @ cells / sum|wg_i| (weighted average) # Affinity between tokens is address AGREEMENT through the K-cell # bottleneck: O(L*M*K) per layer — LINEAR in sequence length. # # MEASURED (A40/4090, fp16 autocast, fwd+bwd, us/token): # vs nn.MultiheadAttention(8 heads, d=512): ~3-6x SLOWER at L<=256, # parity ~L=2048, ~2x FASTER at L=8192 (flat cost vs quadratic). # torch.compile (inductor, Linux) gives a further 3-4x on this module. # Associative recall through splat-sharded heads at equal total cells: # top-1 .9995 @ 2k context / .934 @ 8k where one monolithic codebook # reads .042/.0015 — partition+locality rescues superposition. # # THE FAILURE YOU MUST KNOW ABOUT (measured, structural): with LOCAL # positional membership only (gaussian windows) at short L, supports # shrink to a few tokens, attention degenerates to a local blur, # cross-position transport dies, and a cls-pooled encoder COLLAPSES # (representation erank ~5, retrieval at noise). Two repairs, both # included here: # rotary=True position enters as a RoPE rotation of the ADDRESS # QUERY against the frozen codebook: relative # position R(i-j) appears in every affinity, heads # stay GLOBAL, transport exists at every L # (probe: cross-position recall .548 at L=128 where # local-only gave ~0; retrieval decays gracefully # with query/write offset: .64 -> .31 over 0 -> 64). # global_frac>0 reserve a fraction of heads with uniform # membership alongside the local windows. # Defaults below are the SAFE configuration (rotary=True). # # DESIGN CARD (from the measurement battery): # K=4-8 per head (small codebooks saturate their sign-code space at # ~0.5 bits/half-axis; big ones waste it) | M scales with data rank # (rich data pays monotonically to M=2048) | frames born random and # independent — constructed rotations buy nothing; differentiation is # maintained by training pressure itself | overlap sigma/spacing in # [1,2] when using local windows | composition by budget, never by # softmax over heads (comparative composition measurably loses ~.10) | # storage capacity scales with TOTAL cells regardless of partition — # address capacity and memory capacity are different resources. # # USAGE # from splat_attention import SplatAttention # attn = SplatAttention(d_model=512, M=64, K=8, rotary=True) # y = attn(x) # x: (B, L, d), y: (B, L, d) # y = attn(x, key_padding_mask=kpm) # kpm: (B, L) True = pad # python splat_attention.py # runs the demo + a small speed bench # # TRAINING NOTE (measured 2026-08-06/07, 500k-caption encoder screens): # the all-frozen configuration COLLAPSES when trained inside a trunk # (representation erank ~5) — parameter-free routing deforms token # states into address basins. The configuration that CONVERGES: # SplatAttention(..., rotary=True, addr_proj=True, # train_codebooks=True) # (routing-owned parameters: a learned address frame + living # codebooks). It reaches ~90% of a standard block's training-gauge # performance at matched small budget and was still climbing at # cutoff — functional, slower to organize, endpoint parity unproven. # Frozen-everything remains fine for INFERENCE-style play and the # static properties above. # # Status: research prototype. Trained-at-scale results pending; treat # every number above as what it is — a measurement on the stated probe. # ========================================================================= import math import torch import torch.nn as nn import torch.nn.functional as F def rope_rotate(x, pos, base=10000.0): """RoPE rotation of the address query. pos: (L,) float positions.""" D = x.shape[-1] half = D // 2 freqs = base ** (-torch.arange(half, device=x.device, dtype=torch.float32) / half) ang = pos.unsqueeze(-1) * freqs c, sn = torch.cos(ang), torch.sin(ang) x1, x2 = x[..., :half], x[..., half:] return torch.cat([x1 * c - x2 * sn, x1 * sn + x2 * c], dim=-1) class SplatAttention(nn.Module): """Aleph-addressed attention. See module docstring. Args: d_model: model width M: number of heads (tiny codebooks) K: anchors per head (4-8 recommended) tau: address temperature (0.1) dropout: output dropout rotary: position via RoPE on the address query; heads global (RECOMMENDED — see the failure note) global_frac: fraction of heads with uniform membership when rotary=False (transport insurance for windows) overlap: window sigma as a multiple of spacing (local mode) sigma_floor: minimum window sigma in tokens (local mode) head_gates: learnable per-head attenuation, born at identity train_centers: learnable window centers/widths (local mode) — only functions on a transport-capable geometry train_codebooks: unfreeze the anchor frames mchunk: heads per computation chunk (memory control) checkpoint_chunks: recompute chunks in backward (training-time memory saver; needs torch.utils.checkpoint) """ def __init__(self, d_model, M=64, K=8, tau=0.1, dropout=0.0, rotary=True, global_frac=0.0, overlap=1.5, sigma_floor=0.0, head_gates=False, train_centers=False, train_codebooks=False, addr_proj=False, mchunk=16, checkpoint_chunks=False): super().__init__() self.M, self.K, self.tau = M, K, tau self.rotary = rotary self.global_frac = global_frac self.overlap, self.sigma_floor = overlap, sigma_floor self.mchunk = mchunk self.checkpoint_chunks = checkpoint_chunks book = F.normalize(torch.randn(M * K, d_model), dim=-1) if train_codebooks: self.codebook = nn.Parameter(book) else: self.register_buffer("codebook", book) self.w_v = nn.Linear(d_model, d_model) self.w_o = nn.Linear(d_model, d_model) if addr_proj: # learned address frame, born at identity — routing-owned # parameters (see TRAINING note in the module docstring) self.w_a = nn.Linear(d_model, d_model, bias=False) nn.init.eye_(self.w_a.weight) else: self.w_a = None self.drop = nn.Dropout(dropout) if train_centers: self.center_off = nn.Parameter(torch.zeros(M)) self.log_sig = nn.Parameter(torch.zeros(M)) else: self.center_off = self.log_sig = None if head_gates: self.head_gate = nn.Parameter(torch.zeros(M)) else: self.head_gate = None def _book(self): return (F.normalize(self.codebook, dim=-1) if isinstance(self.codebook, nn.Parameter) else self.codebook) def _membership(self, L, device, dtype): if self.rotary: return torch.ones(self.M, L, device=device, dtype=dtype) pos = torch.arange(L, device=device, dtype=torch.float32) frac = torch.linspace(0, 1, self.M, device=device) if self.center_off is not None: frac = (frac + self.center_off.float()).clamp(0, 1) sig = (self.overlap * max(L / self.M, 1.0) * torch.exp(self.log_sig.float()).unsqueeze(1)) sig = sig.clamp(min=max(self.sigma_floor, 1e-6)) else: sig = max(self.overlap * max(L / self.M, 1.0), self.sigma_floor, 1e-6) centers = frac * (L - 1) g = torch.exp(-0.5 * ((pos.unsqueeze(0) - centers.unsqueeze(1)) / sig) ** 2) g = g / g.sum(dim=0, keepdim=True).clamp(min=1e-9) n_glob = int(round(self.M * self.global_frac)) if n_glob > 0: g[:n_glob] = 1.0 return g.to(dtype) def _chunk(self, xn, v, g_c, live, c0, Mc): sl = self._book()[c0 * self.K:(c0 + Mc) * self.K] u = (xn @ sl.T).view(*xn.shape[:2], Mc, self.K) / self.tau m = u.abs().amax(dim=-1, keepdim=True) ep, en = torch.exp(u - m), torch.exp(-u - m) w = (ep - en) / (ep + en).sum(dim=-1, keepdim=True) wg = w * g_c.T.unsqueeze(0).unsqueeze(-1) if self.head_gate is not None: gam = 2 * torch.sigmoid(self.head_gate[c0:c0 + Mc]) wg = wg * gam.view(1, 1, -1, 1) wg = (wg * live.unsqueeze(-1).unsqueeze(-1)).to(v.dtype) cells = torch.einsum("blmk,bld->bmkd", wg, v) # write part = torch.einsum("blmk,bmkd->bld", wg, cells) # read den = wg.abs().sum(dim=(2, 3)) return part, den def forward(self, x, key_padding_mask=None): B, L, d = x.shape xa = self.w_a(x) if self.w_a is not None else x xn = F.normalize(xa, dim=-1) if self.rotary: pos = torch.arange(L, device=x.device, dtype=torch.float32) xn = F.normalize(rope_rotate(xn.float(), pos), dim=-1).to(xn.dtype) g = self._membership(L, x.device, x.dtype) live = ((~key_padding_mask).to(x.dtype) if key_padding_mask is not None else torch.ones(B, L, device=x.device, dtype=x.dtype)) v = self.w_v(x) out = torch.zeros_like(x) den = torch.zeros(B, L, device=x.device, dtype=x.dtype) for c0 in range(0, self.M, self.mchunk): Mc = min(self.mchunk, self.M - c0) if (self.checkpoint_chunks and self.training and torch.is_grad_enabled()): import torch.utils.checkpoint as _ck part, dpart = _ck.checkpoint( self._chunk, xn, v, g[c0:c0 + Mc], live, c0, Mc, use_reentrant=False, preserve_rng_state=False) else: part, dpart = self._chunk(xn, v, g[c0:c0 + Mc], live, c0, Mc) out = out.add_(part) den = den.add_(dpart) out = out / den.unsqueeze(-1).clamp_min(1e-9) return self.drop(self.w_o(out)) def _demo(): torch.manual_seed(0) dev = "cuda" if torch.cuda.is_available() else "cpu" print(f"SplatAttention demo (device={dev})") attn = SplatAttention(d_model=256, M=32, K=8, rotary=True).to(dev) x = torch.randn(2, 128, 256, device=dev) y = attn(x) print(f" forward: {tuple(x.shape)} -> {tuple(y.shape)}") y.sum().backward() print(f" backward OK; trainable params: " f"{sum(p.numel() for p in attn.parameters() if p.requires_grad):,}" f" (+ frozen codebook {attn._book().numel():,})") if dev == "cuda": import time mha = nn.MultiheadAttention(256, 8, batch_first=True).to(dev) for L, B in [(128, 32), (2048, 2)]: xx = torch.randn(B, L, 256, device=dev, requires_grad=True) def t(fn, n=10): for _ in range(3): fn() torch.cuda.synchronize() t0 = time.time() for _ in range(n): fn() torch.cuda.synchronize() return (time.time() - t0) / n / (B * L) * 1e6 ts = t(lambda: attn(xx).sum().backward()) tm = t(lambda: mha(xx, xx, xx, need_weights=False)[0].sum().backward()) print(f" L={L}: splat {ts:.2f} vs MHA {tm:.2f} us/token") print(" (try rotary=False, global_frac=0.25 for windowed mode, " "head_gates=True for learnable attenuation)") if __name__ == "__main__": _demo()