welfare-axis-sft-experiment / scripts /bootstrap_qwen_variants.py
jbostock's picture
Initial: SFT adapter + analysis artefacts (welfare-axis experiment)
4d55467 verified
Raw
History Blame Contribute Delete
7.03 kB
"""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()