Meridian-Tiny / model.py
hashtagg1's picture
Upload 4 files
fb9477a verified
Raw History Blame Contribute Delete
23.9 kB
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
try:
from fla.ops.kda import chunk_kda
except ImportError:
chunk_kda = None
USE_FLA_KDA = True
@dataclass
class MeridianConfig:
# core
vocab_size: int = 32768
hidden_size: int = 512
num_layers: int = 12
max_seq_len: int = 2048
rms_eps: float = 1e-6
# attention (MLA): full attention every 4th layer
full_attn_every: int = 4
num_heads: int = 8
q_lora_rank: int = 256
kv_lora_rank: int = 128
qk_nope_head_dim: int = 64
qk_rope_head_dim: int = 32
v_head_dim: int = 64
rope_theta: float = 10000.0
# linear attention (KDA)
kda_head_dim: int = 64
kda_conv_size: int = 4
kda_gate_lower_bound: float = -5.0
# MoE
num_experts: int = 32
experts_per_token: int = 4
num_shared_experts: int = 1
moe_latent_size: int = 384
moe_intermediate_size: int = 256
dense_intermediate_size: int = 2048
first_k_dense: int = 1
num_expert_groups: int = 4
topk_groups: int = 2
routed_scaling: float = 2.5
swiglu_limit: float = 10.0
moe_capacity_factor: float = 2.0
# mHC + AttnRes
hc_mult: int = 4
hc_sinkhorn_iters: int = 20
attn_res_block_size: int = 4
# Engram
engram_layers: tuple = (1, 6)
engram_max_ngram: int = 3
engram_table_size: int = 524_287
engram_heads: int = 4
engram_head_dim: int = 64
def is_full_attn(self, layer_idx):
return (layer_idx + 1) % self.full_attn_every == 0
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
rms = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return (x.float() * rms * self.weight.float()).type_as(x)
def build_rope_cache(seq_len, dim, theta, device=None):
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, device=device).float () / dim))
t = torch.arange(seq_len, device=device).float()
freqs = torch.outer(t, inv_freq)
emb = torch.cat([freqs, freqs], dim=-1)
return emb.cos(), emb.sin()
def rotate_half(x):
x1, x2 = x.chunk(2, dim=-1)
return torch.cat([-x2, x1], dim=-1)
def apply_rope(x, cos, sin):
return x * cos + rotate_half(x) * sin
class MLA(nn.Module):
def __init__(self, cfg):
super().__init__()
self.num_heads = cfg.num_heads
self.nope = cfg.qk_nope_head_dim
self.rope = cfg.qk_rope_head_dim
self.v_dim = cfg.v_head_dim
self.kv_rank = cfg.kv_lora_rank
h = cfg.num_heads
self.q_a = nn.Linear(cfg.hidden_size, cfg.q_lora_rank, bias=False)
self.q_norm = RMSNorm(cfg.q_lora_rank, cfg.rms_eps)
self.q_b = nn.Linear(cfg.q_lora_rank, h * (self.nope + self.rope), bias=False)
self.kv_a = nn.Linear(cfg.hidden_size, self.kv_rank + self.rope, bias=False)
self.kv_norm = RMSNorm(self.kv_rank, cfg.rms_eps)
self.kv_b = nn.Linear(self.kv_rank, h * (self.nope + self.v_dim), bias=False)
self.gate = nn.Linear(cfg.hidden_size, h * self.v_dim, bias=False)
self.o_proj = nn.Linear(h * self.v_dim, cfg.hidden_size, bias=False)
def forward(self, x, cos, sin):
b, s, _ = x.shape
h = self.num_heads
q = self.q_b(self.q_norm(self.q_a(x)))
q = q.view(b, s, h, self.nope + self.rope).transpose(1, 2)
q_nope, q_rope = q.split([self.nope, self.rope], dim=-1)
c_kv, k_rope = self.kv_a(x).split([self.kv_rank, self.rope], dim=-1)
kv = self.kv_b(self.kv_norm(c_kv))
kv = kv.view(b, s, h, self.nope + self.v_dim).transpose(1, 2)
k_nope, v = kv.split([self.nope, self.v_dim], dim=-1)
q_rope = apply_rope(q_rope, cos, sin)
k_rope = apply_rope(k_rope.unsqueeze(1), cos, sin).expand(b, h, s, self.rope)
q = torch.cat([q_nope, q_rope], dim=-1)
k = torch.cat([k_nope, k_rope], dim=-1)
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
out = out.transpose(1, 2).reshape(b, s, h * self.v_dim)
out = out * torch.sigmoid(self.gate(x))
return self.o_proj(out)
def kda_recurrent(q, k, v, a, beta):
out_dtype = v.dtype
q, k, v, a, beta = (x.float() for x in (q, k, v, a, beta))
b, s, h, dk = q.shape
dv = v.shape[-1]
S = torch.zeros(b, h, dk, dv, device=q.device)
outs = []
for t in range(s):
S = S * a[:, t, :, :, None]
pred = torch.einsum("bhk,bhkv->bhv", k[:, t], S)
S = S + beta[:, t, :, None, None] * k[:, t, :, :, None] * (v[:, t] - pred)[:, :, None, :]
outs.append(torch.einsum("bhk,bhkv->bhv", q[:, t], S))
return torch.stack(outs, dim=1).to(out_dtype)
class KDA(nn.Module):
def __init__(self, cfg):
super().__init__()
self.num_heads = cfg.num_heads
self.head_dim = cfg.kda_head_dim
self.lower_bound = cfg.kda_gate_lower_bound
inner = cfg.num_heads * cfg.kda_head_dim
self.qkv = nn.Linear(cfg.hidden_size, 3 * inner, bias=False)
self.conv = nn.Conv1d(3 * inner, 3 * inner, cfg.kda_conv_size,
groups=3 * inner, padding=cfg.kda_conv_size - 1, bias=False)
self.a_proj = nn.Linear(cfg.hidden_size, inner, bias=False)
self.b_proj = nn.Linear(cfg.hidden_size, cfg.num_heads, bias=False)
self.o_norm = RMSNorm(cfg.kda_head_dim, cfg.rms_eps)
self.gate = nn.Linear(cfg.hidden_size, inner, bias=False)
self.o_proj = nn.Linear(inner, cfg.hidden_size, bias=False)
def forward(self, x):
b, s, _ = x.shape
h, d = self.num_heads, self.head_dim
qkv = self.conv(self.qkv(x).transpose(1, 2))[..., :s].transpose(1, 2)
q, k, v = F.silu(qkv).chunk(3, dim=-1)
q = F.normalize(q.view(b, s, h, d), dim=-1) * d ** -0.5
k = F.normalize(k.view(b, s, h, d), dim=-1)
v = v.view(b, s, h, d)
log_a = (self.lower_bound * torch.sigmoid(self.a_proj(x).float())).view(b, s, h, d)
beta = torch.sigmoid(self.b_proj(x))
if x.is_cuda and chunk_kda is not None and USE_FLA_KDA:
dt = v.dtype
o, _ = chunk_kda(q=q.to(dt), k=k.to(dt), v=v.to(dt), g=log_a, beta=beta.to(dt), scale=1.0)
else:
o = kda_recurrent(q, k, v, log_a.exp(), beta)
o = self.o_norm(o)
o = o.reshape(b, s, h * d) * torch.sigmoid(self.gate(x))
return self.o_proj(o)
class SwiGLU(nn.Module):
def __init__(self, dim, hidden, limit):
super().__init__()
self.limit = limit
self.gate = nn.Linear(dim, hidden, bias=False)
self.up = nn.Linear(dim, hidden, bias=False)
self.down = nn.Linear(hidden, dim, bias=False)
def forward(self, x):
g = self.gate(x).clamp(max=self.limit)
u = self.up(x).clamp(-self.limit, self.limit)
return self.down(F.silu(g) * u)
class MoE(nn.Module):
def __init__(self, cfg):
super().__init__()
E, L, I = cfg.num_experts, cfg.moe_latent_size, cfg.moe_intermediate_size
self.num_experts = E
self.k = cfg.experts_per_token
self.n_groups = cfg.num_expert_groups
self.per_group = E // cfg.num_expert_groups
self.topk_groups = cfg.topk_groups
self.scaling = cfg.routed_scaling
self.limit = cfg.swiglu_limit
self.router = nn.Linear(cfg.hidden_size, E, bias=False)
self.register_buffer("route_bias", torch.zeros(E))
self.register_buffer("expert_counts", torch.zeros(E), persistent=False)
self.latent_down = nn.Linear(cfg.hidden_size, L, bias=False)
self.latent_norm = RMSNorm(L, cfg.rms_eps)
self.latent_up = nn.Linear(L, cfg.hidden_size, bias=False)
self.w_gate = nn.Parameter(torch.randn(E, L, I) * 0.02)
self.w_up = nn.Parameter(torch.randn(E, L, I) * 0.02)
self.w_down = nn.Parameter(torch.randn(E, I, L) * 0.02)
self.shared = SwiGLU(cfg.hidden_size, I * cfg.num_shared_experts, cfg.swiglu_limit)
self.capacity_factor = cfg.moe_capacity_factor
def route(self, x):
scores = torch.sigmoid(self.router(x).float())
biased = scores + self.route_bias
grouped = biased.view(-1, self.n_groups, self.per_group)
group_scores = grouped.topk(2, dim=-1).values.sum(-1)
top_groups = group_scores.topk(self.topk_groups, dim=-1).indices
group_mask = torch.zeros_like(group_scores).scatter_(1, top_groups, 1.0)
expert_mask = group_mask.unsqueeze(-1).expand_as(grouped).reshape(-1, self.num_experts)
masked = biased.masked_fill(expert_mask == 0, float("-inf"))
idx = masked.topk(self.k, dim=-1).indices
weights = scores.gather(1, idx)
weights = weights / weights.sum(-1, keepdim=True) * self.scaling
return weights, idx
def forward(self, x):
b, s, d = x.shape
x_flat = x.reshape(-1, d)
n = x_flat.shape[0]
weights, idx = self.route(x_flat)
flat_e = idx.reshape(-1)
counts = torch.zeros(self.num_experts, device=x.device)
counts.scatter_add_(0, flat_e, torch.ones_like(flat_e, dtype=counts.dtype))
if self.training:
self.expert_counts += counts
cap = int(self.capacity_factor * n * self.k / self.num_experts) + 1
order = torch.argsort(flat_e, stable=True)
starts = (torch.cumsum(counts, 0) - counts).long()
rank_sorted = torch.arange(flat_e.numel(), device=x.device) - starts[flat_e[order]]
rank = torch.empty_like(rank_sorted).scatter_(0, order, rank_sorted)
slot = rank.clamp(max=cap)
keep = (rank < cap).to(weights.dtype)
z = self.latent_norm(self.latent_down(x_flat))
buf = z.new_zeros(self.num_experts, cap + 1, z.shape[-1])
buf[flat_e, slot] = z.repeat_interleave(self.k, dim=0)
g = torch.bmm(buf, self.w_gate).clamp(max=self.limit)
u = torch.bmm(buf, self.w_up).clamp(-self.limit, self.limit)
h = torch.bmm(F.silu(g) * u, self.w_down)
w = (weights.reshape(-1) * keep).unsqueeze(-1).to(h.dtype)
out = (h[flat_e, slot] * w).view(n, self.k, -1).sum(1).to(z.dtype)
y = self.latent_up(out) + self.shared(x_flat)
return y.view(b, s, d)
def forward_reference(self, x):
b, s, d = x.shape
x_flat = x.reshape(-1, d)
weights, idx = self.route(x_flat)
if self.training:
self.expert_counts += torch.bincount(idx.flatten(), minlength=self.num_experts).float()
z = self.latent_norm(self.latent_down(x_flat))
out = torch.zeros_like(z)
for e in range(self.num_experts):
token_ids, slot = (idx == e).nonzero(as_tuple=True)
if token_ids.numel() == 0:
continue
ze = z[token_ids]
g = (ze @ self.w_gate[e]).clamp(max=self.limit)
u = (ze @ self.w_up[e]).clamp(-self.limit, self.limit)
h = (F.silu(g) * u) @ self.w_down[e]
out.index_add_(0, token_ids, (h * weights[token_ids, slot, None]).to(out.dtype))
y = self.latent_up(out) + self.shared(x_flat)
return y.view(b, s, d)
@torch.no_grad()
def update_bias(self, rate=1e-3):
load = self.expert_counts
self.route_bias += rate * torch.sign(load.mean() - load)
self.expert_counts.zero_()
def sinkhorn(logits, iters):
m = (logits - logits.amax(dim=(-1, -2), keepdim=True)).exp()
for _ in range(iters):
m = m / m.sum(-1, keepdim=True)
m = m / m.sum(-2, keepdim=True)
return m
_sinkhorn_compiled = torch.compile(sinkhorn)
USE_COMPILED_SINKHORN = True
def sinkhorn_fast(logits, iters):
if logits.shape[-1] == 1:
return torch.ones_like(logits)
if logits.is_cuda and USE_COMPILED_SINKHORN:
return _sinkhorn_compiled(logits, iters)
return sinkhorn(logits, iters)
class HyperConnection(nn.Module):
def __init__(self, cfg):
super().__init__()
n, d = cfg.hc_mult, cfg.hidden_size
self.n = n
self.iters = cfg.hc_sinkhorn_iters
self.norm = RMSNorm(n * d, cfg.rms_eps)
self.proj = nn.Linear(n * d, n + n + n * n, bias=False)
nn.init.zeros_(self.proj.weight)
self.alpha = nn.Parameter(torch.full((3,), 0.01))
self.b_pre = nn.Parameter(torch.full((n,), -1.0986))
self.b_post = nn.Parameter(torch.zeros(n))
self.b_res = nn.Parameter(torch.eye(n) * 4.0)
def forward(self, streams, fn):
b, s, n, d = streams.shape
h = self.proj(self.norm(streams.reshape(b, s, n * d))).float()
h_pre, h_post, h_res = h.split([n, n, n * n], dim=-1)
pre = torch.sigmoid(self.alpha[0] * h_pre + self.b_pre)
post = 2 * torch.sigmoid(self.alpha[1] * h_post + self.b_post)
res = sinkhorn_fast(self.alpha[2] * h_res.view(b, s, n, n) + self.b_res, self.iters)
x_in = torch.einsum("bsn,bsnd->bsd", pre.to(streams.dtype), streams)
y = fn(x_in)
mixed = torch.einsum("bsmn,bsnd->bsmd", res.to(streams.dtype), streams)
return mixed + post.to(streams.dtype).unsqueeze(-1) * y.unsqueeze(2)
class AttnRes(nn.Module):
def __init__(self, cfg):
super().__init__()
self.norm = RMSNorm(cfg.hidden_size, cfg.rms_eps)
self.query = nn.Parameter(torch.zeros(cfg.hidden_size))
def forward(self, snapshots):
stack = torch.stack(snapshots, dim=0)
keys = self.norm(stack.mean(dim=3))
scores = torch.einsum("kbsd,d->kbs", keys.float(), self.query.float())
w = torch.softmax(scores, dim=0)
return torch.einsum("kbs,kbsnd->bsnd", w.to(stack.dtype), stack)
class Engram(nn.Module):
def __init__(self, cfg):
super().__init__()
self.orders = list(range(2, cfg.engram_max_ngram + 1))
self.heads = cfg.engram_heads
self.table_size = cfg.engram_table_size
n_lookups = len(self.orders) * cfg.engram_heads
mem_dim = n_lookups * cfg.engram_head_dim
self.table = nn.Embedding(cfg.engram_table_size, cfg.engram_head_dim)
nn.init.normal_(self.table.weight, std=0.02)
g = torch.Generator().manual_seed(1234)
mults = torch.randint(1, 2**30, (n_lookups, cfg.engram_max_ngram), generator=g) * 2 + 1
self.register_buffer("mults", mults)
self.key = nn.Linear(mem_dim, cfg.hidden_size, bias=False)
self.value = nn.Linear(mem_dim, cfg.hidden_size, bias=False)
self.h_norm = RMSNorm(cfg.hidden_size, cfg.rms_eps)
self.k_norm = RMSNorm(cfg.hidden_size, cfg.rms_eps)
d = cfg.hidden_size
self.conv = nn.Conv1d(d, d, 4, groups=d, padding=3, bias=False)
def hash_ids(self, ids):
b, s = ids.shape
max_n = max(self.orders)
padded = F.pad(ids, (max_n - 1, 0), value=0)
shifted = [padded[:, max_n - 1 - j : max_n - 1 - j + s] for j in range(max_n)]
out = []
i = 0
for n in self.orders:
for _ in range(self.heads):
h = torch.zeros_like(ids)
for j in range(n):
h = h ^ (shifted[j] * self.mults[i, j])
out.append(h % self.table_size)
i += 1
return torch.stack(out, dim=-1)
def forward(self, h, ids):
b, s, _ = h.shape
mem = self.table(self.hash_ids(ids)).reshape(b, s, -1)
k = self.k_norm(self.key(mem))
v = self.value(mem)
gate = torch.sigmoid((self.h_norm(h) * k).sum(-1, keepdim=True) / h.shape[-1] ** 0.5)
out = gate * v
return out + self.conv(out.transpose(1, 2))[..., :s].transpose(1, 2)
class Block(nn.Module):
def __init__(self, cfg, idx):
super().__init__()
d = cfg.hidden_size
self.full_attn = cfg.is_full_attn(idx)
self.engram = Engram(cfg) if idx in cfg.engram_layers else None
if self.engram is not None:
self.hc_engram = HyperConnection(cfg)
self.mixer_norm = RMSNorm(d, cfg.rms_eps)
self.mixer = MLA(cfg) if self.full_attn else KDA(cfg)
self.hc_mixer = HyperConnection(cfg)
self.ffn_norm = RMSNorm(d, cfg.rms_eps)
if idx < cfg.first_k_dense:
self.ffn = SwiGLU(d, cfg.dense_intermediate_size, cfg.swiglu_limit)
else:
self.ffn = MoE(cfg)
self.hc_ffn = HyperConnection(cfg)
def forward(self, streams, ids, cos, sin):
if self.engram is not None:
streams = self.hc_engram(streams, lambda z: self.engram(z, ids))
if self.full_attn:
streams = self.hc_mixer(streams, lambda z: self.mixer(self.mixer_norm(z), cos, sin))
else:
streams = self.hc_mixer(streams, lambda z: self.mixer(self.mixer_norm(z)))
return self.hc_ffn(streams, lambda z: self.ffn(self.ffn_norm(z)))
class MiniMeridian(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
d = cfg.hidden_size
self.embed = nn.Embedding(cfg.vocab_size, d)
self.blocks = nn.ModuleList([Block(cfg, i) for i in range(cfg.num_layers)])
bs = cfg.attn_res_block_size
self.boundaries = list(range(bs, cfg.num_layers, bs)) if bs > 0 else []
self.attn_res = nn.ModuleList([AttnRes(cfg) for _ in self.boundaries])
self.final_norm = RMSNorm(d, cfg.rms_eps)
self.lm_head = nn.Linear(d, cfg.vocab_size, bias=False)
cos, sin = build_rope_cache(cfg.max_seq_len, cfg.qk_rope_head_dim, cfg.rope_theta)
self.register_buffer("rope_cos", cos, persistent=False)
self.register_buffer("rope_sin", sin, persistent=False)
self.apply(self._init_weights)
for m in self.modules():
if isinstance(m, HyperConnection):
nn.init.zeros_(m.proj.weight)
def _init_weights(self, m):
if isinstance(m, (nn.Linear, nn.Embedding)):
nn.init.normal_(m.weight, std=0.02)
def forward(self, ids, targets=None):
b, s = ids.shape
x = self.embed(ids)
streams = x.unsqueeze(2).expand(-1, -1, self.cfg.hc_mult, -1).contiguous()
cos, sin = self.rope_cos[:s], self.rope_sin[:s]
snapshots = [streams]
for i, block in enumerate(self.blocks):
if i in self.boundaries:
snapshots.append(streams)
streams = self.attn_res[self.boundaries.index(i)](snapshots)
streams = block(streams, ids, cos, sin)
h = self.final_norm(streams.mean(dim=2))
logits = self.lm_head(h)
if targets is None:
return logits
loss = F.cross_entropy(logits.view(-1, logits.size(-1)).float(), targets.view(-1))
return logits, loss
if __name__ == "__main__":
cfg = MeridianConfig()
print("full-attn layers:", [i for i in range(cfg.num_layers) if cfg.is_full_attn(i)])
norm = RMSNorm(cfg.hidden_size)
x = torch.randn(2, 16, cfg.hidden_size)
print("norm out:", norm(x).shape)
cos, sin = build_rope_cache(16, cfg.qk_rope_head_dim, cfg.rope_theta)
q = torch.randn(2, cfg.num_heads, 16, cfg.qk_rope_head_dim)
q_rot = apply_rope(q, cos, sin)
print("rope out:", q_rot.shape)
print("length preserved:", torch.allclose(q.norm(dim=-1), q_rot.norm(dim=-1), atol=1e-5))
attn = MLA(cfg)
y = attn(x, cos, sin)
print("mla out:", y.shape)
x2 = x.clone()
x2[:, -1] = torch.randn(2, cfg.hidden_size)
y2 = attn(x2, cos, sin)
print("causal:", torch.allclose(y[:, :-1], y2[:, :-1], atol=1e-5))
n_params = sum(p.numel() for p in attn.parameters())
print(f"mla params: {n_params:,}")
kda = KDA(cfg)
y = kda(x)
print("kda out:", y.shape)
y2 = kda(x2)
print("kda causal:", torch.allclose(y[:, :-1], y2[:, :-1], atol=1e-5))
n_params = sum(p.numel() for p in kda.parameters())
print(f"kda params: {n_params:,}")
moe = MoE(cfg)
y = moe(x)
print("moe out:", y.shape)
print("tokens routed:", int(moe.expert_counts.sum()))
w, idx = moe.route(x.reshape(-1, cfg.hidden_size))
groups = idx // (cfg.num_experts // cfg.num_expert_groups)
used = torch.zeros(idx.shape[0], cfg.num_expert_groups).scatter_(1, groups, 1.0).sum(-1)
print("max groups per token:", int(used.max()))
print("weight sums:", [round(v, 3) for v in w.sum(-1)[:4].tolist()])
n_params = sum(p.numel() for p in moe.parameters())
print(f"moe params: {n_params:,}")
moe.capacity_factor = 100.0
ref = moe.forward_reference(x)
fast = moe(x)
print("moe fast matches reference:", torch.allclose(ref, fast, atol=1e-5))
hc = HyperConnection(cfg)
streams = x.unsqueeze(2).expand(-1, -1, cfg.hc_mult, -1).contiguous()
out = hc(streams, lambda z: torch.zeros_like(z))
print("hc out:", out.shape)
print("hc identity:", torch.allclose(out, streams, atol=1e-4))
m = sinkhorn(torch.randn(5, 4, 4), cfg.hc_sinkhorn_iters)
print("row sums:", [round(v, 3) for v in m.sum(-1)[0].tolist()])
print("col sums:", [round(v, 3) for v in m.sum(-2)[0].tolist()])
if torch.cuda.is_available():
lg = torch.randn(4, 2048, 4, 4, device="cuda")
same = torch.allclose(sinkhorn(lg, 20), sinkhorn_fast(lg, 20), atol=1e-5)
print("compiled sinkhorn matches:", same)
n_params = sum(p.numel() for p in hc.parameters())
print(f"hc params: {n_params:,}")
ar = AttnRes(cfg)
snaps = [torch.randn(2, 16, cfg.hc_mult, cfg.hidden_size) for _ in range(3)]
out = ar(snaps)
print("attnres out:", out.shape)
print("attnres uniform at init:", torch.allclose(out, torch.stack(snaps).mean(0), atol=1e-5))
with torch.no_grad():
ar.query.copy_(ar.norm(snaps[0].mean(dim=2))[0, 0] * 10)
w_first = torch.softmax(torch.einsum("kbsd,d->kbs", ar.norm(torch.stack(snaps).mean(3)), ar.query), dim=0)
print("attnres can focus:", round(w_first[0, 0, 0].item(), 3))
n_params = sum(p.numel() for p in ar.parameters())
print(f"attnres params: {n_params:,}")
eng = Engram(cfg)
ids = torch.randint(0, cfg.vocab_size, (2, 16))
out = eng(x, ids)
print("engram out:", out.shape)
ids2 = ids.clone()
ids2[:, -1] = (ids2[:, -1] + 1) % cfg.vocab_size
out2 = eng(x, ids2)
print("engram causal:", torch.allclose(out[:, :-1], out2[:, :-1], atol=1e-5))
rep = torch.tensor([[5, 6, 7, 1, 5, 6, 7]])
rows = eng.hash_ids(rep)
print("same trigram, same rows:", torch.equal(rows[0, 2], rows[0, 6]))
n_params = sum(p.numel() for p in eng.parameters())
print(f"engram params: {n_params:,}")
model = MiniMeridian(cfg)
ids = torch.randint(0, cfg.vocab_size, (2, 16))
targets = torch.randint(0, cfg.vocab_size, (2, 16))
logits, loss = model(ids, targets)
print("model logits:", logits.shape)
print("initial loss:", round(loss.item(), 3))
loss.backward()
missing = [n for n, p in model.named_parameters() if p.grad is None]
print("params without grads:", len(missing))
total = sum(p.numel() for p in model.parameters())
print(f"total params: {total:,}")