andyoneal's picture
Ecce Vectors: 29 control vectors for Gemma 4 31B roleplay fine-tunes, with metrics, blind-test results, examples and the fitting code
e51b495 verified
Raw History Blame Contribute Delete
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