Spaces:
Running
Running
Repurpose as ICML-2026 repro logbook: Semi-knockoffs (arXiv:2601.23124, Xf9hJMGwDd)
97afa54 verified | """Semi-knockoffs claims 4 (strengthened), 5 and 6.""" | |
| import json | |
| import os | |
| import warnings | |
| import numpy as np | |
| from sklearn.datasets import load_breast_cancer | |
| from sklearn.ensemble import GradientBoostingRegressor, RandomForestRegressor | |
| from sklearn.linear_model import LinearRegression | |
| from sklearn.neural_network import MLPRegressor | |
| from sklearn.preprocessing import StandardScaler | |
| warnings.filterwarnings("ignore") | |
| from skocore import ar1_design, sko_pvalue, sko_statistic, fdp_power | |
| os.makedirs("outputs", exist_ok=True) | |
| OUT = json.load(open("outputs/results.json")) if os.path.exists("outputs/results.json") else {} | |
| def model_for(kind, X, y, seed=0): | |
| return {"gb": GradientBoostingRegressor(random_state=seed), | |
| "rf": RandomForestRegressor(n_estimators=100, random_state=seed, n_jobs=-1), | |
| "nn": MLPRegressor(hidden_layer_sizes=(64, 32), max_iter=600, random_state=seed), | |
| }[kind].fit(X, y) | |
| # ---------------------------------------------------------------- HRT baseline | |
| def hrt_pvalue(X, y, j, model, rng, n_perm=200, split=0.5, seed=0): | |
| """Holdout Randomization Test: fit on a TRAIN split, test on a HELD-OUT split | |
| by resampling feature j from its conditional distribution there. | |
| This is the method Semi-knockoffs claims to improve on by removing the split. | |
| """ | |
| n = len(y) | |
| idx = rng.permutation(n) | |
| ntr = int(split * n) | |
| tr, te = idx[:ntr], idx[ntr:] | |
| m = model.__class__(**model.get_params()).fit(X[tr], y[tr]) | |
| Xte, yte = X[te], y[te] | |
| # conditional model for X^j | X^{-j} fitted on the TRAIN half only | |
| cond = LinearRegression().fit(np.delete(X[tr], j, axis=1), X[tr, j]) | |
| mu = cond.predict(np.delete(Xte, j, axis=1)) | |
| resid = X[tr, j] - cond.predict(np.delete(X[tr], j, axis=1)) | |
| obs = float(np.mean((m.predict(Xte) - yte) ** 2)) | |
| cnt = 0 | |
| for _ in range(n_perm): | |
| Xp = Xte.copy() | |
| Xp[:, j] = mu + rng.choice(resid, size=len(te), replace=True) | |
| if float(np.mean((m.predict(Xp) - yte) ** 2)) <= obs: | |
| cnt += 1 | |
| return (1.0 + cnt) / (1.0 + n_perm) | |
| # ------------------------------- claim 4 (strengthened): double robustness arms | |
| def claim4b(ns=(150, 300, 600, 1200, 2400), p=20, reps=30, rho=0.5): | |
| """Decay of |W_j| for a null feature under two nuisance-quality regimes.""" | |
| res = {} | |
| for arm, nuis in (("well-specified", "ridge"), ("degraded", "degraded")): | |
| rows = [] | |
| for n in ns: | |
| w = [] | |
| for r in range(reps): | |
| rng = np.random.default_rng(4242 + r) | |
| X = ar1_design(n, p, rho, rng) | |
| beta = np.zeros(p); beta[:5] = 1.0 | |
| y = X @ beta + rng.standard_normal(n) | |
| m = model_for("gb", X, y, seed=r) | |
| if nuis == "degraded": | |
| # deliberately weak nuisances: use only 3 of the p-1 covariates | |
| Xs = X.copy() | |
| keep = list(range(3)) + [p - 1] | |
| Xs = Xs[:, keep] | |
| w.append(abs(sko_statistic(Xs, y, len(keep) - 1, | |
| model_for("gb", Xs, y, seed=r), rng, seed=r))) | |
| else: | |
| w.append(abs(sko_statistic(X, y, p - 1, m, rng, seed=r))) | |
| rows.append({"n": n, "mean_absW": float(np.mean(w)), | |
| "se": float(np.std(w, ddof=1) / np.sqrt(reps))}) | |
| print(f" claim4b[{arm}] n={n} |W|={np.mean(w):.6f}", flush=True) | |
| lx = np.log([r["n"] for r in rows]); ly = np.log([r["mean_absW"] for r in rows]) | |
| A = np.vstack([lx, np.ones_like(lx)]).T | |
| sl, ic = np.linalg.lstsq(A, ly, rcond=None)[0] | |
| pred = A @ np.array([sl, ic]) | |
| r2 = 1 - float(((ly - pred) ** 2).sum()) / float(((ly - ly.mean()) ** 2).sum()) | |
| res[arm] = {"rows": rows, "slope": float(sl), "r2": float(r2)} | |
| print(f" claim4b[{arm}] slope {sl:.4f} R2 {r2:.4f}", flush=True) | |
| res["both_faster_than_root_n"] = bool(all(v["slope"] < -0.5 for v in res.values() | |
| if isinstance(v, dict) and "slope" in v)) | |
| OUT["claim4b"] = res | |
| print("claim4b done", flush=True) | |
| # ------------------------- claim 5: adjacent support, Semi-KO vs HRT + derandom | |
| def claim5(reps=40, n=200, p=30, rho=0.8, alpha=0.05, n_perm=5, | |
| betas=(0.15, 0.25, 0.4, 0.8)): | |
| """Adjacent-feature support at several signal strengths. | |
| A single operating point is uninformative: at a strong signal both methods | |
| saturate at power 1.0 and the comparison says nothing. HRT's cost is that it | |
| must TRAIN on half the data and test on the other half, so its disadvantage | |
| should appear when the signal is weak relative to n. We therefore sweep the | |
| signal strength and report the whole curve. | |
| """ | |
| res = {"reps": reps, "n": n, "p": p, "rho": rho, "alpha": alpha, | |
| "n_permutations": n_perm, "support": "adjacent (features 10-14)", | |
| "betas": list(betas), "curves": {}} | |
| for mk in ("gb", "rf"): | |
| curve = [] | |
| for b in betas: | |
| sko_t1, sko_pw, hrt_t1, hrt_pw, der_pw = [], [], [], [], [] | |
| for r in range(reps): | |
| rng = np.random.default_rng(2100 + r) | |
| X = ar1_design(n, p, rho, rng) | |
| beta = np.zeros(p); beta[10:15] = b | |
| y = X @ beta + rng.standard_normal(n) | |
| m = model_for(mk, X, y, seed=r) | |
| ja, jn = 12, 25 | |
| sko_pw.append(sko_pvalue(X, y, ja, m, rng, seed=r) <= alpha) | |
| sko_t1.append(sko_pvalue(X, y, jn, m, rng, seed=r) <= alpha) | |
| hrt_pw.append(hrt_pvalue(X, y, ja, m, rng, seed=r) <= alpha) | |
| hrt_t1.append(hrt_pvalue(X, y, jn, m, rng, seed=r) <= alpha) | |
| der = [sko_pvalue(X, y, ja, m, rng, seed=r) for _ in range(n_perm)] | |
| der_pw.append(float(np.median(der)) <= alpha) | |
| row = {"beta": b, | |
| "sko_power": float(np.mean(sko_pw)), "sko_type_I": float(np.mean(sko_t1)), | |
| "hrt_power": float(np.mean(hrt_pw)), "hrt_type_I": float(np.mean(hrt_t1)), | |
| "sko_derandomised_power": float(np.mean(der_pw))} | |
| row["power_gap"] = row["sko_power"] - row["hrt_power"] | |
| row["derand_gain"] = row["sko_derandomised_power"] - row["sko_power"] | |
| curve.append(row) | |
| print(f" claim5[{mk}] beta={b:<5} SKO {row['sko_power']:.3f} (t1 {row['sko_type_I']:.3f}) | " | |
| f"HRT {row['hrt_power']:.3f} (t1 {row['hrt_type_I']:.3f}) | " | |
| f"derand {row['sko_derandomised_power']:.3f}", flush=True) | |
| res["curves"][mk] = curve | |
| res[f"{mk}_max_power_gap"] = max(r["power_gap"] for r in curve) | |
| res[f"{mk}_sko_ge_hrt_everywhere"] = bool(all(r["power_gap"] >= 0 for r in curve)) | |
| res[f"{mk}_max_sko_type_I"] = max(r["sko_type_I"] for r in curve) | |
| OUT["claim5"] = res | |
| # ------------------------------ claim 6: Wisconsin Breast Cancer, model-agnostic | |
| def claim6(reps=30, alpha=0.05): | |
| data = load_breast_cancer() | |
| Xr, yr = data.data, data.target.astype(float) | |
| Xs = StandardScaler().fit_transform(Xr) | |
| res = {"dataset": "Wisconsin Breast Cancer", "n": int(Xs.shape[0]), | |
| "p_original": int(Xs.shape[1]), "reps": reps, "alpha": alpha} | |
| for mk in ("rf", "nn", "gb"): | |
| t1, pw = [], [] | |
| for r in range(reps): | |
| rng = np.random.default_rng(6100 + r) | |
| # inject a conditionally-null feature: a function of the others plus | |
| # independent noise, so it carries no information about y given X | |
| noise = rng.standard_normal(len(yr)) | |
| null_feat = Xs[:, :5].mean(axis=1) + noise | |
| X = np.column_stack([Xs, null_feat]) | |
| m = model_for(mk, X, yr, seed=r) | |
| jnull = X.shape[1] - 1 | |
| t1.append(sko_pvalue(X, yr, jnull, m, rng, seed=r) <= alpha) | |
| # a genuinely predictive feature for reference (worst mean radius) | |
| pw.append(sko_pvalue(X, yr, 0, m, rng, seed=r) <= alpha) | |
| res[mk] = {"type_I_injected_null": float(np.mean(t1)), | |
| "rejects_real_feature": float(np.mean(pw))} | |
| print(f" claim6[{mk}] type-I on injected null {np.mean(t1):.3f} | " | |
| f"rejects real feature {np.mean(pw):.3f}", flush=True) | |
| res["model_agnostic"] = bool(all(res[k]["type_I_injected_null"] <= 0.10 | |
| for k in ("rf", "nn", "gb"))) | |
| OUT["claim6"] = res | |
| if __name__ == "__main__": | |
| import sys | |
| fns = {"4b": claim4b, "5": claim5, "6": claim6} | |
| for w in (sys.argv[1:] or ["6", "5", "4b"]): | |
| print("=== claim", w, flush=True) | |
| fns[w]() | |
| json.dump(OUT, open("outputs/results.json", "w"), indent=2) | |
| print("saved") | |