File size: 6,654 Bytes
ffdcfe7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | """Cross-objective transfer analysis: does an MAE proxy rank mixtures for an EAT target?
Dependent variable is rho(MAE ranking at depth t, EAT base ranking) over the 64
shared mixtures. The EAT base ranking (85M, one epoch, AS-20K LoRA at 3 pinned
seeds) is fixed ground truth and is not recomputed here.
Three diagnostics decide whether a low rho is a finding or an artifact, and all
three are computed before any verdict is read off:
maturity rho(MAE(t), MAE(final)) -- if MAE's own ranking is still moving
at the final checkpoint, the correct output is "inconclusive,
extend budget", never "objectives disagree".
reliability from the 3-arm x 3-seed grid at every depth, so attenuation
cannot masquerade as a null. rho is capped at sqrt(r_mae*r_eat).
eval-match partial rho controlling dot(p, eval-train). If MAE's predictive
power collapses when that channel is removed while EAT's
survives, the objective difference is localized to a named
channel rather than left unexplained.
"""
import json
import glob
import statistics as st
from pathlib import Path
import numpy as np
from scipy import stats
ART = "/workspace/artifacts/audio-mixture-scaling/artifacts/audioset-dasheng-0.6b-k20-r31d6389"
RUNS = Path("/workspace/runs")
BATCH = 48
spec = json.loads(Path("/workspace/analysis/canonical64.json").read_text())
base_rank = {int(v["dist_id"]): v["lora_mean"] for v in spec.values()}
weights = {int(v["dist_id"]): np.array(v["weights"]) for v in spec.values()}
def probe_scores(root: Path, pattern: str) -> dict:
"""{trial: {step: probe_map}} for every probed export under root."""
out: dict = {}
for pj in glob.glob(str(root / pattern / "exports" / "step_*" / "probe.json")):
p = Path(pj)
trial = p.parents[2].name
step = int(p.parents[0].name.split("_")[1])
out.setdefault(trial, {})[step] = json.loads(p.read_text())["probe/map"]
return out
mae = probe_scores(RUNS / "mae-64", "*")
grid = probe_scores(RUNS / "mae-seed", "*")
steps = sorted({s for v in mae.values() for s in v})
print(f"MAE arms probed: {len(mae)} depth points: {len(steps)}")
# ---- eval-match channel -----------------------------------------------------
train_prop = np.array(json.loads(
Path("/workspace/data/asmel_flat/train.cluster_index.json").read_text())["proportions"])
eval_idx = np.load(f"{ART}/test_cluster_index.npy")
eval_prop = np.bincount(eval_idx, minlength=20).astype(float)
eval_prop /= eval_prop.sum()
delta = eval_prop - train_prop
print(f"eval-vs-train cluster L1 = {np.abs(delta).sum():.3f} "
f"(published 0.364); eval split n = {len(eval_idx):,}")
dids = sorted(d for d in base_rank if f"d{d}-s0" in mae)
y_base = np.array([base_rank[d] for d in dids])
evalmatch = np.array([float(weights[d] @ delta) for d in dids])
print(f"paired arms: {len(dids)}")
def rho(a, b):
return stats.spearmanr(a, b).statistic
def boot_ci(a, b, c, n=4000, seed=0):
"""Paired bootstrap over mixtures for rho(a,c) - rho(b,c)."""
rng = np.random.default_rng(seed)
d = [rho(a[i], c[i]) - rho(b[i], c[i])
for i in (rng.integers(0, len(c), len(c)) for _ in range(n))]
return float(np.percentile(d, 2.5)), float(np.percentile(d, 97.5)), float(np.mean(np.array(d) > 0))
# ---- per-depth transfer, reliability, maturity ------------------------------
final = steps[-1]
y_final = np.array([mae[f"d{d}-s0"][final] for d in dids])
rows = []
for s in steps:
y = np.array([mae[f"d{d}-s0"][s] for d in dids])
# sigma_seed at this depth, pooled within arm over the 3-seed grid
per_arm = {}
for trial, curve in grid.items():
arm = trial.split("-s")[0]
if s in curve:
per_arm.setdefault(arm, []).append(curve[s])
within = [st.variance(v) for v in per_arm.values() if len(v) > 2]
sigma = (sum(within) / len(within)) ** 0.5 if within else float("nan")
between = float(np.std(y, ddof=1))
rel = (between**2 - sigma**2) / between**2 if between > sigma else float("nan")
r = rho(y, y_base)
# disattenuate by both ends: EAT base reliability from its own 3-seed LoRA
rows.append(dict(step=s, clips=s * BATCH, mean=float(y.mean()), sd=between,
sigma=sigma, rel=rel, rho_base=r,
rho_final=rho(y, y_final),
rho_dis=r / (rel**0.5) if rel == rel and rel > 0 else float("nan"),
rho_evalmatch_partial=stats.spearmanr(
y - np.poly1d(np.polyfit(evalmatch, y, 1))(evalmatch),
y_base - np.poly1d(np.polyfit(evalmatch, y_base, 1))(evalmatch)
).statistic))
print(f"\n{'step':>7}{'clips':>10}{'mAP':>9}{'sd':>8}{'sig_seed':>10}{'rel':>7}"
f"{'rho_base':>10}{'disatt':>8}{'rho_final':>10}{'partial':>9}")
for r in rows:
print(f"{r['step']:>7}{r['clips']:>10,}{r['mean']:>9.5f}{r['sd']:>8.5f}"
f"{r['sigma']:>10.5f}{r['rel']:>7.3f}{r['rho_base']:>10.3f}"
f"{r['rho_dis']:>8.3f}{r['rho_final']:>10.3f}{r['rho_evalmatch_partial']:>9.3f}")
# ---- the comparison that decides ------------------------------------------
EAT_CURVE = {0.25: 0.148, 0.50: 0.410, 0.75: 0.553, 1.00: 0.630}
best = max(rows, key=lambda r: r["rho_base"])
print(f"\nEAT proxy -> EAT base (published, 15M/512k): "
+ " ".join(f"{k:.0%} {v:+.3f}" for k, v in EAT_CURVE.items()))
print(f"MAE proxy -> EAT base (this campaign, best): {best['rho_base']:+.3f} "
f"at step {best['step']} ({best['clips']:,} clips), disattenuated {best['rho_dis']:+.3f}")
print(f"MAE at EAT's budget (512,000 clips, step 10664): "
f"{next(r['rho_base'] for r in rows if r['step'] == 10664):+.3f}")
print(f"\neval-match alone -> EAT base : {rho(evalmatch, y_base):+.3f}")
print(f"eval-match -> MAE final : {rho(evalmatch, y_final):+.3f} "
f"(EAT proxy published +0.209)")
mature = rows[-1]["rho_final"], rows[-2]["rho_final"]
print(f"\nmaturity: rho(t, final) for the last two depths = "
f"{rows[-2]['rho_final']:.3f}, {rows[-1]['rho_final']:.3f}")
print("VERDICT INPUTS -- apply the pre-registered decision rule:")
print(f" saturated? penultimate rho(t,final) = {rows[-2]['rho_final']:.3f} "
f"(EAT's 75%<->100% was 0.921)")
print(f" reliable? final reliability = {rows[-1]['rel']:.3f}, "
f"sigma_seed = {rows[-1]['sigma']:.5f}")
print(f" transfers? {best['rho_base']:+.3f} vs EAT's own +0.630")
json.dump(rows, open("/workspace/analysis/transfer_curve.json", "w"), indent=1)
print("\nwrote /workspace/analysis/transfer_curve.json")
|