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:,}")