Skip to content

Commit 723db80

Browse files
committed
feat(v4): ROI 6 — Engram primitives as minimal-edit port of TileKernels torch/engram.py
Vendor deepseek-ai/TileKernels tile_kernels/torch/engram.py verbatim and apply only MLX edits (torch->mx, in-place .bitwise_xor_/.cumsum_/.clone dropped or functional, .float/.bfloat16 mapped, einsum used directly via mx.einsum). Function names, signatures, dtypes, and algorithm preserved. cppmega_v4/nn/_external/tilekernels_engram.py: - make_offsets(vocab_sizes) — exclusive prefix-sum offsets per ngram layer - engram_hash_ref(token_ids, mult, vs, off) — per-layer per-ngram bitwise-xor hash with modulo-into-vocab indexing; produces final embedding indices - engram_gate_ref(x, k, v, wh, we, clamp, eps[, save_for_backward]) — fused RMSNorm(x, wh) . RMSNorm(k, we) -> scaled-dot -> signed-sqrt -> sigmoid -> additive gate. Returns bfloat16 output, optionally with saved intermediates (dot, gate_score, rstd_x, rstd_k) for backward. 7 new tests in tests/v4/test_engram_v4.py: - make_offsets shape + exact-prefix-sum values - engram_hash_ref shape contract - engram_gate_ref shape, bfloat16 dtype, save_for_backward tuple - *** Parity vs PyTorch tile_kernels.torch.engram (3 functions, all match upstream within int-exact for hash/offsets and bfloat16 atol 5e-2 for gate) *** 147/147 regression tests green (v4 + extensions + engram + nam pattern + MTP). Pipeline: Implementer (TileKernels torch port verbatim) -> Code Review (MIT attribution, no invented math) -> Perf (none — Path A reference) -> Regression Tests (parity passes).
1 parent 9b335f4 commit 723db80

2 files changed

Lines changed: 322 additions & 0 deletions

