geolip-bytelex / proto_frame /proto_phase0.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
57e06d7 verified
Raw
History Blame Contribute Delete
8.53 kB
"""Frame prototype Phase 0 — instruments + full extraction (v0.1).
1. T5 walk table of C (instrument-first: entry loss needs whole-walk
census; only the wordpiece walk existed).
2. Boundary-context profiles per state (prev/next byte distributions
at attested sites) for the byte-structural codebook init.
3. Extraction v2: ALL sites, no cap. Unmasked anchor readouts
(byte-anchored first subtoken + last subtoken, k per side) AND
masked-span readouts (bert: [MASK]*k, read first mask; t5:
<extra_id_0> replacing span, read sentinel in encoder) — the
reading-vs-guessing control. CUDA, <18GB.
"""
import json
import sys
sys.path.insert(0, r"E:\mirel\geolip-bytelex")
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
import numpy as np
import torch
import transformers
from transformers import AutoTokenizer, AutoModel, T5EncoderModel
import transformers.utils.logging as hlog
hlog.set_verbosity_error()
CODEX = r"E:\mirel\data\bytelex\codex_v1.txt"
WORDS = r"E:\mirel\data\bytelex\words_of_C.json"
OUTDIR = r"E:\mirel\data\bytelex\proto_frame"
SEPS = set(b" \t\n\r.,;:!?\"'()[]{}")
L_BERT, L_T5 = 8, 6
DEV = "cuda"
BATCH = 256
states = json.load(open(WORDS, encoding="utf-8"))
sid_of = {s["text"]: i for i, s in enumerate(states)}
tkB = AutoTokenizer.from_pretrained("bert-base-uncased")
tkA = AutoTokenizer.from_pretrained("google/flan-t5-small")
# ---- 1. t5 walk table of C
t5_walk = []
for s in states:
ids = tkA(s["text"], add_special_tokens=False)["input_ids"]
toks = tkA.convert_ids_to_tokens(ids)
t5_walk.append({"text": s["text"], "count": s["count"],
"k": len(ids), "seg": "|".join(toks),
"ids": ids,
"whole": len(ids) == 1})
n_whole = sum(w["whole"] for w in t5_walk)
with open(rf"{OUTDIR}\t5_walk_of_C.json", "w", encoding="utf-8") as f:
json.dump(t5_walk, f, indent=0)
print(f"[P0] t5 walk of C: {n_whole}/999 whole", flush=True)
# ---- sites (ALL, no cap) + 2. boundary-context profiles
lines = open(CODEX, "rb").read().decode("ascii").split("\n")
sites = []
prev_ctx = np.zeros((999, 256))
next_ctx = np.zeros((999, 256))
for li, ln in enumerate(lines):
if not ln:
continue
raw = ln.encode("ascii") + b" "
lo = None
for j, ch in enumerate(raw):
if ch in SEPS:
if lo is not None:
sid = sid_of.get(ln[lo:j])
if sid is not None:
sites.append((li, lo, j, sid))
prev_ctx[sid, raw[lo - 1] if lo else 32] += 1
next_ctx[sid, ch] += 1
lo = None
elif lo is None:
lo = j
np.savez_compressed(rf"{OUTDIR}\ctx_profiles.npz",
prev_ctx=prev_ctx, next_ctx=next_ctx)
print(f"[P0] {len(sites)} sites (uncapped), ctx profiles saved",
flush=True)
mB = AutoModel.from_pretrained("bert-base-uncased").to(DEV).eval()
mT = T5EncoderModel.from_pretrained("google/flan-t5-small").to(DEV).eval()
embT = mT.get_input_embeddings()
SENT = tkA.convert_tokens_to_ids("<extra_id_0>")
MASK = tkB.mask_token_id
N = len(sites)
H_B = np.zeros((N, 2, 768), dtype=np.float16) # first,last
H_T = np.zeros((N, 2, 512), dtype=np.float32)
M_B = np.zeros((N, 768), dtype=np.float16) # masked readout
M_T = np.zeros((N, 512), dtype=np.float32)
E_T = np.zeros((N, 512), dtype=np.float32) # t5 input emb (first)
KK = np.zeros((N, 2), dtype=np.int16) # k_B, k_T
# ---- unmasked pass, batched by line
by_line = {}
for k, (li, lo, hi, sid) in enumerate(sites):
by_line.setdefault(li, []).append(k)
line_ids = sorted(by_line)
tokcacheB, tokcacheT = {}, {}
with torch.no_grad():
for bs in range(0, len(line_ids), BATCH):
chunk = line_ids[bs:bs + BATCH]
texts = [lines[li] for li in chunk]
eb = tkB(texts, return_offsets_mapping=True, padding=True,
return_tensors="pt")
et = tkA(texts, return_offsets_mapping=True, padding=True,
return_tensors="pt")
hb = mB(input_ids=eb["input_ids"].to(DEV),
attention_mask=eb["attention_mask"].to(DEV),
output_hidden_states=True).hidden_states[L_BERT].cpu()
ht = mT(input_ids=et["input_ids"].to(DEV),
attention_mask=et["attention_mask"].to(DEV),
output_hidden_states=True).hidden_states[L_T5].cpu()
em = embT(et["input_ids"].to(DEV)).cpu()
for r, li in enumerate(chunk):
offB = eb["offset_mapping"][r].tolist()
offT = et["offset_mapping"][r].tolist()
idsB = eb["input_ids"][r].tolist()
idsT = et["input_ids"][r].tolist()
tokcacheB[li] = (idsB, offB)
tokcacheT[li] = (idsT, offT)
for k in by_line[li]:
_, lo, hi, sid = sites[k]
ixB = [i for i, (s, t) in enumerate(offB)
if t > s and s < hi and t > lo]
ixT = [i for i, (s, t) in enumerate(offT)
if t > s and s < hi and t > lo]
if not ixB or not ixT:
KK[k] = (0, 0)
continue
H_B[k, 0] = hb[r, ixB[0]].numpy()
H_B[k, 1] = hb[r, ixB[-1]].numpy()
H_T[k, 0] = ht[r, ixT[0]].numpy()
H_T[k, 1] = ht[r, ixT[-1]].numpy()
E_T[k] = em[r, ixT[0]].numpy()
KK[k] = (len(ixB), len(ixT))
if (bs // BATCH) % 5 == 0:
print(f"[P0-unmasked] {bs}/{len(line_ids)} lines", flush=True)
print("[P0] unmasked pass done", flush=True)
# ---- masked pass: one sequence per SITE, id-spliced
def masked_batchB(ks):
seqs, poss = [], []
for k in ks:
li, lo, hi, sid = sites[k]
idsB, offB = tokcacheB[li]
ix = [i for i, (s, t) in enumerate(offB)
if t > s and s < hi and t > lo]
pre = [idsB[i] for i, (s, t) in enumerate(offB)
if t > s and t <= lo]
post = [idsB[i] for i, (s, t) in enumerate(offB)
if t > s and s >= hi]
seqs.append([tkB.cls_token_id] + pre + [MASK] * max(len(ix), 1)
+ post + [tkB.sep_token_id])
poss.append(1 + len(pre))
return seqs, poss
def masked_batchT(ks):
seqs, poss = [], []
for k in ks:
li, lo, hi, sid = sites[k]
ln = lines[li]
src = ln[:lo] + "<extra_id_0>" + ln[hi:]
ids = tkA(src, add_special_tokens=False)["input_ids"]
try:
p = ids.index(SENT)
except ValueError:
p = 0
seqs.append(ids)
poss.append(p)
return seqs, poss
def run_masked(model, seqs, poss, layer, pad_id):
mx = max(len(s) for s in seqs)
ids = torch.full((len(seqs), mx), pad_id, dtype=torch.long)
att = torch.zeros((len(seqs), mx), dtype=torch.long)
for i, s in enumerate(seqs):
ids[i, :len(s)] = torch.tensor(s)
att[i, :len(s)] = 1
with torch.no_grad():
h = model(input_ids=ids.to(DEV), attention_mask=att.to(DEV),
output_hidden_states=True).hidden_states[layer].cpu()
return h[torch.arange(len(seqs)), torch.tensor(poss)]
order = [k for k in range(N) if KK[k, 0] > 0]
for bs in range(0, len(order), BATCH):
ks = order[bs:bs + BATCH]
sq, ps = masked_batchB(ks)
out = run_masked(mB, sq, ps, L_BERT, tkB.pad_token_id)
for j, k in enumerate(ks):
M_B[k] = out[j].numpy()
sq, ps = masked_batchT(ks)
out = run_masked(mT, sq, ps, L_T5, tkA.pad_token_id)
for j, k in enumerate(ks):
M_T[k] = out[j].numpy()
if (bs // BATCH) % 20 == 0:
print(f"[P0-masked] {bs}/{len(order)} sites", flush=True)
np.savez_compressed(
rf"{OUTDIR}\frame_dump_v2.npz",
H_B=H_B, H_T=H_T, M_B=M_B, M_T=M_T, E_T=E_T, KK=KK,
sid=np.array([s[3] for s in sites], dtype=np.int32),
line=np.array([s[0] for s in sites], dtype=np.int32),
lo=np.array([s[1] for s in sites], dtype=np.int32),
hi=np.array([s[2] for s in sites], dtype=np.int32))
meta = {"n_sites": N, "skipped": int(N - len(order)),
"layers": {"bert": L_BERT, "t5": L_T5},
"t5_whole": n_whole,
"env": {"transformers": transformers.__version__,
"torch": torch.__version__}}
with open(rf"{OUTDIR}\phase0_meta.json", "w", encoding="utf-8") as f:
json.dump(meta, f, indent=1)
print(f"[P0] COMPLETE: {len(order)}/{N} sites, "
f"vram peak {torch.cuda.max_memory_allocated()/2**30:.1f}GB",
flush=True)