File size: 4,958 Bytes
e51b495
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
"""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