File tree

Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,149 @@
1+
# Verbatim port of deepseek-ai/TileKernels/tile_kernels/torch/engram.py with
2+
# the smallest possible edits to run on MLX instead of PyTorch.
3+
#
4+
# Upstream source:
5+
# /Users/dave/sources/TileKernels/tile_kernels/torch/engram.py
6+
# Upstream license: MIT (deepseek-ai/TileKernels)
7+
#
8+
# Edits made (kept minimal so a diff against the upstream file is short):
9+
# - import torch -> import mlx.core as mx
10+
# - torch.Tensor annotations -> mx.array
11+
# - x.view(-1) -> x.reshape(-1)
12+
# - x.unsqueeze(0|1|-1) -> x[None] / x[..., None]
13+
# - .to(torch.int32) / int64 -> .astype(mx.int32) / int64
14+
# - .clone() -> dropped (functional)
15+
# - hashes.bitwise_xor_(other) -> hashes = mx.bitwise_xor(hashes, other)
16+
# (MLX is functional, no in-place ops)
17+
# - x.cumsum(0, dtype=torch.int32) -> mx.cumsum(x, axis=0).astype(mx.int32)
18+
# - torch.cat([...], dim=...) -> mx.concatenate([...], axis=...)
19+
# - torch.stack([...], dim=...) -> mx.stack([...], axis=...)
20+
# - torch.zeros(1, dtype=..., device=...) -> mx.zeros((1,), dtype=...)
21+
# - x.float() / x.bfloat16() -> x.astype(mx.float32|bfloat16)
22+
# - torch.rsqrt / x.pow(2) / x.sigmoid -> mx.rsqrt / mx.square / mx.sigmoid
23+
# - torch.einsum('...d,...d->...', a, b) -> mx.einsum('...d,...d->...', a, b)
24+
# - dot.abs().clamp_min(v).sqrt() * dot.sign() -> mx.sign(dot) * mx.sqrt(mx.maximum(mx.abs(dot), v))
25+
# Algorithm, signatures, dtypes, return shapes unchanged.
26+
27+
import mlx.core as mx
28+
29+
30+
def make_offsets(vocab_sizes: mx.array) -> mx.array:
31+
"""Compute exclusive prefix-sum offsets from vocab_sizes.
32+
33+
Args:
34+
vocab_sizes: Per-layer per-ngram embedding table sizes of shape
35+
(num_ngram_layers, max_ngram_size - 1, num_embed_table_per_ngram), int32.
36+
37+
Returns:
38+
Offsets of shape (num_ngram_layers, (max_ngram_size - 1) * num_embed_table_per_ngram), int32.
39+
"""
40+
num_ngram_layers = vocab_sizes.shape[0]
41+
offsets_list = []
42+
for layer_idx in range(num_ngram_layers):
43+
flat = vocab_sizes[layer_idx].reshape(-1)
44+
prefix = mx.concatenate(
45+
[
46+
mx.zeros((1,), dtype=mx.int32),
47+
mx.cumsum(flat[:-1], axis=0).astype(mx.int32),
48+
]
49+
)
50+
offsets_list.append(prefix)
51+
return mx.stack(offsets_list, axis=0)
52+
53+
54+
def engram_hash_ref(
55+
ngram_token_ids: mx.array,
56+
multipliers: mx.array,
57+
vocab_sizes: mx.array,
58+
offsets: mx.array,
59+
) -> mx.array:
60+
"""Pure PyTorch reference implementation of engram hash.
61+
62+
Args:
63+
ngram_token_ids: N-gram token IDs of shape (num_tokens, max_ngram_size), int32.
64+
multipliers: Per-layer hash multipliers of shape (num_ngram_layers, max_ngram_size), int64.
65+
vocab_sizes: Per-layer per-ngram embedding table sizes of shape
66+
(num_ngram_layers, max_ngram_size - 1, num_embed_table_per_ngram), int32.
67+
offsets: Per-layer embedding table offsets of shape
68+
(num_ngram_layers, (max_ngram_size - 1) * num_embed_table_per_ngram), int32.
69+
70+
Returns:
71+
Embedding indices of shape (num_ngram_layers, num_tokens, (max_ngram_size - 1) * num_embed_table_per_ngram), int32.
72+
"""
73+
num_ngram_layers = multipliers.shape[0]
74+
max_ngram_size = multipliers.shape[1]
75+
76+
prod = ngram_token_ids.astype(mx.int64)[None] * multipliers[:, None]
77+
78+
ans = [[] for _ in range(num_ngram_layers)]
79+
hashes = prod[:, :, 0]
80+
for i in range(1, max_ngram_size):
81+
hashes = mx.bitwise_xor(hashes, prod[:, :, i])
82+
for layer_idx in range(num_ngram_layers):
83+
ans[layer_idx].append(
84+
(
85+
hashes[layer_idx][..., None]
86+
% vocab_sizes[layer_idx, i - 1].astype(mx.int64)[None]
87+
).astype(mx.int32)
88+
)
89+
90+
for layer_idx in range(num_ngram_layers):
91+
ans[layer_idx] = mx.concatenate(ans[layer_idx], axis=-1)
92+
93+
output = mx.stack(ans, axis=0)
94+
return output + offsets[:, None]
95+
96+
97+
def engram_gate_ref(
98+
hidden_states: mx.array,
99+
k: mx.array,
100+
v: mx.array,
101+
weight_hidden: mx.array,
102+
weight_embed: mx.array,
103+
clamp_value: float,
104+
eps: float,
105+
save_for_backward: bool = False,
106+
):
107+
"""Pure PyTorch reference implementation of engram gate (vectorized, supports autograd).
108+
109+
Computes: output = x + sigmoid(signed_sqrt(dot(RMSNorm(x, wh), RMSNorm(k, we)) * scalar)) * v
110+
111+
Args:
112+
hidden_states: Input of shape (num_tokens, hc_mult, hidden_size), bfloat16.
113+
k: Key embeddings of shape (num_tokens, hc_mult, hidden_size), bfloat16.
114+
v: Value embeddings of shape (num_tokens, hidden_size), bfloat16.
115+
weight_hidden: RMSNorm weight for hidden states, shape (hc_mult, hidden_size), bfloat16.
116+
weight_embed: RMSNorm weight for key embeddings, shape (hc_mult, hidden_size), bfloat16.
117+
clamp_value: Clamp threshold for signed-sqrt gate activation.
118+
eps: Epsilon for RMSNorm numerical stability.
119+
save_for_backward: If True, also return (dot, gate_score, rstd_x, rstd_k).
120+
121+
Returns:
122+
If save_for_backward is False: output tensor of shape (num_tokens, hc_mult, hidden_size), bfloat16.
123+
If save_for_backward is True: tuple of (output, dot, gate_score, rstd_x, rstd_k).
124+
"""
125+
hidden_size = hidden_states.shape[-1]
126+
scalar = hidden_size**-0.5
127+
128+
x = hidden_states.astype(mx.float32)
129+
k_f = k.astype(mx.float32)
130+
wh = weight_hidden.astype(mx.float32)[None]
131+
we = weight_embed.astype(mx.float32)[None]
132+
133+
# RMSNorm
134+
rstd_x = mx.rsqrt(mx.mean(mx.square(x), axis=-1) + eps)
135+
rstd_k = mx.rsqrt(mx.mean(mx.square(k_f), axis=-1) + eps)
136+
137+
# Dot -> sqrt-gate -> sigmoid
138+
# raw_dot is the unnormalized sum(x * wh * k * we), matching the kernel's dot_out
139+
raw_dot = mx.einsum('...d,...d->...', x * wh, k_f * we)
140+
dot = raw_dot * rstd_x * rstd_k * scalar
141+
signed_sqrt = mx.sign(dot) * mx.sqrt(mx.maximum(mx.abs(dot), clamp_value))
142+
gate_score = mx.sigmoid(signed_sqrt)
143+
144+
output = x + gate_score[..., None] * v[..., None, :]
145+
output = output.astype(mx.bfloat16)
146+
147+
if save_for_backward:
148+
return output, raw_dot, gate_score, rstd_x, rstd_k
149+
return output

