Download model.py from Circuits-V2/Meridian-Tiny: direct link, hf CLI and curl.
- Browser
- Download file 23.9 kB
-
https://huggingface.co/Circuits-V2/Meridian-Tiny/resolve/main/model.py
- Command line
-
hf download hf://Circuits-V2/Meridian-Tiny/model.py
-
curl -L -o model.py https://huggingface.co/Circuits-V2/Meridian-Tiny/resolve/main/model.py
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 | |
| 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) | |
| 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:,}") |