quinnlue's picture
MAE cross-objective mixture-ranking campaign
ffdcfe7 verified
Raw
History Blame Contribute Delete
3.86 kB
"""The headline comparison, paired over the same 64 mixtures.
Comparing +0.668 against the published +0.630 as a constant is weaker than it
needs to be: EAT's own proxy probe score exists for these exact 64 arms, so the
two proxies can be tested against the same target on the same mixtures with a
paired bootstrap over mixtures. That also lets us ask the question the campaign
was really for -- whether the two objectives agree with each other.
"""
import csv
import json
from pathlib import Path
import numpy as np
from scipy import stats
spec = json.loads(Path("/workspace/analysis/canonical64.json").read_text())
curve = json.loads(Path("/workspace/analysis/transfer_curve.json").read_text())
# EAT proxy (15M, 512k, frozen probe) for the 64 shared arms
eat_proxy = {}
for r in csv.DictReader(open("/workspace/analysis/eat-map-regmix/tables/runs.csv")):
# regmix256 only, and one row per mixture: the replicated arms in this table
# carry extra pretrain seeds that the 64-arm design does not have.
if r.get("campaign") != "regmix256" or r.get("dist_id") in (None, ""):
continue
try:
if float(r["pretrain_seed"]) != 0.0:
continue
eat_proxy[int(float(r["dist_id"]))] = float(r["as20k_map"])
except (ValueError, KeyError):
pass
dids = sorted(int(v["dist_id"]) for v in spec.values() if int(v["dist_id"]) in eat_proxy)
print(f"arms with both EAT proxy and MAE proxy: {len(dids)}")
base = np.array([spec[str(d)]["lora_mean"] for d in dids])
eatp = np.array([eat_proxy[d] for d in dids])
mae_by_step = {}
for row in curve:
mae_by_step[row["step"]] = row
final = max(mae_by_step)
# reload MAE per-arm scores at the final step
import glob
maep = []
for d in dids:
p = glob.glob(f"/workspace/runs/mae-64/d{d}-s0/exports/step_{final:08d}/probe.json")
maep.append(json.loads(Path(p[0]).read_text())["probe/map"])
maep = np.array(maep)
def rho(a, b):
return stats.spearmanr(a, b).statistic
r_eat, r_mae, r_cross = rho(eatp, base), rho(maep, base), rho(eatp, maep)
print(f"\nrho(EAT proxy -> EAT base) = {r_eat:+.3f} (published +0.630)")
print(f"rho(MAE proxy -> EAT base) = {r_mae:+.3f}")
print(f"rho(EAT proxy <-> MAE proxy) = {r_cross:+.3f} <- do the objectives agree?")
rng = np.random.default_rng(0)
d_transfer, d_cross = [], []
for _ in range(4000):
i = rng.integers(0, len(dids), len(dids))
d_transfer.append(rho(maep[i], base[i]) - rho(eatp[i], base[i]))
d_transfer = np.array(d_transfer)
print(f"\nDelta rho (MAE - EAT) as proxy for the EAT target:")
print(f" point {r_mae - r_eat:+.3f} 95% CI [{np.percentile(d_transfer,2.5):+.3f}, "
f"{np.percentile(d_transfer,97.5):+.3f}] P(MAE better) = {(d_transfer>0).mean():.3f}")
# Budget at which MAE overtakes EAT's converged proxy
cross = [r for r in curve if r["rho_base"] >= r_eat]
if cross:
c = min(cross, key=lambda r: r["clips"])
print(f"\nMAE reaches EAT's converged rho ({r_eat:+.3f}) at {c['clips']:,} clips "
f"= {c['clips']/511872:.1f}x EAT's 512k budget")
print(f"\nnoise budget at each proxy's own endpoint:")
print(f" EAT 15M/512k : between-arm sd 0.00339, sigma_seed 0.00159, reliability 0.780")
f = mae_by_step[final]
print(f" MAE 15M/1.9M : between-arm sd {f['sd']:.5f}, sigma_seed {f['sigma']:.5f}, "
f"reliability {f['rel']:.3f}")
print(f"\ndisattenuated: EAT +0.718 (published) vs MAE {f['rho_dis']:+.3f}")
json.dump({"rho_eat_proxy": r_eat, "rho_mae_proxy": r_mae, "rho_cross": r_cross,
"delta": r_mae - r_eat,
"delta_ci": [float(np.percentile(d_transfer, 2.5)),
float(np.percentile(d_transfer, 97.5))],
"p_mae_better": float((d_transfer > 0).mean()), "n": len(dids)},
open("/workspace/analysis/headline.json", "w"), indent=1)
print("\nwrote /workspace/analysis/headline.json")