tests/v4/test_engram_v4.py

Lines changed: 173 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,173 @@
1+
"""Tests for cppmega_v4.nn._external.tilekernels_engram — port of DeepSeek engram torch ref.
2+
3+
Parity tests against ``~/sources/TileKernels/tile_kernels/torch/engram.py``
4+
when torch is available.
5+
"""
6+
7+
from __future__ import annotations
8+
9+
import sys
10+
from pathlib import Path
11+
12+
import mlx.core as mx
13+
import numpy as np
14+
import pytest
15+
16+
from cppmega_v4.nn._external.tilekernels_engram import (
17+
engram_gate_ref,
18+
engram_hash_ref,
19+
make_offsets,
20+
)
21+
22+
23+
# ----- make_offsets shape + values -----
24+
25+
26+
def test_make_offsets_shape_and_values():
27+
# 2 layers, 3 ngram positions, 2 tables/ngram -> flat per layer has 6 entries.
28+
vocab_sizes = mx.array([[[10, 20], [30, 40], [50, 60]],
29+
[[5, 5], [5, 5], [5, 5]]], dtype=mx.int32)
30+
offsets = make_offsets(vocab_sizes)
31+
assert offsets.shape == (2, 6)
32+
# Layer 0 exclusive prefix sum of [10, 20, 30, 40, 50, 60] = [0,10,30,60,100,150]
33+
np.testing.assert_array_equal(
34+
np.array(offsets[0]), np.array([0, 10, 30, 60, 100, 150], dtype=np.int32)
35+
)
36+
np.testing.assert_array_equal(
37+
np.array(offsets[1]), np.array([0, 5, 10, 15, 20, 25], dtype=np.int32)
38+
)
39+
40+
41+
def test_engram_hash_shape_contract():
42+
num_tokens = 4
43+
max_ngram_size = 3
44+
num_ngram_layers = 2
45+
num_embed = 2
46+
rng = np.random.default_rng(42)
47+
ngram_token_ids = mx.array(
48+
rng.integers(0, 100, size=(num_tokens, max_ngram_size)).astype(np.int32)
49+
)
50+
multipliers = mx.array(
51+
rng.integers(1, 10000, size=(num_ngram_layers, max_ngram_size)).astype(np.int64)
52+
)
53+
vocab_sizes = mx.array(
54+
rng.integers(5, 50, size=(num_ngram_layers, max_ngram_size - 1, num_embed)).astype(np.int32)
55+
)
56+
offsets = make_offsets(vocab_sizes)
57+
out = engram_hash_ref(ngram_token_ids, multipliers, vocab_sizes, offsets)
58+
# out shape (num_ngram_layers, num_tokens, (max_ngram_size - 1) * num_embed_table_per_ngram)
59+
assert out.shape == (num_ngram_layers, num_tokens, (max_ngram_size - 1) * num_embed)
60+
61+
62+
def test_engram_gate_shape_and_dtype():
63+
num_tokens, hc_mult, hidden = 4, 2, 16
64+
rng = np.random.default_rng(7)
65+
hs = mx.array(rng.standard_normal((num_tokens, hc_mult, hidden)).astype(np.float32))
66+
k = mx.array(rng.standard_normal((num_tokens, hc_mult, hidden)).astype(np.float32))
67+
v = mx.array(rng.standard_normal((num_tokens, hidden)).astype(np.float32))
68+
wh = mx.ones((hc_mult, hidden))
69+
we = mx.ones((hc_mult, hidden))
70+
out = engram_gate_ref(
71+
hs.astype(mx.bfloat16), k.astype(mx.bfloat16), v.astype(mx.bfloat16),
72+
wh.astype(mx.bfloat16), we.astype(mx.bfloat16),
73+
clamp_value=1e-4, eps=1e-6,
74+
)
75+
assert out.shape == (num_tokens, hc_mult, hidden)
76+
assert out.dtype == mx.bfloat16
77+
78+
79+
def test_engram_gate_save_for_backward_returns_tuple():
80+
num_tokens, hc_mult, hidden = 2, 1, 8
81+
hs = mx.random.normal((num_tokens, hc_mult, hidden)).astype(mx.bfloat16)
82+
k = mx.random.normal((num_tokens, hc_mult, hidden)).astype(mx.bfloat16)
83+
v = mx.random.normal((num_tokens, hidden)).astype(mx.bfloat16)
84+
wh = mx.ones((hc_mult, hidden)).astype(mx.bfloat16)
85+
we = mx.ones((hc_mult, hidden)).astype(mx.bfloat16)
86+
out, dot, gate, rstd_x, rstd_k = engram_gate_ref(
87+
hs, k, v, wh, we, clamp_value=1e-4, eps=1e-6, save_for_backward=True
88+
)
89+
assert out.shape == (num_tokens, hc_mult, hidden)
90+
assert dot.shape == (num_tokens, hc_mult)
91+
assert gate.shape == (num_tokens, hc_mult)
92+
assert rstd_x.shape == (num_tokens, hc_mult)
93+
assert rstd_k.shape == (num_tokens, hc_mult)
94+
95+
96+
# ----- parity vs PyTorch TileKernels reference -----
97+
98+
99+
@pytest.fixture(scope="module")
100+
def tk_engram_torch():
101+
torch = pytest.importorskip("torch")
102+
repo = Path("/Users/dave/sources/TileKernels")
103+
if not repo.exists():
104+
pytest.skip("TileKernels repo not present at expected path")
105+
if str(repo) not in sys.path:
106+
sys.path.insert(0, str(repo))
107+
try:
108+
from tile_kernels.torch.engram import (
109+
engram_gate_ref as tk_gate,
110+
engram_hash_ref as tk_hash,
111+
make_offsets as tk_offsets,
112+
)
113+
except Exception as exc:
114+
pytest.skip(f"could not import TileKernels torch.engram: {exc}")
115+
return {"torch": torch, "gate": tk_gate, "hash": tk_hash, "offsets": tk_offsets}
116+
117+
118+
def test_parity_make_offsets(tk_engram_torch):
119+
torch = tk_engram_torch["torch"]
120+
rng = np.random.default_rng(11)
121+
vs_np = rng.integers(1, 100, size=(2, 3, 2)).astype(np.int32)
122+
t_out = tk_engram_torch["offsets"](torch.from_numpy(vs_np))
123+
m_out = make_offsets(mx.array(vs_np))
124+
np.testing.assert_array_equal(np.array(m_out), t_out.numpy())
125+
126+
127+
def test_parity_engram_hash(tk_engram_torch):
128+
torch = tk_engram_torch["torch"]
129+
rng = np.random.default_rng(12)
130+
num_tokens, max_ngram, layers, num_embed = 4, 3, 2, 2
131+
ng = rng.integers(0, 100, size=(num_tokens, max_ngram)).astype(np.int32)
132+
mult = rng.integers(1, 10000, size=(layers, max_ngram)).astype(np.int64)
133+
vs = rng.integers(5, 50, size=(layers, max_ngram - 1, num_embed)).astype(np.int32)
134+
off = make_offsets(mx.array(vs))
135+
t_off = tk_engram_torch["offsets"](torch.from_numpy(vs))
136+
t_out = tk_engram_torch["hash"](
137+
torch.from_numpy(ng), torch.from_numpy(mult), torch.from_numpy(vs), t_off
138+
)
139+
m_out = engram_hash_ref(mx.array(ng), mx.array(mult), mx.array(vs), off)
140+
np.testing.assert_array_equal(np.array(m_out), t_out.numpy())
141+
142+
143+
def test_parity_engram_gate(tk_engram_torch):
144+
torch = tk_engram_torch["torch"]
145+
rng = np.random.default_rng(13)
146+
nt, hc, h = 3, 2, 16
147+
hs = rng.standard_normal((nt, hc, h)).astype(np.float32)
148+
k = rng.standard_normal((nt, hc, h)).astype(np.float32)
149+
v = rng.standard_normal((nt, h)).astype(np.float32)
150+
wh = rng.standard_normal((hc, h)).astype(np.float32)
151+
we = rng.standard_normal((hc, h)).astype(np.float32)
152+
t_out = tk_engram_torch["gate"](
153+
torch.from_numpy(hs).bfloat16(),
154+
torch.from_numpy(k).bfloat16(),
155+
torch.from_numpy(v).bfloat16(),
156+
torch.from_numpy(wh).bfloat16(),
157+
torch.from_numpy(we).bfloat16(),
158+
clamp_value=1e-4, eps=1e-6,
159+
)
160+
m_out = engram_gate_ref(
161+
mx.array(hs).astype(mx.bfloat16),
162+
mx.array(k).astype(mx.bfloat16),
163+
mx.array(v).astype(mx.bfloat16),
164+
mx.array(wh).astype(mx.bfloat16),
165+
mx.array(we).astype(mx.bfloat16),
166+
clamp_value=1e-4, eps=1e-6,
167+
)
168+
# bfloat16 path — tolerate ~5e-2 max abs error.
169+
np.testing.assert_allclose(
170+
np.array(m_out.astype(mx.float32)),
171+
t_out.float().numpy(),
172+
atol=5e-2, rtol=5e-2,
173+
)

0 commit comments

Comments
 (0)