"""Toy cross-model KV-cache transfer: RoPE strip/reapply, ridge-regression
mapping, a trained low-rank residual, and partial recomputation -- the same
four-rung ladder TrustAI's research notes climb, on synthetic data small
enough to run with the standard library alone.
python3 kv_transfer_lab.py
Two toy "models" share one layer of representation (the shared boundary) and
diverge after it: the source applies a fixed linear transform, the target
applies its own linear transform PLUS a rank-2 bilinear term with no source
counterpart -- a nonlinearity that makes a *linear* map provably unable to
reach perfect agreement, the same ceiling Research Note 004 hit with real
weights. Positions carry a simplified RoPE: a per-position rotation applied
after the shared boundary, which is why it has to be stripped before mapping
and reapplied after.
"""
import random
import math
random.seed(7)
D = 8 # hidden dim (kept small enough to invert by hand-rolled Gauss-Jordan)
N_CALIB = 500 # calibration sequences, matching the FineWeb-Edu-sized sets in the reports
N_EVAL = 100 # held-out evaluation sequences
SEQ_LEN = 32 # positions per sequence
RIDGE_LAMBDA = 1.0
VOCAB = 200 # toy unembedding size, for a top-1 "agreement" metric
# ---- tiny linear algebra: plain lists of lists, no numpy ------------------
def mat_mul(A, B):
n, k, m = len(A), len(B), len(B[0])
return [[sum(A[i][t] * B[t][j] for t in range(k)) for j in range(m)] for i in range(n)]
def transpose(A):
return [list(row) for row in zip(*A)]
def mat_vec(A, v):
return [sum(A[i][j] * v[j] for j in range(len(v))) for i in range(len(A))]
def add(A, B):
return [[A[i][j] + B[i][j] for j in range(len(A[0]))] for i in range(len(A))]
def scale(A, s):
return [[a * s for a in row] for row in A]
def identity(n):
return [[1.0 if i == j else 0.0 for j in range(n)] for i in range(n)]
def solve(A, B):
"""Gauss-Jordan solve of A X = B for square A. A and B are lists of lists;
B may have multiple columns. Returns X with the same shape as B."""
n = len(A)
M = [row[:] + B[i][:] for i, row in enumerate(A)]
bw = len(B[0])
for col in range(n):
piv = max(range(col, n), key=lambda r: abs(M[r][col]))
M[col], M[piv] = M[piv], M[col]
pv = M[col][col]
M[col] = [x / pv for x in M[col]]
for r in range(n):
if r == col:
continue
f = M[r][col]
if f != 0:
M[r] = [M[r][c] - f * M[col][c] for c in range(n + bw)]
return [row[n:] for row in M]
def ridge_fit(X, Y, lam):
"""W minimising ||XW - Y||^2 + lam*||W||^2, solved per-output-column via
the normal equations (X^T X + lam I) W = X^T Y."""
Xt = transpose(X)
XtX = mat_mul(Xt, X)
reg = add(XtX, scale(identity(len(XtX)), lam))
XtY = mat_mul(Xt, Y)
return solve(reg, XtY)
def rand_vec(d, scale_=1.0):
return [random.gauss(0, scale_) for _ in range(d)]
def rand_mat(n, m, scale_=1.0):
return [[random.gauss(0, scale_) for _ in range(m)] for _ in range(n)]
def tanh_v(v):
return [math.tanh(x) for x in v]
# ---- the two toy models -----------------------------------------------------
# Shared boundary: both models produce the same layer-0 output for a given
# input (a stand-in for a shared embedding + early layers). Source and target
# then each apply their OWN layer-1 transform -- different random weights,
# different models -- source purely linear, target linear plus a bilinear
# term (defined below) with no source counterpart. RoPE is a per-position
# rotation in the first three (x, y) coordinate pairs, applied after layer-1.
W_shared = rand_mat(D, D, 0.5)
W_src1 = rand_mat(D, D, 0.5)
W_tgt1 = rand_mat(D, D, 0.5)
UNEMBED = rand_mat(VOCAB, D, 1.0) # toy "read out a token id from a hidden state"
def rope(v, pos, dim_pairs=3):
"""Rotate the first `dim_pairs` (x, y) pairs of v by an angle proportional
to position -- a simplified stand-in for real per-frequency RoPE."""
out = v[:]
for p in range(dim_pairs):
i, j = 2 * p, 2 * p + 1
theta = pos * (0.05 * (p + 1))
c, s = math.cos(theta), math.sin(theta)
x, y = out[i], out[j]
out[i] = c * x - s * y
out[j] = s * x + c * y
return out
def unrope(v, pos, dim_pairs=3):
return rope(v, -pos, dim_pairs)
# The target's transform has a genuine nonlinear component with no source
# counterpart: a fixed rank-2 bilinear term, ((h0.P) elementwise* (h0.Q)) . Cb.
# No linear map from source to target -- ridge or otherwise -- can represent
# it exactly, which is the toy-model version of Research Note 004's "layers
# 1-3 perturbed" nonlinear drift. It is EXACTLY the functional form the
# factorized-quadratic residual (below) is built to fit, so a correctly
# trained rank-2 residual should recover most of it; a plain linear residual,
# stacked on top of a map that already found the best linear fit, cannot
# recover any of it -- try R_TRUE = 2 with a LINEAR-only residual first.
R_TRUE = 2
P_true = rand_mat(R_TRUE, D, 0.6)
Q_true = rand_mat(R_TRUE, D, 0.6)
C_bilin = rand_mat(D, R_TRUE, 0.5)
def bilinear_term(h0):
p = mat_vec(P_true, h0)
q = mat_vec(Q_true, h0)
hh = [p[k] * q[k] for k in range(R_TRUE)]
return mat_vec(C_bilin, hh)
def layer0(x):
return tanh_v(mat_vec(W_shared, x))
def source_layer1(h0, pos):
h1 = mat_vec(W_src1, h0) # source: linear
return rope(h1, pos)
def target_layer1_exact(h0, pos):
h1 = [a + b for a, b in zip(mat_vec(W_tgt1, h0), bilinear_term(h0))]
return rope(h1, pos)
def token_id(v):
scores = mat_vec(UNEMBED, v)
return max(range(VOCAB), key=lambda i: scores[i])
# ---- build calibration and eval sets ---------------------------------------
def make_sequence(seq_len):
"""One synthetic 'prompt': a random walk of embeddings, each carried
through both models' first layer at its own position."""
src_cache, tgt_cache, positions = [], [], []
for pos in range(seq_len):
x = rand_vec(D, 1.0)
h0 = layer0(x)
src_cache.append(source_layer1(h0, pos))
tgt_cache.append(target_layer1_exact(h0, pos))
positions.append(pos)
return src_cache, tgt_cache, positions
def flatten(seqs):
src_rows, tgt_rows = [], []
for src_cache, tgt_cache, positions in seqs:
for s, t, pos in zip(src_cache, tgt_cache, positions):
src_rows.append(unrope(s, pos)) # strip RoPE before mapping
tgt_rows.append(unrope(t, pos))
return src_rows, tgt_rows
calib = [make_sequence(SEQ_LEN) for _ in range(N_CALIB)]
evalset = [make_sequence(SEQ_LEN) for _ in range(N_EVAL)]
calib_src, calib_tgt = flatten(calib)
W_ridge = ridge_fit(calib_src, calib_tgt, RIDGE_LAMBDA) # content-space linear map
# A FACTORIZED QUADRATIC residual on top of the frozen ridge map -- Research
# Note 005's own adapter formula, verbatim:
#
# z -> z_hat = z + ((z.A) elementwise* (z.B)) . C A,B: D->r C: r->D
#
# The point of the elementwise product is that it is the cheapest possible
# NONLINEAR function of z. A plain linear residual B(A(z)) is still linear in
# z, and ridge regression already found the best linear map -- so a linear
# residual has, by construction, nothing left to correct on the SAME
# objective ridge was fit to. Try that first (see the pitfall below) before
# running this cell, and it reproduces exactly that null result. Trained by
# plain SGD on the ridge map's residual error; the real adapters are trained
# by backpropagating through a multistep rollout, which the math lab in this
# level connects back to a single-step loss.
def ridge_predict(src_vec):
return mat_vec(transpose(W_ridge), src_vec)
R = 2
resid_pairs = [(s, [t[d] - ridge_predict(s)[d] for d in range(D)]) for s, t in zip(calib_src, calib_tgt)]
A_q = rand_mat(R, D, 0.2) # D -> R
B_q = rand_mat(R, D, 0.2) # D -> R
C_q = rand_mat(D, R, 0.2) # R -> D
LR, EPOCHS = 0.003, 25
train_set = resid_pairs[:4000] # subsample for pure-Python training speed
clip = lambda x: max(-1.0, min(1.0, x))
for epoch in range(EPOCHS):
random.shuffle(train_set)
for s, r in train_set:
a = mat_vec(A_q, s) # R-vector, z.A
b = mat_vec(B_q, s) # R-vector, z.B
h = [a[k] * b[k] for k in range(R)] # elementwise product
pred = mat_vec(C_q, h) # D-vector
err = [pred[d] - r[d] for d in range(D)]
for d in range(D):
for k in range(R):
C_q[d][k] -= LR * clip(err[d] * h[k])
dh = [sum(C_q[d][k] * err[d] for d in range(D)) for k in range(R)]
da = [dh[k] * b[k] for k in range(R)]
db = [dh[k] * a[k] for k in range(R)]
for k in range(R):
for i in range(D):
A_q[k][i] -= LR * clip(da[k] * s[i])
B_q[k][i] -= LR * clip(db[k] * s[i])
def residual_predict(src_vec):
a = mat_vec(A_q, src_vec)
b = mat_vec(B_q, src_vec)
h = [a[k] * b[k] for k in range(R)]
return mat_vec(C_q, h)
# ---- evaluation --------------------------------------------------------------
def evaluate(name, predict_fn):
agree, total = 0, 0
sq_err, tgt_sq = 0.0, 0.0
for src_cache, tgt_cache, positions in evalset:
for s_r, t_r, pos in zip(src_cache, tgt_cache, positions):
s = unrope(s_r, pos)
pred_content = predict_fn(s)
pred = rope(pred_content, pos) # reapply target position
t = t_r
if token_id(pred) == token_id(t):
agree += 1
total += 1
sq_err += sum((a - b) ** 2 for a, b in zip(pred, t))
tgt_sq += sum(b * b for b in t)
top1 = agree / total
rel_err = math.sqrt(sq_err / tgt_sq)
return top1, rel_err
results = {}
results['direct (no mapping)'] = evaluate('direct', lambda s: s)
results['ridge map'] = evaluate('ridge', ridge_predict)
results['ridge + rank-2 residual'] = evaluate(
'ridge+resid', lambda s: [a + b for a, b in zip(ridge_predict(s), residual_predict(s))])
# partial recompute: given the source's content-space vector, invert the
# SOURCE's own layer-1 map to recover h0 exactly (source is purely linear, so
# this inverse is exact), then apply the target's own layer-1 exactly (linear
# + bilinear) -- no mapping error anywhere, by construction.
def evaluate_recompute():
agree, total = 0, 0
sq_err, tgt_sq = 0.0, 0.0
W_src1_inv = solve(W_src1, identity(D))
for src_cache, tgt_cache, positions in evalset:
for s_r, t_r, pos in zip(src_cache, tgt_cache, positions):
s = unrope(s_r, pos) # source content vector = W_src1 @ h0
h0 = mat_vec(W_src1_inv, s) # recover the shared boundary exactly
pred_content = [a + b for a, b in zip(mat_vec(W_tgt1, h0), bilinear_term(h0))] # target's OWN exact layer
pred = rope(pred_content, pos)
t = t_r
if token_id(pred) == token_id(t):
agree += 1
total += 1
sq_err += sum((a - b) ** 2 for a, b in zip(pred, t))
tgt_sq += sum(b * b for b in t)
return agree / total, math.sqrt(sq_err / tgt_sq)
results['partial recompute'] = evaluate_recompute()
print(f"{'method':<26}{'top-1 agreement':>18}{'rel. error':>14}")
for name, (top1, rel_err) in results.items():
print(f"{name:<26}{top1*100:>17.1f}%{rel_err:>14.4f}")