""" Engram -- static n-gram memory (arXiv:2601.07372, DeepSeek engram_demo_v1.py), adapted for Veylon/Arya's standard transformer (hc_mult = 1, tiktoken). Fixes vs. the drafted port (all verified in test_engram.py): 1. Hashing ran on CPU/numpy INSIDE forward(): a GPU->CPU sync + torch.compile graph break on every Engram layer, every step (and fatal on XLA). Now the hash is pure int64 torch ops on the model device, compile-friendly. 2. Every Engram layer built its own NgramHashMapping (re-decoding the whole vocab) and each call hashed ALL layers, then kept one. Now ONE NgramHasher lives on GPT, hashes all layers once per forward. 3. pad_id=None crashed np.pad; pad_id=-1 silently indexed the LAST lookup entry. Left-context padding is now a dedicated reserved id (never collides with a real token class). 4. nn.RMSNorm (torch>=2.4 only, eps=finfo(fp16)=1e-3 default, no fp32 upcast) replaced by an fp32-upcast RMSNorm -- fp16/T4 stable like model.RMSNorm. 5. Gate dot-product computed in fp32 (sum over D in fp16 can overflow/lose precision), cast back after the sigmoid. 6. The compressed-vocab lookup is now a persistent buffer, so a checkpoint carries it and inference needs no tokenizer object. 7. sympy dependency dropped (trial-division primality is plenty for table sizes). 8. Special-id discovery no longer assumes TokenizerWrapper._special_ids is a set of ints (dict / missing attr both handled). 9. embed_per_ngram % n_heads != 0 silently produced a wrong-width table -> hard assert. Same for len(engram_vocab_size) < max_ngram-1. """ from __future__ import annotations import math import re import unicodedata from typing import List, Sequence import numpy as np import torch import torch.nn as nn import torch.nn.functional as F # ───────────────────────────────────────────────────────────────────────────── # Primes # ───────────────────────────────────────────────────────────────────────────── def _is_prime(n: int) -> bool: if n < 2: return False if n < 4: return True if n % 2 == 0 or n % 3 == 0: return False i = 5 while i * i <= n: if n % i == 0 or n % (i + 2) == 0: return False i += 6 return True def find_next_prime(start: int, seen_primes: set) -> int: candidate = start + 1 while True: if _is_prime(candidate) and candidate not in seen_primes: return candidate candidate += 1 # ───────────────────────────────────────────────────────────────────────────── # CompressedTokenizer -- build-time only (numpy). Output is a lookup table. # ───────────────────────────────────────────────────────────────────────────── class CompressedTokenizer: """Maps raw token IDs -> compressed IDs (NFKC / strip-accents / lowercase / whitespace-collapse of each token's surface string). Build-time only: the result (`lookup_table`) is handed to the model as a persistent buffer.""" _SENTINEL = "\uE000" _WS_RE = re.compile(r"[ \t\r\n]+") def __init__(self, tokenizer_wrapper): self.tokenizer = tokenizer_wrapper self._special_ids = self._collect_special_ids(tokenizer_wrapper) self.lookup_table, self.num_new_token = self._build_lookup_table() def __len__(self): return self.num_new_token @staticmethod def _collect_special_ids(tw) -> set: ids = set() raw = getattr(tw, "_special_ids", None) if raw is not None: if isinstance(raw, dict): for k, v in raw.items(): for cand in (k, v): if isinstance(cand, (int, np.integer)): ids.add(int(cand)) else: for v in raw: if isinstance(v, (int, np.integer)): ids.add(int(v)) vs = int(tw.vocab_size) for name in ("bos_id", "eos_id", "pad_id"): v = getattr(tw, name, None) if isinstance(v, (int, np.integer)) and 0 <= int(v) < vs: ids.add(int(v)) return ids @classmethod def _normalize(cls, text: str) -> str: text = unicodedata.normalize("NFKC", text) text = unicodedata.normalize("NFD", text) text = "".join(c for c in text if unicodedata.category(c) != "Mn") text = text.lower() text = cls._WS_RE.sub(" ", text) if text == " ": # lone space survives strip() text = cls._SENTINEL text = text.strip() return text.replace(cls._SENTINEL, " ") def _decode_id(self, tid: int) -> str: tw = self.tokenizer try: enc = getattr(tw, "enc", None) if enc is not None: return enc.decode([tid]) return tw.decode([tid]) except Exception: return "" def _build_lookup_table(self): key2new, new_tokens = {}, [] vocab_size = int(self.tokenizer.vocab_size) lookup = np.empty(vocab_size, dtype=np.int64) for tid in range(vocab_size): if tid in self._special_ids: key = f"__special_{tid}" else: text = self._decode_id(tid) if not text or "\ufffd" in text: key = f"__raw_{tid}" else: norm = self._normalize(text) key = norm if norm else f"__raw_{tid}" nid = key2new.get(key) if nid is None: nid = len(new_tokens) key2new[key] = nid new_tokens.append(key) lookup[tid] = nid return lookup, len(new_tokens) # ───────────────────────────────────────────────────────────────────────────── # NgramHasher -- static, deterministic, layer-specific, pure torch int64 # ───────────────────────────────────────────────────────────────────────────── class NgramHasher(nn.Module): """(B, T) raw ids -> {layer_id: (B, T, (max_ngram-1)*n_heads) hash ids}. No parameters, no gradients. Only `lookup` is persistent (state_dict); multipliers / moduli are re-derived from config, so they never drift. """ def __init__(self, layer_ids: Sequence[int], max_ngram: int, vocab_size_per_ngram: Sequence[int], n_heads: int, compressed_vocab: int, raw_vocab_size: int, seed: int, lookup: torch.Tensor | None = None): super().__init__() assert max_ngram >= 2, "engram_max_ngram must be >= 2" assert len(vocab_size_per_ngram) >= max_ngram - 1, ( f"engram_vocab_size needs {max_ngram - 1} entries (ngram 2..{max_ngram}), " f"got {len(vocab_size_per_ngram)}" ) assert compressed_vocab > 0, "engram_compressed_vocab must be set (>0)" self.layer_ids = tuple(int(l) for l in layer_ids) self.max_ngram = int(max_ngram) self.n_heads = int(n_heads) self.n_cols = (self.max_ngram - 1) * self.n_heads self.compressed_vocab = int(compressed_vocab) # Reserved id for left-context padding: distinct from every token class. self.pad_cid = self.compressed_vocab if lookup is None: lookup = torch.zeros(int(raw_vocab_size), dtype=torch.int64) else: lookup = lookup.detach().to(torch.int64).clone() assert lookup.numel() == int(raw_vocab_size), ( f"lookup has {lookup.numel()} entries, vocab_size={raw_vocab_size}") assert int(lookup.max()) < self.compressed_vocab self.register_buffer("lookup", lookup, persistent=True) # Multipliers: r*2+1, bounded so tok*mult can never overflow int64. # (+1 vocab slot for the reserved pad id.) max_long = int(np.iinfo(np.int64).max) M_max = max_long // (self.compressed_vocab + 1) half_bound = max(1, M_max // 2) PRIME_1 = 10007 mult = np.empty((len(self.layer_ids), self.max_ngram), dtype=np.int64) for li, layer_id in enumerate(self.layer_ids): g = np.random.default_rng(int(seed + PRIME_1 * int(layer_id))) r = g.integers(low=0, high=half_bound, size=(self.max_ngram,), dtype=np.int64) mult[li] = r * 2 + 1 self.register_buffer("mult", torch.from_numpy(mult), persistent=False) # Distinct prime moduli per (layer, ngram order, head); unique globally. seen: set = set() self.head_sizes: List[List[int]] = [] # per layer: flat list, len n_cols for _ in self.layer_ids: flat = [] for ngram in range(2, self.max_ngram + 1): search_start = int(vocab_size_per_ngram[ngram - 2]) - 1 for _h in range(self.n_heads): found = find_next_prime(search_start, seen) seen.add(found) flat.append(found) search_start = found self.head_sizes.append(flat) self.register_buffer( "mods", torch.tensor(self.head_sizes, dtype=torch.int64), persistent=False) def compress(self, idx: torch.Tensor) -> torch.Tensor: idx = idx.long() return torch.where(idx >= 0, self.lookup[idx.clamp_min(0)], idx) @torch.no_grad() def forward(self, idx: torch.Tensor): ids = self.compress(idx) # (B, T) T = ids.shape[1] shifted = [ids] + [ F.pad(ids, (k, 0), value=self.pad_cid)[:, :T] for k in range(1, self.max_ngram) ] out = {} for li, layer_id in enumerate(self.layer_ids): m = self.mult[li] cols = [] for n in range(2, self.max_ngram + 1): mix = shifted[0] * m[0] for k in range(1, n): mix = torch.bitwise_xor(mix, shifted[k] * m[k]) lo = (n - 2) * self.n_heads mods = self.mods[li, lo:lo + self.n_heads] # (n_heads,) cols.append(mix.unsqueeze(-1) % mods) # (B, T, n_heads) out[layer_id] = torch.cat(cols, dim=-1) # (B, T, n_cols) return out # ───────────────────────────────────────────────────────────────────────────── # Modules # ───────────────────────────────────────────────────────────────────────────── class _RMSNorm(nn.Module): """fp32-upcast RMSNorm (matches model.RMSNorm's fp16-safe fallback path).""" def __init__(self, dim: int, eps: float = 1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: xf = x.float() rms = torch.sqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps) return (xf / rms).to(x.dtype) * self.weight.to(x.dtype) class ShortConv(nn.Module): """(B, T, D) -> (B, T, D). Depthwise causal conv, dilation = max_ngram.""" def __init__(self, hidden_size: int, kernel_size: int = 4, dilation: int = 1, activation: bool = True): super().__init__() self.activation = activation self.conv = nn.Conv1d( hidden_size, hidden_size, kernel_size=kernel_size, groups=hidden_size, bias=False, padding=(kernel_size - 1) * dilation, dilation=dilation, ) self.norm = _RMSNorm(hidden_size) self.act_fn = nn.SiLU() if activation else None def forward(self, x: torch.Tensor) -> torch.Tensor: T = x.shape[1] y = self.conv(self.norm(x).transpose(1, 2))[..., :T] # causal crop if self.activation: y = self.act_fn(y) return y.transpose(1, 2) class MultiHeadEmbedding(nn.Module): def __init__(self, list_of_N: List[int], D: int): super().__init__() offsets = [0] for n in list_of_N[:-1]: offsets.append(offsets[-1] + n) self.register_buffer("offsets", torch.tensor(offsets, dtype=torch.long), persistent=False) self.embedding = nn.Embedding(sum(list_of_N), D) self.reset_parameters() def reset_parameters(self): nn.init.normal_(self.embedding.weight, std=0.01) def forward(self, hash_ids: torch.Tensor) -> torch.Tensor: return self.embedding(hash_ids + self.offsets) class Engram(nn.Module): """Static n-gram memory with conditional gated injection (hc_mult = 1). forward() returns the DELTA to add to the residual stream: x = x + engram(x, hash_ids) """ def __init__(self, hidden_size: int, head_sizes: List[int], max_ngram: int, embed_per_ngram: int, n_heads: int, kernel_size: int): super().__init__() assert embed_per_ngram % n_heads == 0, ( f"engram_embed_per_ngram={embed_per_ngram} must be divisible by " f"engram_n_heads={n_heads}") self.hidden_size = hidden_size d_head = embed_per_ngram // n_heads engram_hidden = (max_ngram - 1) * embed_per_ngram self.multi_head_embedding = MultiHeadEmbedding(list(head_sizes), d_head) self.value_proj = nn.Linear(engram_hidden, hidden_size, bias=False) self.key_proj = nn.Linear(engram_hidden, hidden_size, bias=False) self.norm1 = _RMSNorm(hidden_size) self.norm2 = _RMSNorm(hidden_size) self.short_conv = ShortConv(hidden_size, kernel_size=kernel_size, dilation=max_ngram) self.reset_parameters() def reset_parameters(self): self.multi_head_embedding.reset_parameters() nn.init.normal_(self.value_proj.weight, std=0.02) nn.init.normal_(self.key_proj.weight, std=0.02) def forward(self, hidden_states: torch.Tensor, hash_ids: torch.Tensor): emb = self.multi_head_embedding(hash_ids).flatten(start_dim=-2) # (B,T,engram_hidden) key = self.key_proj(emb) nk = self.norm1(key).float() nq = self.norm2(hidden_states).float() gate = (nk * nq).sum(dim=-1) / math.sqrt(self.hidden_size) # fp32 gate = gate.abs().clamp_min(1e-6).sqrt() * gate.sign() gate = gate.sigmoid().unsqueeze(-1).to(hidden_states.dtype) # (B,T,1) value = gate * self.value_proj(emb) return value + self.short_conv(value)