Download ecce_vector/model.py from andyoneal/ecce-vectors: direct link, hf CLI and curl.
- Browser
- Download file 4.96 kB
-
https://huggingface.co/andyoneal/ecce-vectors/resolve/main/ecce_vector/model.py
- Command line
-
hf download hf://andyoneal/ecce-vectors/ecce_vector/model.py
-
curl -L -o model.py https://huggingface.co/andyoneal/ecce-vectors/resolve/main/ecce_vector/model.py
4.96 kB
| """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 | |