ArushBuilds commited on
Commit
009509d
Β·
1 Parent(s): 61b358c

Upload ngram.py

Browse files
Files changed (1) hide show
  1. ngram.py +350 -0
ngram.py ADDED
@@ -0,0 +1,350 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Engram -- static n-gram memory (arXiv:2601.07372, DeepSeek engram_demo_v1.py),
3
+ adapted for Veylon/Arya's standard transformer (hc_mult = 1, tiktoken).
4
+
5
+ Fixes vs. the drafted port (all verified in test_engram.py):
6
+
7
+ 1. Hashing ran on CPU/numpy INSIDE forward(): a GPU->CPU sync + torch.compile
8
+ graph break on every Engram layer, every step (and fatal on XLA). Now the
9
+ hash is pure int64 torch ops on the model device, compile-friendly.
10
+ 2. Every Engram layer built its own NgramHashMapping (re-decoding the whole
11
+ vocab) and each call hashed ALL layers, then kept one. Now ONE NgramHasher
12
+ lives on GPT, hashes all layers once per forward.
13
+ 3. pad_id=None crashed np.pad; pad_id=-1 silently indexed the LAST lookup
14
+ entry. Left-context padding is now a dedicated reserved id (never
15
+ collides with a real token class).
16
+ 4. nn.RMSNorm (torch>=2.4 only, eps=finfo(fp16)=1e-3 default, no fp32 upcast)
17
+ replaced by an fp32-upcast RMSNorm -- fp16/T4 stable like model.RMSNorm.
18
+ 5. Gate dot-product computed in fp32 (sum over D in fp16 can overflow/lose
19
+ precision), cast back after the sigmoid.
20
+ 6. The compressed-vocab lookup is now a persistent buffer, so a checkpoint
21
+ carries it and inference needs no tokenizer object.
22
+ 7. sympy dependency dropped (trial-division primality is plenty for table sizes).
23
+ 8. Special-id discovery no longer assumes TokenizerWrapper._special_ids is a
24
+ set of ints (dict / missing attr both handled).
25
+ 9. embed_per_ngram % n_heads != 0 silently produced a wrong-width table ->
26
+ hard assert. Same for len(engram_vocab_size) < max_ngram-1.
27
+ """
28
+
29
+ from __future__ import annotations
30
+
31
+ import math
32
+ import re
33
+ import unicodedata
34
+ from typing import List, Sequence
35
+
36
+ import numpy as np
37
+ import torch
38
+ import torch.nn as nn
39
+ import torch.nn.functional as F
40
+
41
+
42
+ # ─────────────────────────────────────────────────────────────────────────────
43
+ # Primes
44
+ # ─────────────────────────────────────────────────────────────────────────────
45
+
46
+ def _is_prime(n: int) -> bool:
47
+ if n < 2:
48
+ return False
49
+ if n < 4:
50
+ return True
51
+ if n % 2 == 0 or n % 3 == 0:
52
+ return False
53
+ i = 5
54
+ while i * i <= n:
55
+ if n % i == 0 or n % (i + 2) == 0:
56
+ return False
57
+ i += 6
58
+ return True
59
+
60
+
61
+ def find_next_prime(start: int, seen_primes: set) -> int:
62
+ candidate = start + 1
63
+ while True:
64
+ if _is_prime(candidate) and candidate not in seen_primes:
65
+ return candidate
66
+ candidate += 1
67
+
68
+
69
+ # ─────────────────────────────────────────────────────────────────────────────
70
+ # CompressedTokenizer -- build-time only (numpy). Output is a lookup table.
71
+ # ─────────────────────────────────────────────────────────────────────────────
72
+
73
+ class CompressedTokenizer:
74
+ """Maps raw token IDs -> compressed IDs (NFKC / strip-accents / lowercase /
75
+ whitespace-collapse of each token's surface string). Build-time only: the
76
+ result (`lookup_table`) is handed to the model as a persistent buffer."""
77
+
78
+ _SENTINEL = "\uE000"
79
+ _WS_RE = re.compile(r"[ \t\r\n]+")
80
+
81
+ def __init__(self, tokenizer_wrapper):
82
+ self.tokenizer = tokenizer_wrapper
83
+ self._special_ids = self._collect_special_ids(tokenizer_wrapper)
84
+ self.lookup_table, self.num_new_token = self._build_lookup_table()
85
+
86
+ def __len__(self):
87
+ return self.num_new_token
88
+
89
+ @staticmethod
90
+ def _collect_special_ids(tw) -> set:
91
+ ids = set()
92
+ raw = getattr(tw, "_special_ids", None)
93
+ if raw is not None:
94
+ if isinstance(raw, dict):
95
+ for k, v in raw.items():
96
+ for cand in (k, v):
97
+ if isinstance(cand, (int, np.integer)):
98
+ ids.add(int(cand))
99
+ else:
100
+ for v in raw:
101
+ if isinstance(v, (int, np.integer)):
102
+ ids.add(int(v))
103
+ vs = int(tw.vocab_size)
104
+ for name in ("bos_id", "eos_id", "pad_id"):
105
+ v = getattr(tw, name, None)
106
+ if isinstance(v, (int, np.integer)) and 0 <= int(v) < vs:
107
+ ids.add(int(v))
108
+ return ids
109
+
110
+ @classmethod
111
+ def _normalize(cls, text: str) -> str:
112
+ text = unicodedata.normalize("NFKC", text)
113
+ text = unicodedata.normalize("NFD", text)
114
+ text = "".join(c for c in text if unicodedata.category(c) != "Mn")
115
+ text = text.lower()
116
+ text = cls._WS_RE.sub(" ", text)
117
+ if text == " ": # lone space survives strip()
118
+ text = cls._SENTINEL
119
+ text = text.strip()
120
+ return text.replace(cls._SENTINEL, " ")
121
+
122
+ def _decode_id(self, tid: int) -> str:
123
+ tw = self.tokenizer
124
+ try:
125
+ enc = getattr(tw, "enc", None)
126
+ if enc is not None:
127
+ return enc.decode([tid])
128
+ return tw.decode([tid])
129
+ except Exception:
130
+ return ""
131
+
132
+ def _build_lookup_table(self):
133
+ key2new, new_tokens = {}, []
134
+ vocab_size = int(self.tokenizer.vocab_size)
135
+ lookup = np.empty(vocab_size, dtype=np.int64)
136
+ for tid in range(vocab_size):
137
+ if tid in self._special_ids:
138
+ key = f"__special_{tid}"
139
+ else:
140
+ text = self._decode_id(tid)
141
+ if not text or "\ufffd" in text:
142
+ key = f"__raw_{tid}"
143
+ else:
144
+ norm = self._normalize(text)
145
+ key = norm if norm else f"__raw_{tid}"
146
+ nid = key2new.get(key)
147
+ if nid is None:
148
+ nid = len(new_tokens)
149
+ key2new[key] = nid
150
+ new_tokens.append(key)
151
+ lookup[tid] = nid
152
+ return lookup, len(new_tokens)
153
+
154
+
155
+ # ─────────────────────────────────────────────────────────────────────────────
156
+ # NgramHasher -- static, deterministic, layer-specific, pure torch int64
157
+ # ─────────────────────────────────────────────────────────────────────────────
158
+
159
+ class NgramHasher(nn.Module):
160
+ """(B, T) raw ids -> {layer_id: (B, T, (max_ngram-1)*n_heads) hash ids}.
161
+
162
+ No parameters, no gradients. Only `lookup` is persistent (state_dict);
163
+ multipliers / moduli are re-derived from config, so they never drift.
164
+ """
165
+
166
+ def __init__(self, layer_ids: Sequence[int], max_ngram: int,
167
+ vocab_size_per_ngram: Sequence[int], n_heads: int,
168
+ compressed_vocab: int, raw_vocab_size: int, seed: int,
169
+ lookup: torch.Tensor | None = None):
170
+ super().__init__()
171
+ assert max_ngram >= 2, "engram_max_ngram must be >= 2"
172
+ assert len(vocab_size_per_ngram) >= max_ngram - 1, (
173
+ f"engram_vocab_size needs {max_ngram - 1} entries (ngram 2..{max_ngram}), "
174
+ f"got {len(vocab_size_per_ngram)}"
175
+ )
176
+ assert compressed_vocab > 0, "engram_compressed_vocab must be set (>0)"
177
+
178
+ self.layer_ids = tuple(int(l) for l in layer_ids)
179
+ self.max_ngram = int(max_ngram)
180
+ self.n_heads = int(n_heads)
181
+ self.n_cols = (self.max_ngram - 1) * self.n_heads
182
+ self.compressed_vocab = int(compressed_vocab)
183
+ # Reserved id for left-context padding: distinct from every token class.
184
+ self.pad_cid = self.compressed_vocab
185
+
186
+ if lookup is None:
187
+ lookup = torch.zeros(int(raw_vocab_size), dtype=torch.int64)
188
+ else:
189
+ lookup = lookup.detach().to(torch.int64).clone()
190
+ assert lookup.numel() == int(raw_vocab_size), (
191
+ f"lookup has {lookup.numel()} entries, vocab_size={raw_vocab_size}")
192
+ assert int(lookup.max()) < self.compressed_vocab
193
+ self.register_buffer("lookup", lookup, persistent=True)
194
+
195
+ # Multipliers: r*2+1, bounded so tok*mult can never overflow int64.
196
+ # (+1 vocab slot for the reserved pad id.)
197
+ max_long = int(np.iinfo(np.int64).max)
198
+ M_max = max_long // (self.compressed_vocab + 1)
199
+ half_bound = max(1, M_max // 2)
200
+ PRIME_1 = 10007
201
+ mult = np.empty((len(self.layer_ids), self.max_ngram), dtype=np.int64)
202
+ for li, layer_id in enumerate(self.layer_ids):
203
+ g = np.random.default_rng(int(seed + PRIME_1 * int(layer_id)))
204
+ r = g.integers(low=0, high=half_bound, size=(self.max_ngram,), dtype=np.int64)
205
+ mult[li] = r * 2 + 1
206
+ self.register_buffer("mult", torch.from_numpy(mult), persistent=False)
207
+
208
+ # Distinct prime moduli per (layer, ngram order, head); unique globally.
209
+ seen: set = set()
210
+ self.head_sizes: List[List[int]] = [] # per layer: flat list, len n_cols
211
+ for _ in self.layer_ids:
212
+ flat = []
213
+ for ngram in range(2, self.max_ngram + 1):
214
+ search_start = int(vocab_size_per_ngram[ngram - 2]) - 1
215
+ for _h in range(self.n_heads):
216
+ found = find_next_prime(search_start, seen)
217
+ seen.add(found)
218
+ flat.append(found)
219
+ search_start = found
220
+ self.head_sizes.append(flat)
221
+ self.register_buffer(
222
+ "mods", torch.tensor(self.head_sizes, dtype=torch.int64), persistent=False)
223
+
224
+ def compress(self, idx: torch.Tensor) -> torch.Tensor:
225
+ idx = idx.long()
226
+ return torch.where(idx >= 0, self.lookup[idx.clamp_min(0)], idx)
227
+
228
+ @torch.no_grad()
229
+ def forward(self, idx: torch.Tensor):
230
+ ids = self.compress(idx) # (B, T)
231
+ T = ids.shape[1]
232
+ shifted = [ids] + [
233
+ F.pad(ids, (k, 0), value=self.pad_cid)[:, :T]
234
+ for k in range(1, self.max_ngram)
235
+ ]
236
+ out = {}
237
+ for li, layer_id in enumerate(self.layer_ids):
238
+ m = self.mult[li]
239
+ cols = []
240
+ for n in range(2, self.max_ngram + 1):
241
+ mix = shifted[0] * m[0]
242
+ for k in range(1, n):
243
+ mix = torch.bitwise_xor(mix, shifted[k] * m[k])
244
+ lo = (n - 2) * self.n_heads
245
+ mods = self.mods[li, lo:lo + self.n_heads] # (n_heads,)
246
+ cols.append(mix.unsqueeze(-1) % mods) # (B, T, n_heads)
247
+ out[layer_id] = torch.cat(cols, dim=-1) # (B, T, n_cols)
248
+ return out
249
+
250
+
251
+ # ─────────────────────────────────────────────────────────────────────────────
252
+ # Modules
253
+ # ─────────────────────────────────────────────────────────────────────────────
254
+
255
+ class _RMSNorm(nn.Module):
256
+ """fp32-upcast RMSNorm (matches model.RMSNorm's fp16-safe fallback path)."""
257
+
258
+ def __init__(self, dim: int, eps: float = 1e-5):
259
+ super().__init__()
260
+ self.eps = eps
261
+ self.weight = nn.Parameter(torch.ones(dim))
262
+
263
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
264
+ xf = x.float()
265
+ rms = torch.sqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps)
266
+ return (xf / rms).to(x.dtype) * self.weight.to(x.dtype)
267
+
268
+
269
+ class ShortConv(nn.Module):
270
+ """(B, T, D) -> (B, T, D). Depthwise causal conv, dilation = max_ngram."""
271
+
272
+ def __init__(self, hidden_size: int, kernel_size: int = 4,
273
+ dilation: int = 1, activation: bool = True):
274
+ super().__init__()
275
+ self.activation = activation
276
+ self.conv = nn.Conv1d(
277
+ hidden_size, hidden_size, kernel_size=kernel_size,
278
+ groups=hidden_size, bias=False,
279
+ padding=(kernel_size - 1) * dilation, dilation=dilation,
280
+ )
281
+ self.norm = _RMSNorm(hidden_size)
282
+ self.act_fn = nn.SiLU() if activation else None
283
+
284
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
285
+ T = x.shape[1]
286
+ y = self.conv(self.norm(x).transpose(1, 2))[..., :T] # causal crop
287
+ if self.activation:
288
+ y = self.act_fn(y)
289
+ return y.transpose(1, 2)
290
+
291
+
292
+ class MultiHeadEmbedding(nn.Module):
293
+ def __init__(self, list_of_N: List[int], D: int):
294
+ super().__init__()
295
+ offsets = [0]
296
+ for n in list_of_N[:-1]:
297
+ offsets.append(offsets[-1] + n)
298
+ self.register_buffer("offsets", torch.tensor(offsets, dtype=torch.long),
299
+ persistent=False)
300
+ self.embedding = nn.Embedding(sum(list_of_N), D)
301
+ self.reset_parameters()
302
+
303
+ def reset_parameters(self):
304
+ nn.init.normal_(self.embedding.weight, std=0.01)
305
+
306
+ def forward(self, hash_ids: torch.Tensor) -> torch.Tensor:
307
+ return self.embedding(hash_ids + self.offsets)
308
+
309
+
310
+ class Engram(nn.Module):
311
+ """Static n-gram memory with conditional gated injection (hc_mult = 1).
312
+
313
+ forward() returns the DELTA to add to the residual stream:
314
+ x = x + engram(x, hash_ids)
315
+ """
316
+
317
+ def __init__(self, hidden_size: int, head_sizes: List[int], max_ngram: int,
318
+ embed_per_ngram: int, n_heads: int, kernel_size: int):
319
+ super().__init__()
320
+ assert embed_per_ngram % n_heads == 0, (
321
+ f"engram_embed_per_ngram={embed_per_ngram} must be divisible by "
322
+ f"engram_n_heads={n_heads}")
323
+ self.hidden_size = hidden_size
324
+ d_head = embed_per_ngram // n_heads
325
+ engram_hidden = (max_ngram - 1) * embed_per_ngram
326
+
327
+ self.multi_head_embedding = MultiHeadEmbedding(list(head_sizes), d_head)
328
+ self.value_proj = nn.Linear(engram_hidden, hidden_size, bias=False)
329
+ self.key_proj = nn.Linear(engram_hidden, hidden_size, bias=False)
330
+ self.norm1 = _RMSNorm(hidden_size)
331
+ self.norm2 = _RMSNorm(hidden_size)
332
+ self.short_conv = ShortConv(hidden_size, kernel_size=kernel_size,
333
+ dilation=max_ngram)
334
+ self.reset_parameters()
335
+
336
+ def reset_parameters(self):
337
+ self.multi_head_embedding.reset_parameters()
338
+ nn.init.normal_(self.value_proj.weight, std=0.02)
339
+ nn.init.normal_(self.key_proj.weight, std=0.02)
340
+
341
+ def forward(self, hidden_states: torch.Tensor, hash_ids: torch.Tensor):
342
+ emb = self.multi_head_embedding(hash_ids).flatten(start_dim=-2) # (B,T,engram_hidden)
343
+ key = self.key_proj(emb)
344
+ nk = self.norm1(key).float()
345
+ nq = self.norm2(hidden_states).float()
346
+ gate = (nk * nq).sum(dim=-1) / math.sqrt(self.hidden_size) # fp32
347
+ gate = gate.abs().clamp_min(1e-6).sqrt() * gate.sign()
348
+ gate = gate.sigmoid().unsqueeze(-1).to(hidden_states.dtype) # (B,T,1)
349
+ value = gate * self.value_proj(emb)
350
+ return value + self.short_conv(value)