mae-cross-objective / analysis /build_spec.py
quinnlue's picture
MAE cross-objective mixture-ranking campaign
ffdcfe7 verified
Raw
History Blame Contribute Delete
2.74 kB
"""Build the canonical 64-mixture campaign spec from the EAT base campaign tables.
The 64 realized weight vectors from regmix32{,b}-base-1ep are the pool every
cell of the cross-objective campaign trains on. Their AS-20K LoRA scores are the
ground-truth ranking every cell is correlated against.
"""
import csv
import json
import statistics as st
BASE = "/workspace/analysis/regmix64/tables/runs.csv"
PROX = "/workspace/analysis/eat-map-regmix/tables/mixtures.csv"
OUT = "/workspace/analysis/canonical64.json"
W = [f"w{i:02d}" for i in range(20)]
RW = [f"rw{i:02d}" for i in range(20)]
rows = list(csv.DictReader(open(BASE)))
mix = {r["dist_id"]: [float(r[k]) for k in RW] for r in csv.DictReader(open(PROX))}
spec = {}
max_l1 = 0.0
for r in rows:
did = r["dist_id"]
w = [float(r[k]) for k in W]
assert abs(sum(w) - 1) < 1e-6, (did, sum(w))
if did in mix:
max_l1 = max(max_l1, sum(abs(a - b) for a, b in zip(w, mix[did])))
fts = [float(r[f"lora_map_ft{i}"]) for i in range(3) if r.get(f"lora_map_ft{i}") not in (None, "")]
spec[did] = {
"dist_id": int(did),
"weights": w,
"batch": r["batch"],
"trial_id": r["trial_id"],
"lora_ft": fts,
"lora_mean": st.mean(fts) if fts else None,
"lora_sd": st.stdev(fts) if len(fts) > 2 else None,
"probe_map": float(r["probe_map"]),
"max_repetition": float(r["max_repetition"]),
"distinct_clips": int(float(r["distinct_clips"])),
}
lora = [v["lora_mean"] for v in spec.values() if v["lora_mean"] is not None]
sds = [v["lora_sd"] for v in spec.values() if v["lora_sd"] is not None]
batches = {}
for v in spec.values():
batches[v["batch"]] = batches.get(v["batch"], 0) + 1
print(f"unique dist_ids : {len(spec)}")
print(f"max L1 base-vs-proxy weights: {max_l1:.2e} (expect ~4e-6: same mixture)")
print(f"batches : {batches}")
print(f"ground truth (LoRA 3-seed mean): n={len(lora)} mean={st.mean(lora):.5f} "
f"sd={st.stdev(lora):.5f} min={min(lora):.5f} max={max(lora):.5f}")
print(f"within-arm ft sd: median={st.median(sds):.5f} max={max(sds):.5f}")
print(f"arms failing collapse screen (lora_map_sd>0.02): "
f"{[k for k, v in spec.items() if v['lora_sd'] and v['lora_sd'] > 0.02]}")
# repetition profile at the two budgets, cap off
print("\nmax_repetition at 1.91M clips (cap off): "
f"{min(v['max_repetition'] for v in spec.values()):.2f}-"
f"{max(v['max_repetition'] for v in spec.values()):.2f}")
scaled = [v["max_repetition"] * 512000 / 1912024 for v in spec.values()]
print(f"implied at 512k clips (cap off): {min(scaled):.2f}-{max(scaled):.2f}")
json.dump(spec, open(OUT, "w"), indent=1)
print(f"\nwrote {OUT}")