"""Gemma 4 31B decoder layers on MLX, one at a time, plus the output head and the scoring arithmetic. Nothing here builds a whole model. A layer is built from its stored weights (upcast to fp32), run over the corpus, and dropped, so a fit or a score needs about 10 GB of unified memory whatever the checkpoint size. Every projection is wrapped in a `Side`, which can add a LoRA branch (for fitting an adapter) or, when scoring, compute several vector-steered streams as differences from one recorded stock stream. """ from __future__ import annotations import json from pathlib import Path import mlx.core as mx import numpy as np from mlx_lm.models import gemma4_text from .io import N_LAYERS F32 = mx.float32 PROJECTIONS = ("self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj", "self_attn.o_proj", "mlp.gate_proj", "mlp.up_proj", "mlp.down_proj") def model_args(base_dir: str | Path) -> gemma4_text.ModelArgs: cfg = json.loads((Path(base_dir) / "config.json").read_text()) args = gemma4_text.ModelArgs.from_dict(cfg.get("text_config", cfg)) assert args.num_hidden_layers == N_LAYERS and not getattr(args, "enable_moe_block", False), \ "this code is written for the dense 60-layer gemma-4-31B-it (the 26B-A4B shares its model_type)" return args class Side: """Stands in for one projection. mode None: W x, plus the side branch f if set mode "rec": the same, recording input and output (the stock stream) mode "diff": the input holds k streams over the recorded sequences, and W x = y_s + fp16(W) (x - x_s) f: (A, sB), a LoRA branch applied as (x A^T) sB^T """ def __init__(self, inner, diff: bool = True): self.inner, self.f, self.mode = inner, None, None self.xs = self.ys = None self.W16 = inner.weight.astype(mx.float16) if diff else None def __call__(self, x): if self.mode == "diff": k = x.shape[0] // self.xs.shape[0] dx = (x.reshape(k, *self.xs.shape) - self.xs[None]).reshape(x.shape) y = ((dx.astype(mx.float16) @ self.W16.T).astype(F32).reshape(k, *self.ys.shape) + self.ys[None]).reshape(*x.shape[:-1], self.ys.shape[-1]) else: y = self.inner(x) if self.mode == "rec": self.xs, self.ys = x, y if self.f is None: return y A, sB = self.f return y + (x @ A.T) @ sB.T def make_layer(args, L: int, weights: dict[str, np.ndarray], diff: bool = True): """An mlx_lm Gemma 4 decoder layer with the given weights, every projection wrapped in a Side.""" layer = gemma4_text.DecoderLayer(args, L) layer.load_weights([(k, mx.array(v)) for k, v in weights.items()], strict=True) mx.eval(layer.parameters()) sides = {} for fam in PROJECTIONS: parent, attr = fam.split(".") mod = getattr(layer, parent) if hasattr(mod, attr): # the ten full-attention layers have no v_proj sides[fam] = Side(getattr(mod, attr), diff) setattr(mod, attr, sides[fam]) return layer, sides def embed(stock, ids: np.ndarray, width: int): """Scaled input embeddings (bf16 table, times sqrt(d) in fp32).""" table = mx.array(stock.get("embed_tokens.weight")).astype(mx.bfloat16) return table[mx.array(ids)].astype(F32) * (width ** 0.5) def masked_mean(d, mask): """Mean of d (B, T, D) over the positions where mask (B, T) is 1.""" return (d * mask[..., None]).sum((0, 1)) / mask.sum() def gather_rows(x, rows): """The scored positions of x (B, T, D) as one (N, D) array.""" return mx.concatenate([x[i, a:b] for i, a, b in rows], 0) def head(checkpoint): """(final norm weight, tied embedding) of a checkpoint as mx arrays; the embedding stays bf16.""" assert not checkpoint.has("lm_head.weight"), "untied output head: not handled" norm = mx.array(checkpoint.get("norm.weight")) emb = mx.array(checkpoint.get("embed_tokens.weight")).astype(mx.bfloat16) mx.eval(norm, emb) return norm, emb def log_probs(args, x, norm, emb): """Next-token log-probabilities for final hidden states x (N, D): final norm, tied head, softcap.""" xn = mx.fast.rms_norm(x, norm, args.rms_norm_eps) z = mx.concatenate([xn @ emb[v:v + 32768].astype(F32).T for v in range(0, emb.shape[0], 32768)], -1) cap = args.final_logit_softcapping z = mx.tanh(z / cap) * cap lp = z - mx.logsumexp(z, -1, keepdims=True) mx.eval(lp) return lp def compare(target, candidate, spans): """Mean over sequences of KL(target || candidate), top-1 agreement, and per-sequence KL.""" kl = np.array((mx.exp(target) * (target - candidate)).sum(-1)) agree = np.array(target.argmax(-1)) == np.array(candidate.argmax(-1)) per_seq = [float(kl[a:b].mean()) for a, b in spans] return float(np.mean(per_seq)), float(np.mean([agree[a:b].mean() for a, b in spans])), per_seq