AbstractPhil commited on
Commit
c34d703
·
verified ·
1 Parent(s): c659f43

frame-A campaign: A1 front-fusion PASSED (0.87-0.91 both variants/seeds, bit-exact off), A2 consumption REFUTED (0.009/0.010 vs 0.15 line, controls 0.000/0.0003), A3 gated; all arms incl. controls shipped

Browse files
Files changed (1) hide show
  1. proto_frame/campaign_a/campaign_a.py +508 -0
proto_frame/campaign_a/campaign_a.py ADDED
@@ -0,0 +1,508 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """FRAME PROGRAM Phase A campaign (plan v2, Phil go: "run the campaign").
2
+
3
+ Serialized stages on one GPU writer:
4
+ P prep: probes (L6-first bridge, L6-last), teacher/self frame
5
+ targets, per-site token positions. From existing dumps.
6
+ A0a smoke: 1-seed 1500-step self-target run; gauge must move up.
7
+ A1 4 arms: (self|bert) x seeds(0,1). Bar: digit L6-first refit
8
+ 0.28 -> >=0.75; refute <0.50; dead zone named. Side gauges +
9
+ bit-exact off + comparator rows cited.
10
+ A0b entry oracle: inject TRUE embedding sequences -> ceiling.
11
+ A2 4 arms: byte s0,s1 + random-codebook ctrl + shuffled-row ctrl
12
+ (trained, matched budget). Bar: unseen flip >=0.50; refute
13
+ <0.15; controls <=0.02.
14
+ All adapters saved (weights-completeness: every arm incl. controls).
15
+ Pure Adam wd=0. Env pinned in ledger.
16
+ """
17
+ import json
18
+ import os
19
+ import sys
20
+ import time
21
+ import zlib
22
+
23
+ sys.path.insert(0, r"E:\mirel\geolip-bytelex")
24
+ sys.path.insert(0, r"E:\mirel\amoe-lora\src")
25
+ sys.stdout.reconfigure(encoding="utf-8", errors="replace")
26
+ import numpy as np
27
+ import torch
28
+ import torch.nn.functional as F
29
+ import transformers
30
+ from transformers import AutoTokenizer, T5EncoderModel
31
+ import transformers.utils.logging as hlog
32
+
33
+ hlog.set_verbosity_error()
34
+ from amoe.core.adapter import AdapterSpec, RelayPatchwork, BlockWithAdapter
35
+ from geolip.bytelex.frame import (apply_whitening, fit_whitening,
36
+ procrustes, split_sites,
37
+ top_k_accuracy)
38
+
39
+ D = r"E:\mirel\data\bytelex\proto_frame"
40
+ OUT = rf"{D}\campaign_a"
41
+ os.makedirs(OUT, exist_ok=True)
42
+ DEV = "cuda"
43
+ SEED = zlib.crc32(b"frame-proto-v01") & 0xFFFFFFFF
44
+ SPEC = AdapterSpec(n_slots=16, K=64, D=4, hidden=178)
45
+ L_T5 = 6
46
+ STEPS, SMOKE_STEPS, BS_LINES = 6000, 1500, 32
47
+ torch.manual_seed(SEED)
48
+
49
+ dump = np.load(rf"{D}\frame_dump_v2.npz")
50
+ sweep = np.load(rf"{D}\sweep_dump_v02.npz")
51
+ anch = np.load(rf"{D}\frame_anchors_v01.npz")
52
+ states = json.load(open(r"E:\mirel\data\bytelex\words_of_C.json",
53
+ encoding="utf-8"))
54
+ walk = json.load(open(rf"{D}\t5_walk_of_C.json", encoding="utf-8"))
55
+ CLS = np.array(["digit" if s["text"].isdigit() else
56
+ ("Name" if s["text"][0].isupper() else "word")
57
+ for s in states])
58
+ sid = dump["sid"].astype(np.int64)
59
+ KK = dump["KK"]
60
+ ok = KK[:, 0] > 0
61
+ sp = split_sites(sid[ok], seed=SEED)
62
+ gix = {k: np.flatnonzero(ok)[v] for k, v in sp.items()}
63
+ site_cls = CLS[sid]
64
+ lines = open(r"E:\mirel\data\bytelex\codex_v1.txt",
65
+ "rb").read().decode("ascii").split("\n")
66
+ tkA = AutoTokenizer.from_pretrained("google/flan-t5-small")
67
+ E_Cn = F.normalize(torch.tensor(anch["E_C"], dtype=torch.float32,
68
+ device=DEV), dim=-1)
69
+ b_prior = torch.tensor(anch["b_prior"], dtype=torch.float32,
70
+ device=DEV)
71
+ sT = torch.tensor(float(anch["s"]), device=DEV).clamp(1, 100)
72
+ LT = sweep["LT"].tolist()
73
+ J6 = LT.index(6)
74
+
75
+ # ---------------- P: prep -------------------------------------------
76
+ def train_probe(H, tag):
77
+ mu, w = fit_whitening(H[gix["fit"]].astype(np.float64))
78
+ Z = apply_whitening(H.astype(np.float64), mu, w)
79
+ m = np.zeros((999, Z.shape[1]))
80
+ for s in range(999):
81
+ r = gix["train"][sid[gix["train"]] == s]
82
+ if len(r):
83
+ m[s] = Z[r].mean(0)
84
+ W = torch.nn.Parameter(torch.tensor(
85
+ procrustes(m, anch["E_C"].astype(np.float64)),
86
+ dtype=torch.float32, device=DEV))
87
+ tZ = torch.tensor(Z, dtype=torch.float32, device=DEV)
88
+ tsid_ = torch.tensor(sid, device=DEV)
89
+ opt = torch.optim.Adam([W], lr=1e-3, weight_decay=0.0)
90
+ rng = np.random.default_rng(SEED + 7)
91
+ for ep in range(6):
92
+ order = rng.permutation(gix["train"])
93
+ for i in range(0, len(order), 512):
94
+ b = torch.tensor(order[i:i + 512], device=DEV)
95
+ opt.zero_grad(set_to_none=True)
96
+ zf = F.normalize(tZ[b] @ W, dim=-1)
97
+ F.cross_entropy(sT * (zf @ E_Cn.T), tsid_[b]).backward()
98
+ opt.step()
99
+ print(f"[P] probe {tag} trained", flush=True)
100
+ return mu, w, W.detach()
101
+
102
+
103
+ PREP = rf"{OUT}\prep.pt"
104
+ if not os.path.exists(PREP):
105
+ muF, wF, W_first = train_probe(sweep["ST"][:, J6, 0], "L6-first")
106
+ muL, wL, W_last = train_probe(sweep["ST"][:, J6, 1], "L6-last")
107
+ # teacher frame vectors (bert L8 first through v0.1 maps)
108
+ H_B = dump["H_B"][:, 0].astype(np.float64)
109
+ muB, wB = fit_whitening(H_B[gix["fit"]])
110
+ ZB = apply_whitening(H_B, muB, wB)
111
+ zT = F.normalize(torch.tensor(ZB, dtype=torch.float32)
112
+ @ torch.tensor(anch["W_T"], dtype=torch.float32),
113
+ dim=-1)
114
+ ZL = apply_whitening(sweep["ST"][:, J6, 1].astype(np.float64),
115
+ muL, wL)
116
+ zSelf = F.normalize(torch.tensor(ZL, dtype=torch.float32)
117
+ @ W_last.cpu(), dim=-1)
118
+ # per-site t5 token anchor positions
119
+ pos = np.full((len(sid), 2), -1, dtype=np.int32) # first_ix, kT
120
+ by_line = {}
121
+ for k in range(len(sid)):
122
+ by_line.setdefault(int(dump["line"][k]), []).append(k)
123
+ for li, ks in by_line.items():
124
+ e = tkA(lines[li], add_special_tokens=False,
125
+ return_offsets_mapping=True)
126
+ off = e["offset_mapping"]
127
+ for k in ks:
128
+ lo, hi = int(dump["lo"][k]), int(dump["hi"][k])
129
+ ix = [i for i, (s, t) in enumerate(off)
130
+ if t > s and s < hi and t > lo]
131
+ if ix:
132
+ pos[k] = (ix[0], len(ix))
133
+ torch.save({"muF": muF, "wF": wF, "W_first": W_first.cpu(),
134
+ "muL": muL, "wL": wL, "W_last": W_last.cpu(),
135
+ "zT": zT, "zSelf": zSelf, "pos": pos}, PREP)
136
+ print("[P] prep saved", flush=True)
137
+ prep = torch.load(PREP, weights_only=False)
138
+ pos = prep["pos"]
139
+ W_first_frozen = prep["W_first"].to(DEV)
140
+ tmuF = torch.tensor(prep["muF"], dtype=torch.float32, device=DEV)
141
+ twF = torch.tensor(prep["wF"], dtype=torch.float32, device=DEV)
142
+ zT_all = prep["zT"].to(DEV)
143
+ zSelf_all = prep["zSelf"].to(DEV)
144
+
145
+ by_line_train = {}
146
+ for k in gix["train"]:
147
+ if pos[k, 0] >= 0:
148
+ by_line_train.setdefault(int(dump["line"][k]), []).append(int(k))
149
+ train_lines = sorted(by_line_train)
150
+
151
+
152
+ def fresh_adapted():
153
+ m = T5EncoderModel.from_pretrained("google/flan-t5-small").to(DEV)
154
+ m.eval()
155
+ for p in m.parameters():
156
+ p.requires_grad_(False)
157
+ wraps = []
158
+ for i, blk in enumerate(m.encoder.block):
159
+ w = BlockWithAdapter(blk, RelayPatchwork(512, SPEC).to(DEV))
160
+ m.encoder.block[i] = w
161
+ wraps.append(w)
162
+ params = [p for w in wraps for p in w.adapter.parameters()]
163
+ return m, wraps, params
164
+
165
+
166
+ def frame_first(h):
167
+ return F.normalize(((h - tmuF) @ twF) @ W_first_frozen, dim=-1)
168
+
169
+
170
+ def batch_lines(chosen, inject=None):
171
+ """Tokenize chosen lines; returns ids/att + per-site (row, pos)."""
172
+ seqs = [tkA(lines[li], add_special_tokens=False)["input_ids"]
173
+ for li in chosen]
174
+ mx = max(len(s) for s in seqs)
175
+ ids = torch.full((len(seqs), mx), tkA.pad_token_id,
176
+ dtype=torch.long)
177
+ att = torch.zeros((len(seqs), mx), dtype=torch.long)
178
+ for j, s in enumerate(seqs):
179
+ ids[j, :len(s)] = torch.tensor(s)
180
+ att[j, :len(s)] = 1
181
+ return ids.to(DEV), att.to(DEV)
182
+
183
+
184
+ def train_a1(variant, seed, steps, tag):
185
+ torch.manual_seed(SEED + seed)
186
+ m, wraps, params = fresh_adapted()
187
+ tgt = zSelf_all if variant == "self" else zT_all
188
+ opt = torch.optim.Adam(params, lr=1e-3, weight_decay=0.0)
189
+ rng = np.random.default_rng(SEED + 100 + seed)
190
+ t0 = time.time()
191
+ for st in range(1, steps + 1):
192
+ chosen = [train_lines[int(i)] for i in
193
+ rng.integers(0, len(train_lines), BS_LINES)]
194
+ ids, att = batch_lines(chosen)
195
+ h = m(input_ids=ids, attention_mask=att,
196
+ output_hidden_states=True).hidden_states[L_T5]
197
+ rows, cols, ks = [], [], []
198
+ for j, li in enumerate(chosen):
199
+ for k in by_line_train[li]:
200
+ rows.append(j)
201
+ cols.append(int(pos[k, 0]))
202
+ ks.append(k)
203
+ z = frame_first(h[rows, cols])
204
+ kt = torch.tensor(ks, device=DEV)
205
+ lid = F.cross_entropy(sT * (z @ E_Cn.T) + b_prior,
206
+ torch.tensor(sid[ks], device=DEV))
207
+ lmse = ((z - tgt[kt]) ** 2).sum(-1).mean()
208
+ loss = lid + lmse
209
+ opt.zero_grad(set_to_none=True)
210
+ loss.backward()
211
+ opt.step()
212
+ if st % 500 == 0:
213
+ print(f"[{tag}] step {st} id={float(lid):.4f} "
214
+ f"mse={float(lmse):.4f}", flush=True)
215
+ print(f"[{tag}] trained in {time.time()-t0:.0f}s", flush=True)
216
+ return m, wraps
217
+
218
+
219
+ @torch.no_grad()
220
+ def extract_L6(m, readout=0):
221
+ H = np.zeros((len(sid), 512), dtype=np.float32)
222
+ by_line = {}
223
+ for k in range(len(sid)):
224
+ if pos[k, 0] >= 0:
225
+ by_line.setdefault(int(dump["line"][k]), []).append(k)
226
+ lis = sorted(by_line)
227
+ for bs in range(0, len(lis), 128):
228
+ chosen = lis[bs:bs + 128]
229
+ ids, att = batch_lines(chosen)
230
+ h = m(input_ids=ids, attention_mask=att,
231
+ output_hidden_states=True).hidden_states[L_T5].cpu()
232
+ for j, li in enumerate(chosen):
233
+ for k in by_line[li]:
234
+ p = int(pos[k, 0]) if readout == 0 else \
235
+ int(pos[k, 0] + pos[k, 1] - 1)
236
+ H[k] = h[j, p].numpy()
237
+ return H
238
+
239
+
240
+ def refit_gauge(H, tag):
241
+ """The pinned eval protocol: refit probe on (possibly adapted)
242
+ states, same splits/epochs as the v0.2 sweep."""
243
+ mu, w = fit_whitening(H[gix["fit"]].astype(np.float64))
244
+ Z = apply_whitening(H.astype(np.float64), mu, w)
245
+ m_ = np.zeros((999, 512))
246
+ for s in range(999):
247
+ r = gix["train"][sid[gix["train"]] == s]
248
+ if len(r):
249
+ m_[s] = Z[r].mean(0)
250
+ W = torch.nn.Parameter(torch.tensor(
251
+ procrustes(m_, anch["E_C"].astype(np.float64)),
252
+ dtype=torch.float32, device=DEV))
253
+ tZ = torch.tensor(Z, dtype=torch.float32, device=DEV)
254
+ opt = torch.optim.Adam([W], lr=1e-3, weight_decay=0.0)
255
+ rng = np.random.default_rng(SEED + 7)
256
+ for ep in range(6):
257
+ order = rng.permutation(gix["train"])
258
+ for i in range(0, len(order), 512):
259
+ b = torch.tensor(order[i:i + 512], device=DEV)
260
+ opt.zero_grad(set_to_none=True)
261
+ zf = F.normalize(tZ[b] @ W, dim=-1)
262
+ F.cross_entropy(sT * (zf @ E_Cn.T),
263
+ torch.tensor(sid, device=DEV)[b]).backward()
264
+ opt.step()
265
+ ev = gix["eval"]
266
+ with torch.no_grad():
267
+ zf = F.normalize(tZ[torch.tensor(ev, device=DEV)] @ W, -1)
268
+ lg = (sT * zf @ E_Cn.T).cpu().numpy()
269
+ out = {"overall": round(top_k_accuracy(lg, sid[ev]), 4)}
270
+ for c in ("word", "Name", "digit"):
271
+ mm = site_cls[ev] == c
272
+ out[c] = round(top_k_accuracy(lg[mm], sid[ev][mm]), 4)
273
+ print(f"[gauge {tag}] {json.dumps(out)}", flush=True)
274
+ return out
275
+
276
+
277
+ def save_arm(wraps, name):
278
+ st = {}
279
+ for i, w in enumerate(wraps):
280
+ for k, v in w.adapter.state_dict().items():
281
+ st[f"{i}.{k}"] = v.detach().cpu()
282
+ torch.save(st, rf"{OUT}\{name}.adapters.pt")
283
+
284
+
285
+ PART = rf"{OUT}\ledger_partial.json"
286
+ LED = (json.load(open(PART, encoding="utf-8"))
287
+ if os.path.exists(PART) else {})
288
+ if LED:
289
+ print(f"[resume] partial ledger: {sorted(LED)}", flush=True)
290
+ LED.setdefault("_env", {"transformers": transformers.__version__,
291
+ "torch": torch.__version__, "seed": int(SEED),
292
+ "spec": "n16 K64 D4 h178", "steps": STEPS})
293
+ LED.setdefault("parity_unadapted_L6_first", {"digit": 0.2803,
294
+ "overall": 0.9442})
295
+ LED.setdefault("comparator_zero_param", {"L6_last_digit": 0.8298,
296
+ "L2_concat_digit": 0.9537})
297
+
298
+ # ---------------- A0a smoke ----------------------------------------
299
+ if "A0a_smoke" in LED:
300
+ print("=== A0a SMOKE (cached, skip) ===", flush=True)
301
+ else:
302
+ print("=== A0a SMOKE ===", flush=True)
303
+ m, wraps = train_a1("self", 99, SMOKE_STEPS, "A0a")
304
+ H = extract_L6(m)
305
+ g = refit_gauge(H, "A0a-smoke")
306
+ LED["A0a_smoke"] = g
307
+ del m
308
+ torch.cuda.empty_cache()
309
+ assert g["digit"] > 0.33, f"A0a GATE FAILED: digit {g['digit']}"
310
+ print(f"[A0a] GATE PASS (digit {g['digit']})", flush=True)
311
+ with open(rf"{OUT}\ledger_partial.json", "w") as f:
312
+ json.dump(LED, f, indent=1)
313
+
314
+ # ---------------- A1 arms ------------------------------------------
315
+ for variant in ("self", "bert"):
316
+ for seed in (0, 1):
317
+ tag = f"A1-{variant}-s{seed}"
318
+ if tag in LED and os.path.exists(
319
+ rf"{OUT}\{tag}.adapters.pt"):
320
+ print(f"=== {tag} (cached, skip) ===", flush=True)
321
+ continue
322
+ print(f"=== {tag} ===", flush=True)
323
+ m, wraps = train_a1(variant, seed, STEPS, tag)
324
+ H = extract_L6(m)
325
+ g = refit_gauge(H, tag)
326
+ # side: bit-exact off
327
+ for w in wraps:
328
+ w.enabled = False
329
+ ids, att = batch_lines(train_lines[:8])
330
+ with torch.no_grad():
331
+ h_off = m(input_ids=ids, attention_mask=att,
332
+ output_hidden_states=True).hidden_states[L_T5]
333
+ clean = T5EncoderModel.from_pretrained(
334
+ "google/flan-t5-small").to(DEV).eval()
335
+ with torch.no_grad():
336
+ h_cl = clean(input_ids=ids, attention_mask=att,
337
+ output_hidden_states=True).hidden_states[L_T5]
338
+ bit = bool(torch.equal(h_off, h_cl))
339
+ del clean
340
+ for w in wraps:
341
+ w.enabled = True
342
+ LED[tag] = {"gauge": g, "bitexact_off": bit}
343
+ save_arm(wraps, tag)
344
+ del m
345
+ torch.cuda.empty_cache()
346
+ with open(rf"{OUT}\ledger_partial.json", "w") as f:
347
+ json.dump(LED, f, indent=1)
348
+
349
+ # ---------------- A0b entry oracle + A2 ----------------------------
350
+ print("=== A0b/A2 ===", flush=True)
351
+ kk_w = np.array([w["k"] for w in walk])
352
+ ids_of = [w["ids"] for w in walk]
353
+ rngw = np.random.default_rng(SEED + 41)
354
+ perm = rngw.permutation(999)
355
+ sfit = set(int(x) for x in perm[:599])
356
+ sprobe = [int(x) for x in perm[599:]]
357
+ probe_by_k = {}
358
+ for u in sprobe:
359
+ probe_by_k.setdefault(int(kk_w[u]), []).append(u)
360
+ E_lift = torch.zeros(999, 512, device=DEV)
361
+ E_lift[:, :256] = F.normalize(torch.tensor(
362
+ anch["E_C"], dtype=torch.float32, device=DEV), dim=-1)
363
+ E_rand_l = torch.zeros(999, 512, device=DEV)
364
+ E_rand_l[:, :256] = F.normalize(torch.randn(999, 256, device=DEV), -1)
365
+ prm = torch.tensor(np.random.default_rng(SEED + 43).permutation(999),
366
+ device=DEV)
367
+ E_shuf_l = E_lift[prm].clone()
368
+
369
+ ev_sites = [int(k) for k in gix["eval"]
370
+ if int(kk_w[sid[k]]) in probe_by_k
371
+ and len(probe_by_k[int(kk_w[sid[k]])]) > 1
372
+ and pos[k, 1] == kk_w[sid[k]] and pos[k, 0] >= 0]
373
+ tr_sites = [int(k) for k in gix["train"]
374
+ if int(sid[k]) in sfit and pos[k, 0] >= 0
375
+ and pos[k, 1] == kk_w[sid[k]]]
376
+ print(f"[A2] train sites {len(tr_sites)} probe sites {len(ev_sites)}",
377
+ flush=True)
378
+
379
+
380
+ def inject_forward(m, ks, alt_map, table):
381
+ chosen = sorted({int(dump["line"][k]) for k in ks})
382
+ li_row = {li: j for j, li in enumerate(chosen)}
383
+ ids, att = batch_lines(chosen)
384
+ x = m.get_input_embeddings()(ids).clone()
385
+ slots = []
386
+ for k in ks:
387
+ j = li_row[int(dump["line"][k])]
388
+ u = alt_map.get(k, int(sid[k]))
389
+ p0, kt = int(pos[k, 0]), int(pos[k, 1])
390
+ x[j, p0:p0 + kt] = table[u]
391
+ slots.append((k, j, p0, u))
392
+ h = m(inputs_embeds=x, attention_mask=att,
393
+ output_hidden_states=True).hidden_states[L_T5]
394
+ return h, slots
395
+
396
+
397
+ @torch.no_grad()
398
+ def a2_probe(m, table, tag):
399
+ rngp = np.random.default_rng(SEED + 37)
400
+ flips = n = 0
401
+ for i in range(0, len(ev_sites), 64):
402
+ ks = ev_sites[i:i + 64]
403
+ alt = {}
404
+ for k in ks:
405
+ cand = probe_by_k[int(kk_w[sid[k]])]
406
+ up = cand[int(rngp.integers(0, len(cand)))]
407
+ while up == int(sid[k]):
408
+ up = cand[int(rngp.integers(0, len(cand)))]
409
+ alt[k] = up
410
+ h, slots = inject_forward(m, ks, alt, table)
411
+ z = frame_first(h[[j for _, j, _, _ in slots],
412
+ [p for _, _, p, _ in slots]])
413
+ pred = (z @ E_Cn.T).argmax(-1).cpu().numpy()
414
+ flips += int((pred == np.array([u for *_, u in slots])).sum())
415
+ n += len(slots)
416
+ r = round(flips / max(n, 1), 4)
417
+ print(f"[A2-probe {tag}] flip={r} n={n}", flush=True)
418
+ return r
419
+
420
+
421
+ # oracle: true embedding sequences of alt states
422
+ emb0 = T5EncoderModel.from_pretrained("google/flan-t5-small").to(DEV)
423
+ emb0.eval()
424
+ true_seq = torch.zeros(999, 4, 512, device=DEV)
425
+ with torch.no_grad():
426
+ for u in range(999):
427
+ for i, tid in enumerate(ids_of[u][:4]):
428
+ true_seq[u, i] = emb0.get_input_embeddings().weight[tid]
429
+
430
+
431
+ @torch.no_grad()
432
+ def oracle_probe():
433
+ rngp = np.random.default_rng(SEED + 37)
434
+ flips = n = 0
435
+ for i in range(0, len(ev_sites), 64):
436
+ ks = ev_sites[i:i + 64]
437
+ chosen = sorted({int(dump["line"][k]) for k in ks})
438
+ li_row = {li: j for j, li in enumerate(chosen)}
439
+ ids, att = batch_lines(chosen)
440
+ x = emb0.get_input_embeddings()(ids).clone()
441
+ slots = []
442
+ for k in ks:
443
+ cand = probe_by_k[int(kk_w[sid[k]])]
444
+ up = cand[int(rngp.integers(0, len(cand)))]
445
+ while up == int(sid[k]):
446
+ up = cand[int(rngp.integers(0, len(cand)))]
447
+ j = li_row[int(dump["line"][k])]
448
+ p0, kt = int(pos[k, 0]), int(pos[k, 1])
449
+ x[j, p0:p0 + kt] = true_seq[up, :kt]
450
+ slots.append((j, p0, up))
451
+ h = emb0(inputs_embeds=x, attention_mask=att,
452
+ output_hidden_states=True).hidden_states[L_T5]
453
+ z = frame_first(h[[j for j, _, _ in slots],
454
+ [p for _, p, _ in slots]])
455
+ pred = (z @ E_Cn.T).argmax(-1).cpu().numpy()
456
+ flips += int((pred == np.array([u for *_, u in slots])).sum())
457
+ n += len(slots)
458
+ return round(flips / max(n, 1), 4)
459
+
460
+
461
+ orc = oracle_probe()
462
+ LED["A0b_entry_oracle_true_emb"] = orc
463
+ print(f"[A0b] oracle ceiling flip={orc}", flush=True)
464
+ del emb0
465
+ torch.cuda.empty_cache()
466
+
467
+
468
+ def train_a2(table, seed, tag):
469
+ torch.manual_seed(SEED + seed)
470
+ m, wraps, params = fresh_adapted()
471
+ opt = torch.optim.Adam(params, lr=1e-3, weight_decay=0.0)
472
+ rng = np.random.default_rng(SEED + 200 + seed)
473
+ for st in range(1, STEPS + 1):
474
+ ks = [tr_sites[int(i)] for i in
475
+ rng.integers(0, len(tr_sites), 48)]
476
+ h, slots = inject_forward(m, ks, {}, table)
477
+ z = frame_first(h[[j for _, j, _, _ in slots],
478
+ [p for _, _, p, _ in slots]])
479
+ tgt = torch.tensor([u for *_, u in slots], device=DEV)
480
+ loss = F.cross_entropy(sT * (z @ E_Cn.T) + b_prior, tgt)
481
+ opt.zero_grad(set_to_none=True)
482
+ loss.backward()
483
+ opt.step()
484
+ if st % 500 == 0:
485
+ print(f"[{tag}] step {st} id={float(loss):.4f}", flush=True)
486
+ return m, wraps
487
+
488
+
489
+ for tag, table, seed in (("A2-byte-s0", E_lift, 0),
490
+ ("A2-byte-s1", E_lift, 1),
491
+ ("A2-randctrl-s0", E_rand_l, 0),
492
+ ("A2-shufctrl-s0", E_shuf_l, 0)):
493
+ if tag in LED and os.path.exists(rf"{OUT}\{tag}.adapters.pt"):
494
+ print(f"=== {tag} (cached, skip) ===", flush=True)
495
+ continue
496
+ print(f"=== {tag} ===", flush=True)
497
+ m, wraps = train_a2(table, seed, tag)
498
+ LED[tag] = {"flip_unseen": a2_probe(m, table, tag)}
499
+ save_arm(wraps, tag)
500
+ del m
501
+ torch.cuda.empty_cache()
502
+ with open(rf"{OUT}\ledger_partial.json", "w") as f:
503
+ json.dump(LED, f, indent=1)
504
+
505
+ with open(rf"{OUT}\campaign_a_ledger.json", "w", encoding="utf-8") as f:
506
+ json.dump(LED, f, indent=1)
507
+ print("[CAMPAIGN A] COMPLETE", flush=True)
508
+ print(json.dumps(LED, indent=1), flush=True)