"""
mlpipeline.py — numerical companion to the glossary entry 'ML pipeline'.

Purpose
-------
Shows why a stage of an ML pipeline that has its own parameters must be
fitted on the training set only.  A hand-designed feature extraction
stage (keeping the features that agree best with the label) fitted on
all available data leaks information from the validation set into the
learned hypothesis, so that the validation error underestimates the
risk; the same stage fitted on the training set only gives a validation
error that matches the risk.  Self-contained (numpy and matplotlib
only), fixed seed.

Setup
-----
Raw data: 500 features per data point, all drawn independently of the
label, so that no hypothesis can do better than the risk 1.0 (the
label's variance).  Training set of 100 data points, validation set of
100, and a test set of 10000 data points on which the risk of a learned
hypothesis is estimated.  Pipeline: feature extraction (keep the k
features whose agreement with the label is largest), then a linear
model fitted on the training set, then prediction.  The feature
extraction stage is fitted either on training and validation set
together (leaked) or on the training set only (proper), for
k = 1, 2, 5, 10, 20; every number is an average over 20 random draws
of the data.

Blocks
------
[B-leaked]  Feature extraction fitted on all 200 labelled data points:
            the validation error of the pipeline drops below 0.9 for
            k = 20, while the risk of the same hypothesis exceeds 1.4.
[B-proper]  Feature extraction fitted on the training set only: the
            validation error stays within 0.1 of the risk for every k.
[B-compare] The amount by which the leaked validation error
            underestimates the risk grows with k and exceeds 0.3 for
            k >= 10; the proper pipeline's never exceeds 0.1.

Outputs
-------
mlpipeline_errors.csv : per k, the validation error and the risk of the
                        leaked pipeline and of the proper pipeline.
mlpipeline.png        : matplotlib preview of that figure (checking only).
"""
# /// script
# requires-python = ">=3.10"
# dependencies = ["numpy==2.3.4", "matplotlib==3.10.7"]
# ///


import numpy as np
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt

from pathlib import Path

OUT_DIR = Path(__file__).parent

report = []


def check(name, ok):
    report.append((name, bool(ok)))
    print(f"  [{'ok' if ok else 'FAIL'}] {name}")


rng = np.random.default_rng(0)

D, M_TRAIN, M_VAL, M_TEST, REPS = 500, 100, 100, 10000, 20
KS = [1, 2, 5, 10, 20]
tr = slice(0, M_TRAIN); va = slice(M_TRAIN, M_TRAIN + M_VAL)
te = slice(M_TRAIN + M_VAL, None)


def fit_extraction(Xf, yf, k):
    """Stage 1: keep the k features agreeing best with the label."""
    score = np.abs((Xf - Xf.mean(0)).T @ (yf - yf.mean())) / len(yf)
    return np.argsort(score)[::-1][:k]


def fit_model(Xf, yf):
    """Stage 2: linear model with intercept, least squares."""
    A = np.c_[Xf, np.ones(len(yf))]
    return np.linalg.lstsq(A, yf, rcond=None)[0]


def predict(w, Xf):
    return np.c_[Xf, np.ones(len(Xf))] @ w


def run(X, y, k, extraction_rows):
    cols = fit_extraction(X[extraction_rows], y[extraction_rows], k)
    w = fit_model(X[tr][:, cols], y[tr])              # model: training set only
    valerr = float(np.mean((y[va] - predict(w, X[va][:, cols])) ** 2))
    risk = float(np.mean((y[te] - predict(w, X[te][:, cols])) ** 2))
    return valerr, risk


all_rows = np.arange(M_TRAIN + M_VAL)
train_rows = np.arange(M_TRAIN)
acc = np.zeros((len(KS), 4))                            # vl, rl, vp, rp
for _ in range(REPS):
    X = rng.standard_normal((M_TRAIN + M_VAL + M_TEST, D))
    y = rng.standard_normal(M_TRAIN + M_VAL + M_TEST)   # independent of X
    for i, k in enumerate(KS):
        acc[i, 0:2] += run(X, y, k, all_rows)
        acc[i, 2:4] += run(X, y, k, train_rows)
acc /= REPS
leaked = [(acc[i, 0], acc[i, 1]) for i in range(len(KS))]
proper = [(acc[i, 2], acc[i, 3]) for i in range(len(KS))]

# --------------------------------------------------------------- [B-leaked]
v20, r20 = leaked[KS.index(20)]
check(f"[B-leaked]  extraction fitted on all 200 points, k=20: validation "
      f"error {v20:.2f}, risk {r20:.2f}", v20 < 0.9 and r20 > 1.4)

# --------------------------------------------------------------- [B-proper]
gap_proper = max(abs(v - r) for v, r in proper)
check(f"[B-proper]  extraction fitted on the training set only: validation "
      f"error within {gap_proper:.2f} of the risk for every k",
      gap_proper < 0.1)

# -------------------------------------------------------------- [B-compare]
under_leaked = [r - v for v, r in leaked]
under_proper = [r - v for v, r in proper]
check("[B-compare] risk minus validation error per k: leaked "
      + ", ".join(f"{u:.2f}" for u in under_leaked) + "; proper "
      + ", ".join(f"{u:.2f}" for u in under_proper),
      all(a < b for a, b in zip(under_leaked, under_leaked[1:]))
      and all(u > 0.3 for u, k in zip(under_leaked, KS) if k >= 10)
      and all(abs(u) < 0.1 for u in under_proper))

# ---------------------------------------------------------------- CSV
with open(OUT_DIR / "mlpipeline_errors.csv", "w") as fh:
    fh.write("k,valerr_leaked,risk_leaked,valerr_proper,risk_proper\n")
    for k, (vl, rl), (vp, rp) in zip(KS, leaked, proper):
        fh.write(f"{k},{vl:.4f},{rl:.4f},{vp:.4f},{rp:.4f}\n")

# -------------------------------------------------------------- preview
fig, ax = plt.subplots(figsize=(5.6, 3.8))
ax.plot(KS, [v for v, _ in leaked], "ks--", ms=5, label="validation error, extraction fitted on all data")
ax.plot(KS, [r for _, r in leaked], "ks:", ms=5, mfc="none", label="risk of that hypothesis")
ax.plot(KS, [v for v, _ in proper], "ko-", ms=5, label="validation error, extraction fitted on the training set")
ax.plot(KS, [r for _, r in proper], "ko:", ms=5, mfc="none", label="risk of that hypothesis")
ax.set_xscale("log"); ax.set_xticks(KS); ax.set_xticklabels([str(k) for k in KS])
ax.set_xlabel("number of features kept by the extraction stage, k")
ax.set_ylabel("average squared error")
ax.set_title("a stage fitted on all data makes the validation error optimistic")
ax.legend(frameon=False, fontsize=7)
fig.tight_layout()
fig.savefig(OUT_DIR / "mlpipeline.png", dpi=110)

n_ok = sum(ok for _, ok in report)
print(f"\n{n_ok}/{len(report)} checks pass")
print(f"wrote {OUT_DIR / 'mlpipeline_errors.csv'}, {OUT_DIR / 'mlpipeline.png'}")
if n_ok != len(report):
    raise SystemExit(1)
