jon1012's picture
Add files using upload-large-folder tool
bb70bd3 verified
Raw
History Blame
8.14 kB
from dataclasses import dataclass
import numpy as np
import torch
from sympy import isprime
from torch import nn
def find_next_prime(start: int, seen_primes: set[int]) -> int:
"""The smallest prime above `start` that has not been handed out yet."""
candidate = start + 1
while not isprime(candidate) or candidate in seen_primes:
candidate += 1
return candidate
def build_compressed_token_map(tokenizer) -> tuple[list[int], int]:
"""Map every token id onto a smaller id space where tokens that normalize alike collapse together.
N-grams are hashed over these compressed ids, so " The", "the" and "THE" all hash the same way.
Returns the lookup plus the size of the compressed vocab -- and that size matters beyond bounds
checking, because every hash multiplier is derived from it.
"""
from tokenizers import Regex, normalizers
# a private-use char, so a token that is exactly one space survives Strip() instead of
# collapsing to the empty string and merging with unrelated tokens
sentinel = "\ue000"
normalizer = normalizers.Sequence(
[
normalizers.NFKC(),
normalizers.NFD(),
normalizers.StripAccents(),
normalizers.Lowercase(),
normalizers.Replace(Regex(r"[ \t\r\n]+"), " "),
normalizers.Replace(Regex(r"^ $"), sentinel),
normalizers.Strip(),
normalizers.Replace(sentinel, " "),
]
)
# the raw Rust tokenizer, matching what training decodes with (no clean_up_tokenization_spaces)
backend = tokenizer.backend_tokenizer
key_to_new: dict[str, int] = {}
lookup = [0] * len(tokenizer)
for token_id in range(len(tokenizer)):
text = backend.decode([token_id], skip_special_tokens=False)
if "\ufffd" in text:
# a partial UTF-8 byte token: nothing to normalize, so key it by its raw form
key = backend.id_to_token(token_id)
else:
normalized = normalizer.normalize_str(text)
key = normalized if normalized else text
new_id = key_to_new.get(key)
if new_id is None:
new_id = len(key_to_new)
key_to_new[key] = new_id
lookup[token_id] = new_id
return lookup, len(key_to_new)
def compute_hash_multipliers(
layer_ids: tuple[int, ...], max_ngram_size: int, tokenizer_vocab_size: int
) -> torch.Tensor:
"""One multiplier per (layer, lookback), from a per-layer RNG so layers hash differently.
Kept odd, and bounded so that `token_id * multiplier` cannot overflow int64.
"""
max_long = np.iinfo(np.int64).max
multiplier_bound = max(1, (max_long // tokenizer_vocab_size) // 2)
rows = []
for layer_id in layer_ids:
generator = np.random.default_rng(10007 * layer_id)
values = generator.integers(
low=0,
high=multiplier_bound,
size=(max_ngram_size,),
dtype=np.int64,
)
rows.append(torch.tensor(values * 2 + 1))
return torch.stack(rows)
@dataclass(frozen=True)
class EngramLayout:
"""Bucket layout of the n-gram hash tables.
A position is hashed as `max_ngram_size - 1` n-grams (2-gram .. max_ngram_size-gram), each split
over `n_heads` heads. Every (n-gram size, head) pair owns its own prime-sized bucket range in the
layer's table; the primes are drawn in order and never reused, which keeps the ranges disjoint.
"""
max_ngram_size: int
layer_ids: tuple[int, ...]
num_embeddings: tuple[int, ...] # table rows, per engram layer
primes: tuple[tuple[tuple[int, ...], ...], ...] # [layer][n-gram size][head] bucket modulus
n_heads: int
head_dim: int
@classmethod
def from_args(cls, args) -> "EngramLayout | None":
layer_ids = tuple(args.engram_layer_ids)
if not layer_ids:
return None
max_ngram_size, n_heads = args.engram_max_ngram_size, args.engram_n_heads
primes, seen = [], set()
for _ in layer_ids:
per_ngram = []
for _ in range(max_ngram_size - 1):
sizes, current = [], args.engram_vocab_size - 1
for _ in range(n_heads):
current = find_next_prime(current, seen)
seen.add(current)
sizes.append(current)
per_ngram.append(tuple(sizes))
primes.append(tuple(per_ngram))
return cls(
max_ngram_size=max_ngram_size,
layer_ids=layer_ids,
num_embeddings=tuple(args.engram_num_embeddings),
primes=tuple(primes),
n_heads=n_heads,
head_dim=args.engram_head_dim,
)
class NgramHashState(nn.Module):
"""Maps each position to the hash ids of the n-grams ending there.
Ids go through the compressed table, then each position is hashed with the `max_ngram_size - 1`
tokens before it. Look-back stops at the start of the sequence and at any dead token (an image
span, cached as DEAD), so an n-gram never spans one. The cache carries all of this across the
prefill/decode split.
"""
DEAD = -1
def __init__(self, args, layout: EngramLayout, tokenizer):
super().__init__()
self.layout = layout
# every hash multiplier derives from the compressed vocab size, so a mismatch there would
# silently rehash the whole table
token_map, vocab_size = build_compressed_token_map(tokenizer)
assert vocab_size == args.engram_compressed_vocab_size, (vocab_size, args.engram_compressed_vocab_size)
self.pad_id = token_map[args.engram_pad_id]
flat = [[p for per_ngram in layer for p in per_ngram] for layer in layout.primes]
offsets = [np.cumsum([0, *sizes[:-1]]) for sizes in flat]
multipliers = compute_hash_multipliers(layout.layer_ids, layout.max_ngram_size, vocab_size)
self.register_buffer("primes", torch.tensor(layout.primes), persistent=False)
self.register_buffer("offsets", torch.tensor(np.array(offsets)), persistent=False)
self.register_buffer("multipliers", multipliers, persistent=False)
self.register_buffer("token_map", torch.tensor(token_map), persistent=False)
self.register_buffer(
"cache", torch.empty(args.max_batch_size, args.max_seq_len, dtype=torch.int64), persistent=False
)
@torch.inference_mode()
def forward(self, input_ids: torch.Tensor, start_pos: int, token_mask: torch.Tensor | None = None) -> torch.Tensor:
"""token_mask: [B, L], False for tokens that take no part in an n-gram (image spans).
Returns the hash ids, shaped [B, L, n_engram_layers, n_hash_cols]."""
batch, seqlen = input_ids.shape
compressed = self.token_map[input_ids]
if token_mask is not None:
compressed = torch.where(token_mask, compressed, self.DEAD)
self.cache[:batch, start_pos : start_pos + seqlen] = compressed
positions = torch.arange(start_pos, start_pos + seqlen, device=input_ids.device).expand(batch, seqlen)
tokens, blocked = [], torch.zeros_like(positions, dtype=torch.bool)
for shift in range(self.layout.max_ngram_size):
source = self.cache[:batch].gather(1, (positions - shift).clamp_min(0))
blocked = blocked | (positions < shift) | (source == self.DEAD)
tokens.append(torch.where(blocked, self.pad_id, source))
tokens = torch.stack(tokens, dim=-1) # [B, L, max_ngram_size]
# XOR the multiplied ids together one lookback at a time, so the running value after step i
# is the hash of the (i+1)-gram; each lands in its own prime-sized bucket range
products = tokens.unsqueeze(2) * self.multipliers # [B, L, n_engram_layers, max_ngram_size]
rolling, hashes = products[..., 0], []
for i in range(1, self.layout.max_ngram_size):
rolling = torch.bitwise_xor(rolling, products[..., i])
hashes.append(rolling.unsqueeze(-1) % self.primes[:, i - 1])
return torch.cat(hashes, dim=-1) + self.offsets