"""
hyperparameter.py — numerical companion to the glossary entry
'hyperparameter'.

Purpose
-------
Shows why a hyperparameter is chosen by the validation error and not by
the training error: for a polynomial model, the training error falls
with every increase of the degree, while the validation error falls and
then rises again, so only the validation error selects a degree.
Self-contained (numpy/matplotlib only), fixed seed.

Setup
-----
Training set: m = 20 data points with feature x drawn uniformly from
[0, 1] and label y = sin(2 pi x) + Gaussian noise scaled by 0.3;
validation set: 100 further points from the same source.  Model: a
polynomial of degree p in x, with the degree as the hyperparameter;
the model parameters are the coefficients, learned by least squares
(ERM with the squared error) on the training set.  Degrees 0 to
9 are compared by their average squared error on the training set
(training error) and on the validation set (validation error).

Blocks
------
[B-data]   The training set and the validation set.
[B-sweep]  Training and validation error for each degree: the training
           error never increases with the degree, the validation error
           has an interior minimum, and the degree that minimizes it
           (3) is not the degree that minimizes the training error (9).

Outputs
-------
hyperparameter_errors.csv : degree, trainerr, valerr.
hyperparameter.png        : matplotlib preview of the entry's 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(1)

# ---------------------------------------------------------------- [B-data]
m, m_val, sigma = 20, 100, 0.3
x_tr = np.sort(rng.uniform(0.0, 1.0, m))
y_tr = np.sin(2 * np.pi * x_tr) + sigma * rng.standard_normal(m)
x_va = np.sort(rng.uniform(0.0, 1.0, m_val))
y_va = np.sin(2 * np.pi * x_va) + sigma * rng.standard_normal(m_val)
check(f"[B-data]  training set of {m} and validation set of {m_val} points",
      len(x_tr) == m and len(x_va) == m_val)

# --------------------------------------------------------------- [B-sweep]
degrees = np.arange(0, 10)
tr_err, va_err = [], []
for p in degrees:
    A = np.vander(x_tr, p + 1)                     # least squares on the trainset
    w = np.linalg.lstsq(A, y_tr, rcond=None)[0]
    tr_err.append(float(np.mean((y_tr - A @ w) ** 2)))
    va_err.append(float(np.mean((y_va - np.vander(x_va, p + 1) @ w) ** 2)))
tr_err, va_err = np.array(tr_err), np.array(va_err)
p_tr, p_va = int(degrees[np.argmin(tr_err)]), int(degrees[np.argmin(va_err)])
print("  degree -> training error / validation error: " +
      ", ".join(f"{p}: {a:.3f}/{b:.3f}" for p, a, b in zip(degrees, tr_err, va_err)))
check("[B-sweep] the training error never increases with the degree",
      np.all(np.diff(tr_err) <= 1e-12))
check(f"[B-sweep] the validation error has an interior minimum at degree "
      f"{p_va}", 0 < p_va < degrees[-1])
check(f"[B-sweep] the training error would pick degree {p_tr}, the "
      f"validation error degree {p_va}", p_tr == 9 and p_va == 3)

# ---------------------------------------------------------------- CSV
with open(OUT_DIR / "hyperparameter_errors.csv", "w") as fh:
    fh.write("degree,trainerr,valerr\n")
    for p, a, b in zip(degrees, tr_err, va_err):
        fh.write(f"{p},{a:.4f},{b:.4f}\n")

# -------------------------------------------------------------- preview
fig, ax = plt.subplots(figsize=(5.0, 3.6))
ax.plot(degrees, tr_err, "k-", marker="o", ms=4, label="training error")
ax.plot(degrees, va_err, "k--", marker="s", ms=4, mfc="none",
        label="validation error")
ax.axvline(p_va, color="0.6", ls=":", lw=1)
ax.set_yscale("log")
ax.set_xlabel("polynomial degree (the hyperparameter)")
ax.set_ylabel("average squared error")
ax.set_title("the validation error selects the degree, the training error cannot")
ax.legend(frameon=False, fontsize=8)
fig.tight_layout()
fig.savefig(OUT_DIR / "hyperparameter.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 / 'hyperparameter_errors.csv'}, "
      f"{OUT_DIR / 'hyperparameter.png'}")
if n_ok != len(report):
    raise SystemExit(1)
