geolip-bytelex / proto_frame /proto_train_frame.py
AbstractPhil's picture
frame prototype v0.1: byte-state arm beats index twin (digits 0.275 vs 0.045, twin has only 200 keys for 999 states); RSA says the codebook kept C ancestry; entry channel real vs failing random control
3dcca87 verified
Raw
History Blame Contribute Delete
13.2 kB
"""Frame prototype v0.1 — training runner (phases 1, 2, 2b + twin).
Implements the v0.1 spec (plan: 2026-08-18_lexicon_translation_
learning_DRAFT.md). Byte-state arm: E_C init from C's own byte-
structural features; teacher aligns TO the codebook (Route B);
logit-adjusted cosine InfoNCE (prior in frozen bias, never in
geometry); student phase against frozen E_C with unit-sphere MSE.
Index-space twin: same losses/budget, codebook keyed by student
anchor token id, state posterior via P(state|key). Gauges: per-class
identification (prior-free + prior-added), masked-reading delta,
teacher ceiling, oracle-codebook control, parity baseline, drift,
RSA frame attribution, gallery decay, counterfactual entry.
"""
import json
import sys
import zlib
sys.path.insert(0, r"E:\mirel\geolip-bytelex")
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
import numpy as np
import torch
import torch.nn.functional as F
import transformers
from transformers import AutoTokenizer
from geolip.bytelex.frame import (apply_whitening, fit_whitening,
gallery_decay, split_sites,
state_byte_features, procrustes,
top_k_accuracy)
D = r"E:\mirel\data\bytelex\proto_frame"
WORDS = r"E:\mirel\data\bytelex\words_of_C.json"
DEV = "cuda"
SEED = zlib.crc32(b"frame-proto-v01") & 0xFFFFFFFF
EPOCHS_T, EPOCHS_S, BS = 6, 6, 512
torch.manual_seed(SEED)
states = json.load(open(WORDS, encoding="utf-8"))
t5_walk = json.load(open(rf"{D}\t5_walk_of_C.json", encoding="utf-8"))
dump = np.load(rf"{D}\frame_dump_v2.npz")
ctx = np.load(rf"{D}\ctx_profiles.npz")
sid = dump["sid"].astype(np.int64)
KK = dump["KK"]
ok = KK[:, 0] > 0
NS = 999
# walk-based class axis (frame-purity fix 7a)
def cls_of(s):
t = s["text"]
if t.isdigit():
return "digit"
if t[0].isupper():
return "Name"
return "word"
CLS = np.array([cls_of(s) for s in states])
site_cls = CLS[sid]
sp = split_sites(sid[ok], seed=SEED)
gix = {k: np.flatnonzero(ok)[v] for k, v in sp.items()}
print(f"[F] sites fit/train/eval = "
f"{len(gix['fit'])}/{len(gix['train'])}/{len(gix['eval'])}",
flush=True)
# ---------- Phase 0 stats (fp64, fit/train only)
H_B = dump["H_B"][:, 0].astype(np.float64) # byte-anchored first
H_T = dump["H_T"][:, 0].astype(np.float64)
muB, wB = fit_whitening(H_B[gix["fit"]])
muT, wT = fit_whitening(H_T[gix["fit"]])
ZB = apply_whitening(H_B, muB, wB)
ZT = apply_whitening(H_T, muT, wT)
evB = np.linalg.eigvalsh(np.cov(H_B[gix["fit"]] - muB, rowvar=False))
prB = float((evB.sum() ** 2) / (evB ** 2).sum() / len(evB))
# byte-structural codebook init (PCA-whiten fp64 -> 256, unit rows)
feats = state_byte_features([s["text"].encode() for s in states],
ctx["prev_ctx"], ctx["next_ctx"])
fc = feats - feats.mean(0)
u, sv, vt = np.linalg.svd(fc.astype(np.float64), full_matrices=False)
E0 = u[:, :256] * 1.0 # decorrelated rows
E0 = E0 / np.linalg.norm(E0, axis=1, keepdims=True)
# teacher whitened state-means on TRAIN
def state_means(z, ix):
m = np.zeros((NS, z.shape[1]))
for s in range(NS):
r = ix[sid[ix] == s]
if len(r):
m[s] = z[r].mean(0)
return m
mT_train = state_means(ZB, gix["train"])
R_T = procrustes(mT_train, E0) # 768 -> 256, fp64
# priors (frozen, TRAIN)
cnt = np.bincount(sid[gix["train"]], minlength=NS).astype(np.float64)
b_prior = np.log(np.maximum(cnt, 0.5) / cnt.sum())
# ---------- torch setup
tZB = torch.tensor(ZB, dtype=torch.float32, device=DEV)
tZT = torch.tensor(ZT, dtype=torch.float32, device=DEV)
tsid = torch.tensor(sid, device=DEV)
E_C = torch.nn.Parameter(torch.tensor(E0, dtype=torch.float32,
device=DEV))
W_Tm = torch.nn.Parameter(torch.tensor(R_T, dtype=torch.float32,
device=DEV))
s_T = torch.nn.Parameter(torch.tensor(10.0, device=DEV))
bp = torch.tensor(b_prior, dtype=torch.float32, device=DEV)
def id_logits(z, W, E, s, prior):
zf = F.normalize(z @ W, dim=-1)
lg = s.clamp(1, 100) * (zf @ F.normalize(E, dim=-1).T)
return lg + bp if prior else lg
def train_phase(params, loss_fn, ixs, epochs, tag):
opt = torch.optim.Adam(params, lr=1e-3, weight_decay=0.0)
rng = np.random.default_rng(SEED + 7)
for ep in range(epochs):
order = rng.permutation(ixs)
tot = nb = 0
for i in range(0, len(order), BS):
b = torch.tensor(order[i:i + BS], device=DEV)
opt.zero_grad(set_to_none=True)
l = loss_fn(b)
l.backward()
opt.step()
with torch.no_grad():
E_C.data = F.normalize(E_C.data, dim=-1)
tot += float(l)
nb += 1
print(f"[F {tag}] ep{ep} loss={tot / nb:.4f}", flush=True)
# ---------- Phase 1: teacher anchors (L_id^T)
def loss_T(b):
lg = id_logits(tZB[b], W_Tm, E_C, s_T, True)
return F.cross_entropy(lg, tsid[b])
E_before = E_C.detach().cpu().numpy().copy()
train_phase([W_Tm, E_C, s_T], loss_T, gix["train"], EPOCHS_T, "P1")
E_C.requires_grad_(False)
s_T.requires_grad_(False)
# drift by count decile + dead rows
drift = np.linalg.norm(E_C.detach().cpu().numpy() - E_before, axis=1)
dec = np.digitize(cnt, np.quantile(cnt, np.linspace(0, 1, 11)[1:-1]))
drift_dec = [round(float(drift[dec == d].mean()), 4) for d in range(10)]
# ---------- Phase 2: student aligns to frozen center
mS_train = state_means(ZT, gix["train"])
R_S0 = procrustes(mS_train, E_C.detach().cpu().numpy().astype(
np.float64))
W_Sm = torch.nn.Parameter(torch.tensor(R_S0, dtype=torch.float32,
device=DEV))
with torch.no_grad():
zT_frame = F.normalize(tZB @ W_Tm, dim=-1) # frozen teacher
def eval_id(z, W, ix, prior, k=1):
with torch.no_grad():
lg = id_logits(z[torch.tensor(ix, device=DEV)], W, E_C, s_T,
prior).cpu().numpy()
return top_k_accuracy(lg, sid[ix], k), lg
parity1, _ = eval_id(tZT, W_Sm, gix["eval"], False)
def loss_S(b):
lg = id_logits(tZT[b], W_Sm, E_C, s_T, True)
zs = F.normalize(tZT[b] @ W_Sm, dim=-1)
return (F.cross_entropy(lg, tsid[b])
+ ((zs - zT_frame[b]) ** 2).sum(-1).mean())
train_phase([W_Sm], loss_S, gix["train"], EPOCHS_S, "P2")
# ---------- gauges
led = {"_env": {"transformers": transformers.__version__,
"torch": torch.__version__, "seed": int(SEED)},
"splits": {k: int(len(v)) for k, v in gix.items()},
"whitening": {"participation_ratio_B": round(prB, 4)},
"drift_by_decile": drift_dec,
"dead_rows": int((drift < 1e-4).sum()),
"parity_baseline_eval_top1": round(float(parity1), 4)}
ev = gix["eval"]
accs = {}
for tag, z, W in (("teacher", tZB, W_Tm), ("student", tZT, W_Sm)):
for prior in (False, True):
a1, lg = eval_id(z, W, ev, prior)
a5 = top_k_accuracy(lg, sid[ev], 5)
key = f"{tag}_{'prior' if prior else 'balanced'}"
accs[key] = {"top1": round(a1, 4), "top5": round(a5, 4)}
per = {}
for c in ("word", "Name", "digit"):
m = site_cls[ev] == c
if m.any():
per[c] = {"n": int(m.sum()),
"top1": round(top_k_accuracy(
lg[m], sid[ev][m]), 4)}
accs[key]["by_class"] = per
led["identification"] = accs
# masked-reading delta (control): masked readouts through same maps
MB = apply_whitening(dump["M_B"].astype(np.float64), muB, wB)
MT = apply_whitening(dump["M_T"].astype(np.float64), muT, wT)
tMB = torch.tensor(MB, dtype=torch.float32, device=DEV)
tMT = torch.tensor(MT, dtype=torch.float32, device=DEV)
delta = {}
for tag, zu, zm, W in (("teacher", tZB, tMB, W_Tm),
("student", tZT, tMT, W_Sm)):
au, lgu = eval_id(zu, W, ev, False)
am, lgm = eval_id(zm, W, ev, False)
per = {}
for c in ("word", "Name", "digit"):
m = site_cls[ev] == c
if m.any():
per[c] = round(top_k_accuracy(lgu[m], sid[ev][m])
- top_k_accuracy(lgm[m], sid[ev][m]), 4)
delta[tag] = {"overall": round(au - am, 4), "by_class": per}
led["reading_delta"] = delta
# oracle-codebook control (student vs never-trained student means)
mS_unit = mS_train / np.maximum(
np.linalg.norm(mS_train, axis=1, keepdims=True), 1e-9)
scr = ZT[ev] @ mS_unit.T
led["oracle_codebook_student_top1"] = round(
top_k_accuracy(scr, sid[ev]), 4)
# RSA frame attribution: whose geometry is final E_C?
Ef = F.normalize(E_C, dim=-1).cpu().numpy().astype(np.float64)
def rsa(a, b, rng):
iu = np.triu_indices(NS, 1)
pick = rng.choice(len(iu[0]), size=100_000, replace=False)
va = (a @ a.T)[iu][pick]
vb = (b @ b.T)[iu][pick]
ra, rb = np.argsort(np.argsort(va)), np.argsort(np.argsort(vb))
return float(np.corrcoef(ra, rb)[0, 1])
rr = np.random.default_rng(SEED + 13)
mB_unit = mT_train / np.maximum(
np.linalg.norm(mT_train, axis=1, keepdims=True), 1e-9)
led["rsa"] = {"E_C_vs_byte_features": round(rsa(Ef, E0, rr), 4),
"E_C_vs_teacher_means": round(rsa(Ef, mB_unit, rr), 4)}
# gallery decay (prior-free, student)
_, lgS = eval_id(tZT, W_Sm, ev, False)
sub = rr.choice(len(ev), size=min(2000, len(ev)), replace=False)
led["gallery_decay_student"] = gallery_decay(lgS[sub], sid[ev][sub],
seed=SEED + 17)
# ---------- index-space twin (same losses/budget, key = t5 anchor id)
tkA = AutoTokenizer.from_pretrained("google/flan-t5-small")
lines = open(r"E:\mirel\data\bytelex\codex_v1.txt",
"rb").read().decode("ascii").split("\n")
anchor_id = np.full(len(sid), -1, dtype=np.int64)
by_line = {}
for k in range(len(sid)):
by_line.setdefault(int(dump["line"][k]), []).append(k)
for li, ks in by_line.items():
e = tkA(lines[li], add_special_tokens=False,
return_offsets_mapping=True)
off = e["offset_mapping"]
for k in ks:
lo, hi = int(dump["lo"][k]), int(dump["hi"][k])
ix = [i for i, (s, t) in enumerate(off)
if t > s and s < hi and t > lo]
if ix:
anchor_id[k] = e["input_ids"][ix[0]]
keys, key_inv = np.unique(anchor_id[anchor_id >= 0],
return_inverse=False), None
key_of = {int(a): i for i, a in enumerate(keys)}
ksite = np.array([key_of.get(int(a), -1) for a in anchor_id])
NK = len(keys)
# P(state|key) from TRAIN
post = np.zeros((NK, NS))
for k in gix["train"]:
if ksite[k] >= 0:
post[ksite[k], sid[k]] += 1
post = post / np.maximum(post.sum(1, keepdims=True), 1)
mK_train = np.zeros((NK, 768))
for kk in range(NK):
r = gix["train"][ksite[gix["train"]] == kk]
if len(r):
mK_train[kk] = ZB[r].mean(0)
uk, sk, vk = np.linalg.svd(
(mK_train - mK_train.mean(0)).astype(np.float64),
full_matrices=False)
EK0 = uk[:, :256]
EK0 = EK0 / np.maximum(np.linalg.norm(EK0, axis=1, keepdims=True),
1e-9)
E_K = torch.nn.Parameter(torch.tensor(EK0, dtype=torch.float32,
device=DEV))
R_K = procrustes(mK_train, EK0) # same-quality init
W_Ki = torch.nn.Parameter(torch.tensor(R_K, dtype=torch.float32,
device=DEV))
s_K = torch.nn.Parameter(torch.tensor(10.0, device=DEV))
tks = torch.tensor(ksite, device=DEV)
cntK = np.bincount(ksite[gix["train"]][ksite[gix["train"]] >= 0],
minlength=NK).astype(np.float64)
bK = torch.tensor(np.log(np.maximum(cntK, .5) / cntK.sum()),
dtype=torch.float32, device=DEV)
trK = gix["train"][ksite[gix["train"]] >= 0]
def loss_K(b):
zf = F.normalize(tZB[b] @ W_Ki, dim=-1)
lg = s_K.clamp(1, 100) * (zf @ F.normalize(E_K, dim=-1).T) + bK
return F.cross_entropy(lg, tks[b])
train_phase([W_Ki, E_K, s_K], loss_K, trK, EPOCHS_T, "TWIN")
evK = ev[ksite[ev] >= 0]
with torch.no_grad():
zf = F.normalize(tZB[torch.tensor(evK, device=DEV)] @ W_Ki, -1)
lgK = (s_K.clamp(1, 100) * zf @ F.normalize(E_K, -1).T).cpu().numpy()
state_scores = np.exp(lgK - lgK.max(1, keepdims=True)) @ post
twin = {"overall": round(top_k_accuracy(state_scores, sid[evK]), 4),
"n_keys": int(NK)}
for c in ("word", "Name", "digit"):
m = site_cls[evK] == c
if m.any():
twin[c] = {"n": int(m.sum()),
"top1": round(top_k_accuracy(state_scores[m],
sid[evK][m]), 4)}
led["index_twin_state_id"] = twin
with open(rf"{D}\frame_ledger_v01.json", "w", encoding="utf-8") as f:
json.dump(led, f, indent=1)
np.savez(rf"{D}\frame_anchors_v01.npz",
E_C=E_C.detach().cpu().numpy(),
W_T=W_Tm.detach().cpu().numpy(),
W_S=W_Sm.detach().cpu().numpy(),
s=float(s_T), b_prior=b_prior, E0=E0)
print(json.dumps(led["identification"], indent=1)[:1500], flush=True)
print("[F] TWIN:", json.dumps(twin), flush=True)
print("[F] RSA:", json.dumps(led["rsa"]),
"| reading_delta:", json.dumps(led["reading_delta"]), flush=True)
print("[F] PHASES 1-2 + TWIN COMPLETE", flush=True)