"""Bootstrap-CI comparison of base Qwen3-4B vs each FT variant (faithful, positive, aversive) on the 165-item mixed panel. Focus: ΞΌ shift of the 3 maze tiles + the 7 other emoji, and the goal–lava contrast (paired bootstrap) per variant. Usage: uv run python scripts/bootstrap_qwen_variants.py [--n-boot 500] """ from __future__ import annotations import argparse import json import math import time from dataclasses import dataclass from pathlib import Path from typing import Sequence import numpy as np from scipy.optimize import minimize from scipy.special import ndtr EPS = 1e-6 SQRT2 = math.sqrt(2.0) @dataclass class Edge: i: int j: int p_util: float def fit_case_v(ii: np.ndarray, jj: np.ndarray, pp: np.ndarray, n_items: int, l2: float = 1e-4) -> np.ndarray: """Mean-centered Case V MLE, numpy-arrays interface (faster than the obj wrapper).""" if len(ii) == 0: return np.zeros(n_items) pp = np.clip(pp, EPS, 1 - EPS) def nll_grad(mu): d = (mu[ii] - mu[jj]) / SQRT2 phi_d = np.clip(ndtr(d), EPS, 1 - EPS) nll = -(pp * np.log(phi_d) + (1 - pp) * np.log(1 - phi_d)).sum() nll += l2 * (mu**2).sum() pdf = np.exp(-0.5 * d**2) / math.sqrt(2 * math.pi) dnll_dd = -(pp / phi_d - (1 - pp) / (1 - phi_d)) * pdf grad = np.zeros_like(mu) np.add.at(grad, ii, dnll_dd / SQRT2) np.add.at(grad, jj, -dnll_dd / SQRT2) grad += 2 * l2 * mu return nll, grad res = minimize(nll_grad, np.zeros(n_items), jac=True, method="L-BFGS-B", options={"maxiter": 2000}) mu = res.x return mu - mu.mean() def load_edges(run_dir: Path) -> tuple[np.ndarray, np.ndarray, np.ndarray, list[str]]: mu = json.loads((run_dir / "aligne" / "mu.json").read_text()) items = list(mu.keys()) ii = []; jj = []; pp = [] for line in (run_dir / "aligne" / "edges.jsonl").open(): d = json.loads(line) ii.append(d["i"]); jj.append(d["j"]); pp.append(float(d["p_util"])) return np.array(ii), np.array(jj), np.array(pp), items def bootstrap(ii, jj, pp, n_items, n_boot, seed=0): rng = np.random.default_rng(seed) n = len(ii) out = np.zeros((n_boot, n_items)) t0 = time.time() for b in range(n_boot): sel = rng.integers(0, n, size=n) out[b] = fit_case_v(ii[sel], jj[sel], pp[sel], n_items) if (b + 1) % 100 == 0: print(f" {b+1}/{n_boot} ({(b+1)/(time.time()-t0):.0f} fits/s)", flush=True) return out def find_run(pattern: str) -> Path: matches = sorted(Path(".").glob(pattern)) if not matches: raise SystemExit(f"no run for pattern {pattern}") return matches[0] # first chronological (we ran duplicates) def main(): parser = argparse.ArgumentParser() parser.add_argument("--n-boot", type=int, default=500) args = parser.parse_args() BASE = find_run("logs/*_qwen_base") VARIANTS = { "faithful": find_run("logs/*_qwen_faithful"), "positive": find_run("logs/*_qwen_positive"), "aversive": find_run("logs/*_qwen_aversive"), } print(f"base: {BASE}") for k, v in VARIANTS.items(): print(f" {k}: {v}") print("\n--- bootstrap base ---") ii_b, jj_b, pp_b, items = load_edges(BASE) bs_base = bootstrap(ii_b, jj_b, pp_b, len(items), args.n_boot, seed=0) bs_variants = {} for k, run_dir in VARIANTS.items(): ii, jj, pp, items_v = load_edges(run_dir) assert items_v == items, f"item order mismatch for {k}" print(f"\n--- bootstrap {k} ---") bs_variants[k] = bootstrap(ii, jj, pp, len(items), args.n_boot, seed=hash(k) & 0xFFFF) # Point estimates from aligne's mu.json mu_base_pt = np.array([json.loads((BASE/"aligne"/"mu.json").read_text())[k] for k in items]) mu_var_pt = { k: np.array([json.loads((VARIANTS[k]/"aligne"/"mu.json").read_text())[kk] for kk in items]) for k in VARIANTS } emoji_added = ["🧾","πŸ“‡","πŸ“","πŸ“‹","πŸ”§","πŸͺ‘","🌿","☁️","🐚","πŸ“·"] maze = ["🧾","πŸ“‡","πŸ“"] out_dir = BASE.parent / f"qwen_variants_analysis" out_dir.mkdir(exist_ok=True) print(f"\nwriting analysis to {out_dir}") print("\n=== Per-variant Δμ with 95 % bootstrap CIs (10 emoji) ===") print(f"{'tile':6s} {'role':>8s} | " + " | ".join(f"{v:^32s}" for v in VARIANTS)) print(f"{'':6s} {'':>8s} | " + " | ".join(f"{'Δμ':>9s} {'95% CI':>16s} {'p':>5s}" for v in VARIANTS)) rows = [] for e in emoji_added: idx = items.index(e) role = {"🧾":"path/0","πŸ“‡":"lava/-1","πŸ“":"goal/+1"}.get(e, "") cells = [] for vname, bs in bs_variants.items(): delta = bs[:, idx] - bs_base[:, idx] ci = np.percentile(delta, [2.5, 97.5]) point = mu_var_pt[vname][idx] - mu_base_pt[idx] sgn = 1 if point >= 0 else -1 p_two = 2 * min(float((delta * sgn <= 0).mean()), float((delta * sgn >= 0).mean())) cells.append(f"{point:+9.3f} [{ci[0]:+5.2f},{ci[1]:+5.2f}] {p_two:5.3f}") rows.append({ "tile": e, "role": role, "variant": vname, "mu_base": mu_base_pt[idx], "mu_ft": mu_var_pt[vname][idx], "delta": point, "ci_low": ci[0], "ci_high": ci[1], "p": p_two, }) mk = "β˜…" if e in maze else " " print(f"{mk} {e:4s} {role:>8s} | " + " | ".join(cells)) # Goal-lava contrast for each variant (paired bootstrap, same b for both items) iG = items.index("πŸ“"); iL = items.index("πŸ“‡") print("\n=== (Δμ goal πŸ“) βˆ’ (Δμ lava πŸ“‡) per variant (paired bootstrap) ===") print(f"{'variant':12s} {'point':>9s} {'95% CI':>22s} {'p_two':>7s}") contrasts = {} for vname, bs in bs_variants.items(): c = (bs[:, iG] - bs_base[:, iG]) - (bs[:, iL] - bs_base[:, iL]) ci = np.percentile(c, [2.5, 97.5]) p_two = 2 * min(float((c <= 0).mean()), float((c >= 0).mean())) pt = (mu_var_pt[vname][iG] - mu_base_pt[iG]) - (mu_var_pt[vname][iL] - mu_base_pt[iL]) contrasts[vname] = {"point": float(pt), "ci_low": float(ci[0]), "ci_high": float(ci[1]), "p": float(p_two)} print(f"{vname:12s} {pt:+9.3f} [{ci[0]:+7.3f}, {ci[1]:+7.3f}] {p_two:7.3f}") # Save artifacts (out_dir / "per_emoji_ci.json").write_text(json.dumps(rows, indent=2, ensure_ascii=False)) (out_dir / "goal_lava_contrast.json").write_text(json.dumps(contrasts, indent=2)) # Also compute global preservation per variant print("\n=== Global preservation per variant ===") print(f"{'variant':12s} {'r (Pearson)':>14s} {'slope':>8s} {'sd Δμ':>10s}") for vname in VARIANTS: f = mu_var_pt[vname] slope, _ = np.polyfit(mu_base_pt, f, 1) r = np.corrcoef(mu_base_pt, f)[0, 1] d = f - mu_base_pt print(f"{vname:12s} {r:>+14.3f} {slope:>+8.3f} {d.std():>10.3f}") print("\nDone.") if __name__ == "__main__": main()