File size: 13,676 Bytes
1b7a999 | 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 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 | """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")
|