ctokx's picture
Add scripts/
0f27fb6 verified
Raw
History Blame Contribute Delete
9.26 kB
"""Document-level K-fold cross-validation with error bars.
The headline numbers in this repo come from a single 70/15/15 document split.
Two models 0.009 macro-F1 apart on one split is inside the noise of a
151-document corpus, so no ranking claim survives without a variance estimate.
This harness answers the only question that licenses a "beats X" claim:
across independent, leak-free document folds, is the gap larger than its
own spread?
Design
------
* ``GroupKFold`` over ``doc_title`` assigns each of the 151 source reports to
exactly one test fold — no report is ever split across train and test.
* Inside each fold's training documents, a document-disjoint dev set is carved
off for threshold and blend-weight tuning. Test is never touched during
tuning.
* macro-F1 is averaged over the techniques that actually have test support in
that fold, so a technique that happens to land entirely in train does not
drag every model's score toward zero. The same label set is used for every
model within a fold, so the comparison stays fair.
python scripts/06_cv.py --folds 5 --model modernbert
ModernBERT is retrained from scratch inside every fold. On the 4060 this is a
few minutes per fold; the whole run is well under an hour.
"""
import argparse
import gc
import json
import shutil
import sys
import tempfile
import time
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
import numpy as np # noqa: E402
import torch # noqa: E402
from sklearn.model_selection import GroupKFold # noqa: E402
from torch.utils.data import DataLoader # noqa: E402
from cti_attack import baselines, config, data, evaluate, modeling # noqa: E402
def macro_f1_supported(Yte: np.ndarray, Ypred: np.ndarray) -> tuple[float, int]:
"""Mean F1 over the labels that have >=1 positive in this fold's test set."""
supported = np.where(Yte.sum(axis=0) > 0)[0]
fs = []
for j in supported:
yt, yp = Yte[:, j], Ypred[:, j]
tp = int(((yt == 1) & (yp == 1)).sum())
fp = int(((yt == 0) & (yp == 1)).sum())
fn = int(((yt == 1) & (yp == 0)).sum())
p = tp / (tp + fp) if tp + fp else 0.0
r = tp / (tp + fn) if tp + fn else 0.0
fs.append(2 * p * r / (p + r) if p + r else 0.0)
return float(np.mean(fs)) if fs else 0.0, len(supported)
def carve_dev(records, seed, dev_doc_frac=0.2):
"""Split records into (train, dev) by whole documents, dev-disjoint."""
docs = sorted({r["doc_title"] for r in records})
rng = np.random.RandomState(seed)
rng.shuffle(docs)
n_dev = max(1, int(len(docs) * dev_doc_frac))
dev_docs = set(docs[:n_dev])
tr = [r for r in records if r["doc_title"] not in dev_docs]
dv = [r for r in records if r["doc_title"] in dev_docs]
return tr, dv
def eval_regimes(Ydv, s_dv, Yte, s_te):
"""Tune global + per-class thresholds on dev, report both on test."""
gt, _ = evaluate.tune_global_threshold(Ydv, s_dv)
pct = evaluate.tune_per_class_thresholds(Ydv, s_dv)
g, _ = macro_f1_supported(Yte, (s_te >= gt).astype(np.int8))
p, _ = macro_f1_supported(Yte, evaluate.apply_thresholds(s_te, pct))
return g, p
def pick_alpha(Ydv, t_dv, b_dv):
"""Choose blend weight on dev: scores = alpha*tfidf + (1-alpha)*bert."""
best_a, best_f = 0.5, -1.0
for i in range(21):
a = round(0.05 * i, 2)
_, f = evaluate.tune_global_threshold(Ydv, a * t_dv + (1 - a) * b_dv)
if f > best_f:
best_a, best_f = a, f
return best_a
def run_fold(fold, tr, dv, te, labels, model_key, epochs):
Ytr = evaluate.to_matrix(tr, labels)
Ydv = evaluate.to_matrix(dv, labels).astype("int8")
Yte = evaluate.to_matrix(te, labels).astype("int8")
txt = lambda recs: [r["sentence"] for r in recs]
# ---- TF-IDF ------------------------------------------------------------
s = baselines.tfidf_lr_scores(txt(tr), Ytr, {"dev": txt(dv), "test": txt(te)})
t_dv, t_te = s["dev"], s["test"]
tf_g, tf_p = eval_regimes(Ydv, t_dv, Yte, t_te)
# ---- encoder -----------------------------------------------------------
tmp = Path(tempfile.mkdtemp(prefix=f"cvfold{fold}_"))
try:
modeling.train(model_key, "cv", tr, dv, labels, tmp, epochs=epochs)
model, tok = modeling.load_for_inference(tmp)
dev_ = modeling.device()
model.to(dev_)
amp = model_key not in config.FP32_ONLY_MODELS
ds_dv = modeling.SentenceDataset(dv, labels, tok, config.MAX_LENGTH)
ds_te = modeling.SentenceDataset(te, labels, tok, config.MAX_LENGTH)
b_dv = modeling.predict_scores(model, DataLoader(ds_dv, batch_size=32), dev_, amp=amp)
b_te = modeling.predict_scores(model, DataLoader(ds_te, batch_size=32), dev_, amp=amp)
del model
gc.collect()
torch.cuda.empty_cache()
finally:
shutil.rmtree(tmp, ignore_errors=True)
bt_g, bt_p = eval_regimes(Ydv, b_dv, Yte, b_te)
# ---- ensemble ----------------------------------------------------------
alpha = pick_alpha(Ydv, t_dv, b_dv)
en_dv = alpha * t_dv + (1 - alpha) * b_dv
en_te = alpha * t_te + (1 - alpha) * b_te
en_g, en_p = eval_regimes(Ydv, en_dv, Yte, en_te)
return {
"fold": fold,
"n_train": len(tr), "n_dev": len(dv), "n_test": len(te),
"alpha": alpha,
"tfidf_global": tf_g, "tfidf_perclass": tf_p,
f"{model_key}_global": bt_g, f"{model_key}_perclass": bt_p,
"ensemble_global": en_g, "ensemble_perclass": en_p,
}
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--folds", type=int, default=5)
ap.add_argument("--model", default="modernbert", choices=list(config.BASE_MODELS))
ap.add_argument("--epochs", type=int, default=config.EPOCHS)
ap.add_argument("--include-single", action="store_true",
help="merge single_label.json before splitting (Task 6 data lever)")
args = ap.parse_args()
records, labels, _ = data.build(verbose=False, include_single=args.include_single)
groups = [r["doc_title"] for r in records]
gkf = GroupKFold(n_splits=args.folds)
print(f"{'=' * 62}\n document-level {args.folds}-fold CV "
f"(tfidf_lr / {args.model} / ensemble)\n{'=' * 62}")
print(f" records={len(records)} documents={len(set(groups))} labels={len(labels)}\n")
fold_rows = []
t0 = time.time()
for fold, (trval_idx, test_idx) in enumerate(gkf.split(records, groups=groups), 1):
trval = [records[i] for i in trval_idx]
te = [records[i] for i in test_idx]
tr, dv = carve_dev(trval, seed=config.SPLIT_SEED + fold)
print(f" --- fold {fold}/{args.folds} "
f"train={len(tr)} dev={len(dv)} test={len(te)} "
f"({len(set(g['doc_title'] for g in te))} test docs) ---")
row = run_fold(fold, tr, dv, te, labels, args.model, args.epochs)
fold_rows.append(row)
print(f" tfidf={row['tfidf_perclass']:.4f} "
f"{args.model}={row[args.model + '_perclass']:.4f} "
f"ensemble={row['ensemble_perclass']:.4f} (alpha={row['alpha']:.2f}) "
f"[{(time.time() - t0) / 60:.1f} min elapsed]\n")
# ---- aggregate ---------------------------------------------------------
keys = ["tfidf_global", "tfidf_perclass",
f"{args.model}_global", f"{args.model}_perclass",
"ensemble_global", "ensemble_perclass"]
summary = {}
print(f"{'=' * 62}\n {args.folds}-fold CV summary — macro-F1 mean +/- std\n{'=' * 62}")
for k in keys:
vals = np.array([r[k] for r in fold_rows])
summary[k] = {"mean": round(float(vals.mean()), 4),
"std": round(float(vals.std(ddof=1)), 4),
"folds": [round(float(v), 4) for v in vals]}
print(f" {k:26} {vals.mean():.4f} +/- {vals.std(ddof=1):.4f}")
# paired ensemble-vs-tfidf gap and its spread
ens = np.array([r["ensemble_perclass"] for r in fold_rows])
tf = np.array([max(r["tfidf_perclass"], r["tfidf_global"]) for r in fold_rows])
gap = ens - tf
print(f"\n paired gap (ensemble_perclass - best_tfidf), per fold: "
f"{[round(float(g), 4) for g in gap]}")
print(f" mean gap = {gap.mean():+.4f} +/- {gap.std(ddof=1):.4f} "
f"(min {gap.min():+.4f})")
tag = f"{args.model}_single" if args.include_single else args.model
out = config.RESULTS_DIR / f"cv_{args.folds}fold__{tag}.json"
out.write_text(json.dumps({
"folds": args.folds,
"model": args.model,
"include_single_label": args.include_single,
"n_records": len(records),
"n_documents": len(set(groups)),
"per_fold": fold_rows,
"summary": summary,
"ensemble_vs_tfidf_gap": {
"per_fold": [round(float(g), 4) for g in gap],
"mean": round(float(gap.mean()), 4),
"std": round(float(gap.std(ddof=1)), 4),
"min": round(float(gap.min()), 4),
},
}, indent=2), encoding="utf-8")
print(f"\n -> {out.relative_to(config.REPO_ROOT)} "
f"[total {(time.time() - t0) / 60:.1f} min]")
if __name__ == "__main__":
main()