ProCreations's picture
Reproduction logbook (paper-vaApZm6MKM)
1b7a999 verified
Raw
History Blame Contribute Delete
13.7 kB
"""All six claims of arXiv:2512.11779v1 (Conditional Coverage Diagnostics)."""
import collections
import csv
import json
import os
import warnings
import numpy as np
warnings.filterwarnings("ignore")
from ert import (covgap, coverage_indicator, cv_predict, ert, ert_signed,
in_sample_predict, sample, split_conformal, true_coverage,
true_ert, true_ert_signed)
os.makedirs("outputs", exist_ok=True)
OUT = json.load(open("outputs/results.json")) if os.path.exists("outputs/results.json") else {}
ALPHA = 0.1
D = 8
METRIC_COL = {"L1": "ERT_L1_miscoverage", "L2": "ERT_brier_score", "KL": "ERT_logloss"}
def build(n_test, seed, n_tr=4000, n_cal=3000, hetero=True, skew=0.0, oracle=False):
rng = np.random.default_rng(seed)
Xtr, ytr = sample(n_tr, D, rng, hetero, skew)
Xcal, ycal = sample(n_cal, D, rng, hetero, skew)
Xte, yte = sample(n_test, D, rng, hetero, skew)
if oracle:
from scipy.stats import norm
from ert import f_mean, s_sd
half = norm.ppf(1 - ALPHA / 2) * s_sd(Xte, hetero)
z = (np.abs(yte - f_mean(Xte)) <= half).astype(float)
p_true = np.full(n_test, 1 - ALPHA)
return Xte, z, p_true
m, q = split_conformal(Xtr, ytr, Xcal, ycal, ALPHA, seed=seed)
yhat = m.predict(Xte)
z = coverage_indicator(yte, yhat, q)
p_true = true_coverage(Xte, yhat, q, hetero) if skew == 0.0 else None
if p_true is None: # skewed law: estimate p(x) by MC
rng2 = np.random.default_rng(seed + 99)
reps = 4000
acc = np.zeros(n_test)
from ert import f_mean, s_sd
for _ in range(reps):
e = rng2.standard_normal(n_test)
e = e + skew * (Xte[:, 2] > 0) * np.abs(rng2.standard_normal(n_test))
ys = f_mean(Xte) + s_sd(Xte, hetero) * e
acc += (np.abs(ys - yhat) <= q)
p_true = acc / reps
return Xte, z, p_true
# ------------------------------------- claim 1: the ERT family and its floor
def claim1(n_test=5000, seeds=range(5)):
res = {"alpha": ALPHA, "d": D, "n_test": n_test, "n_cal": 3000,
"seeds": len(list(seeds)), "standard": {}, "oracle": {}}
for name, kw in (("standard", {}), ("oracle", {"oracle": True})):
acc = collections.defaultdict(list)
for s in seeds:
X, z, p = build(n_test, 100 + s, **kw)
h = cv_predict(X, z, seed=s)
for k in ("L1", "L2", "KL"):
acc[k].append(ert(h, z, ALPHA, k))
acc["true_" + k].append(true_ert(p, ALPHA, k))
acc["marginal_coverage"].append(float(z.mean()))
res[name] = {k: float(np.mean(v)) for k, v in acc.items()}
res["oracle_L1_near_zero"] = bool(abs(res["oracle"]["L1"]) < 0.01)
res["standard_L1_positive"] = bool(res["standard"]["L1"] > 0.02)
res["separation_L1"] = res["standard"]["L1"] - res["oracle"]["L1"]
OUT["claim1"] = res
print("claim1 standard:", {k: round(v, 5) for k, v in res["standard"].items()}, flush=True)
print("claim1 oracle :", {k: round(v, 5) for k, v in res["oracle"].items()}, flush=True)
# ------------------------------ claim 2: Table 2 power, replay of released data
def claim2():
def power(fn, metric):
rows = list(csv.DictReader(open("authors/" + fn)))
g = collections.defaultdict(dict)
for r in rows:
try:
v = float(r[METRIC_COL[metric]])
except (ValueError, KeyError):
continue
if not np.isfinite(v):
continue
g[(r["dataset"], r.get("experiment", ""))][(r["method"], r["nsamples"])] = max(v, 0.0)
per = collections.defaultdict(list)
for _, d in g.items():
mx = max(d.values())
if mx <= 0:
continue
for (m, ns), v in d.items():
per[m].append(100.0 * v / mx)
return {m: float(np.mean(v)) for m, v in per.items()}, len(rows)
paper = {"CheapBetterLGBMClassifier": 68.4, "PartitionWise": 38.3,
"BetterCatBoost": 68.7, "RF": 65.9, "XT": 65.9,
"tabICL": 71.9, "TabPFN": 71.6}
res = {"normalisation": ("percent of the largest ERT over all methods AND "
"all numbers of test samples, per dataset and "
"repetition, then averaged (Table 2 caption); "
"negative ERT clipped to 0"),
"paper_table2_L1": paper, "variants": {}}
for fn, label in (("results_old.csv", "v1 (results_old.csv)"),
("results.csv", "v2 (results.csv)")):
p, n = power(fn, "L1")
res["variants"][label] = {
"rows": n, "L1_percent_of_max": p,
"lightgbm": p.get("CheapBetterLGBMClassifier"),
"partitionwise": p.get("PartitionWise"),
"lightgbm_abs_err_vs_paper": abs(p.get("CheapBetterLGBMClassifier", 0) - 68.4),
"partitionwise_abs_err_vs_paper": abs(p.get("PartitionWise", 0) - 38.3),
"ordering_lightgbm_above_partitionwise":
bool(p.get("CheapBetterLGBMClassifier", 0) > p.get("PartitionWise", 0))}
print(f"claim2 {label}: LightGBM {p.get('CheapBetterLGBMClassifier'):.2f} "
f"PartitionWise {p.get('PartitionWise'):.2f}", flush=True)
for metric in ("L2", "KL"):
p, _ = power("results.csv", metric)
res.setdefault("other_metrics_v2", {})[metric] = p
OUT["claim2"] = res
# -------------------- claim 3: sample efficiency, ERT vs CovGap vs the truth
def claim3(ns=(500, 1000, 2000, 5000), seeds=range(10)):
rows = []
for n in ns:
acc = collections.defaultdict(list)
for s in seeds:
Xs, zs, ps = build(n, 300 + s)
Xo, zo, po = build(n, 300 + s, oracle=True)
hs = cv_predict(Xs, zs, seed=s)
ho = cv_predict(Xo, zo, seed=s)
acc["ert_std"].append(ert(hs, zs, ALPHA, "L1"))
acc["ert_ora"].append(ert(ho, zo, ALPHA, "L1"))
acc["cg_std"].append(covgap(Xs, zs, ALPHA, seed=s))
acc["cg_ora"].append(covgap(Xo, zo, ALPHA, seed=s))
acc["true_std"].append(true_ert(ps, ALPHA, "L1"))
r = {k: float(np.mean(v)) for k, v in acc.items()}
r["n_test"] = n
r["ert_error_vs_true"] = abs(r["ert_std"] - r["true_std"])
r["covgap_error_vs_true"] = abs(r["cg_std"] - r["true_std"])
r["ert_separation"] = r["ert_std"] - r["ert_ora"]
r["covgap_separation"] = r["cg_std"] - r["cg_ora"]
rows.append(r)
print(f" claim3 n={n}: true L1 {r['true_std']:.5f} | ERT {r['ert_std']:.5f} "
f"(err {r['ert_error_vs_true']:.5f}) | CovGap {r['cg_std']:.5f} "
f"(err {r['covgap_error_vs_true']:.5f}) | sep ERT {r['ert_separation']:.5f} "
f"vs CovGap {r['covgap_separation']:.5f}", flush=True)
last = rows[-1]
OUT["claim3"] = {"rows": rows, "n_seeds": len(list(seeds)),
"ert_closer_to_truth_at_all_n":
bool(all(r["ert_error_vs_true"] < r["covgap_error_vs_true"] for r in rows)),
"separation_ratio_at_max_n":
last["ert_separation"] / max(last["covgap_separation"], 1e-12)}
# ------------------------------------- claim 4: over/under-coverage split
def build_scaled(n_test, seed, scale, n_tr=4000, n_cal=3000):
"""Standard split conformal with the radius deliberately rescaled.
scale > 1 makes the predictor systematically conservative (over-coverage),
scale < 1 systematically aggressive (under-coverage). The true one-sided
deviations are then known, so the decomposition can be checked for
attributing the error to the correct side, not merely for being non-zero.
"""
rng = np.random.default_rng(seed)
Xtr, ytr = sample(n_tr, D, rng)
Xcal, ycal = sample(n_cal, D, rng)
Xte, yte = sample(n_test, D, rng)
m, q = split_conformal(Xtr, ytr, Xcal, ycal, ALPHA, seed=seed)
q = q * scale
yhat = m.predict(Xte)
z = coverage_indicator(yte, yhat, q)
p = true_coverage(Xte, yhat, q)
return Xte, z, p
def claim4(n_test=6000, seeds=range(8)):
res = {"n_test": n_test, "seeds": len(list(seeds)), "scenarios": {}}
for label, scale in (("conservative (radius x1.35)", 1.35),
("aggressive (radius x0.75)", 0.75),
("standard (radius x1)", 1.0)):
acc = collections.defaultdict(list)
for s in seeds:
X, z, p = build_scaled(n_test, 700 + s, scale)
h = cv_predict(X, z, seed=s)
for k in ("L1", "KL"):
o, u = ert_signed(h, z, ALPHA, k)
to, tu = true_ert_signed(p, ALPHA, k)
acc[f"{k}_plus_est"].append(o); acc[f"{k}_minus_est"].append(u)
acc[f"{k}_plus_true"].append(to); acc[f"{k}_minus_true"].append(tu)
acc["marginal_coverage"].append(float(z.mean()))
r = {k: float(np.mean(v)) for k, v in acc.items()}
r["scale"] = scale
r["L1_ratio_plus_over_minus"] = r["L1_plus_est"] / max(r["L1_minus_est"], 1e-9)
r["dominant_side_est"] = "plus" if r["L1_plus_est"] > r["L1_minus_est"] else "minus"
r["dominant_side_true"] = "plus" if r["L1_plus_true"] > r["L1_minus_true"] else "minus"
r["side_attributed_correctly"] = r["dominant_side_est"] == r["dominant_side_true"]
res["scenarios"][label] = r
print(f" claim4 {label}: cov {r['marginal_coverage']:.4f} | "
f"L1+ {r['L1_plus_est']:.5f} (true {r['L1_plus_true']:.5f}) | "
f"L1- {r['L1_minus_est']:.5f} (true {r['L1_minus_true']:.5f}) | "
f"ratio {r['L1_ratio_plus_over_minus']:.2f}", flush=True)
res["all_sides_attributed_correctly"] = all(
v["side_attributed_correctly"] for v in res["scenarios"].values())
c = res["scenarios"]["conservative (radius x1.35)"]
a = res["scenarios"]["aggressive (radius x0.75)"]
res["separation_conservative_vs_aggressive"] = (
c["L1_ratio_plus_over_minus"] / max(a["L1_ratio_plus_over_minus"], 1e-9))
OUT["claim4"] = res
# ------------------------- claim 5: classification decomposition, released rows
def claim5():
rows = list(csv.DictReader(open("authors/results_classification.csv")))
resid = []
agg = collections.defaultdict(lambda: collections.defaultdict(list))
for r in rows:
kl = float(r["ERT_logloss"])
up = float(r["ERT_underconfident_logloss"])
ov = float(r["ERT_overconfident_logloss"])
resid.append(abs(kl - (up + ov)))
key = (r["dataset"], r["method"])
agg[key]["KL"].append(kl)
agg[key]["KL_plus"].append(up) # underconfident -> conservative -> l_+
agg[key]["KL_minus"].append(ov) # overconfident -> aggressive -> l_-
per = {}
for k, v in agg.items():
per[f"{k[0]}|{k[1]}"] = {kk: [float(np.mean(vv)), float(np.std(vv, ddof=1))]
for kk, vv in v.items()}
diverge = sum(1 for v in per.values()
if abs(v["KL_plus"][0] - v["KL_minus"][0])
> 0.3 * max(abs(v["KL_plus"][0]), abs(v["KL_minus"][0]), 1e-12))
OUT["claim5"] = {
"n_rows": len(rows),
"decomposition_max_residual": float(max(resid)),
"per_dataset_method": per,
"n_cells": len(per),
"n_cells_with_divergent_components": diverge,
"naming": ("the released columns are 'underconfident'/'overconfident'; "
"an underconfident predictor makes sets that are too wide, "
"i.e. over-coverage, which is the paper's l_+ component"),
}
print(f"claim5: {len(rows)} rows, decomposition residual {max(resid):.2e}, "
f"{diverge}/{len(per)} cells divergent", flush=True)
# ------------------------------- claim 6: Algorithm 1 cross-fitting is needed
def claim6(n_test=1000, seeds=range(8)):
acc = collections.defaultdict(list)
for s in seeds:
X, z, p = build(n_test, 900 + s)
tru = true_ert(p, ALPHA, "L1")
acc["true"].append(tru)
acc["in_sample_tree"].append(ert(in_sample_predict(X, z, seed=s, model="tree"),
z, ALPHA, "L1"))
acc["cv_tree"].append(ert(cv_predict(X, z, seed=s, model="tree"), z, ALPHA, "L1"))
acc["cv_hgb"].append(ert(cv_predict(X, z, seed=s), z, ALPHA, "L1"))
# oracle scenario: cross-fitting must return ~0, in-sample must not
Xo, zo, po = build(n_test, 900 + s, oracle=True)
acc["oracle_in_sample_tree"].append(
ert(in_sample_predict(Xo, zo, seed=s, model="tree"), zo, ALPHA, "L1"))
acc["oracle_cv_tree"].append(
ert(cv_predict(Xo, zo, seed=s, model="tree"), zo, ALPHA, "L1"))
res = {k: float(np.mean(v)) for k, v in acc.items()}
res["n_test"] = n_test
res["in_sample_inflation_factor"] = res["in_sample_tree"] / max(res["cv_tree"], 1e-12)
res["oracle_in_sample_false_positive"] = res["oracle_in_sample_tree"]
res["oracle_cv_is_near_zero"] = bool(abs(res["oracle_cv_tree"]) < 0.02)
OUT["claim6"] = res
print("claim6:", {k: round(v, 5) for k, v in res.items() if isinstance(v, float)}, flush=True)
if __name__ == "__main__":
import sys
fns = {"1": claim1, "2": claim2, "3": claim3, "4": claim4, "5": claim5, "6": claim6}
for k in (sys.argv[1:] or ["2", "5", "1", "6", "4", "3"]):
print("=== claim", k, flush=True)
fns[k]()
json.dump(OUT, open("outputs/results.json", "w"), indent=2)
print("saved outputs/results.json")