{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Auditing Interpretability Claims for Kolmogorov-Arnold Classification Heads\n### Reproducibility notebook\n\nCompanion code for the manuscript submitted to *Knowledge-Based Systems*.\n\n**Author:** Musa Ataş (Siirt University) · musa.atas@siirt.edu.tr\n\n---\n\n## What each section produces\n\n| Section | Produces | Source | Time |\n|---|---|---|---|\n| 3 | Table 2 — APTOS class distribution | computed | <1 min |\n| 4 | Table 1 — parameter matching | computed | <1 min |\n| 7 | Table 3 — APTOS five-fold CV | released fold predictions | <1 min |\n| 8 | Table 4 — DDR five-fold CV | released fold predictions | <1 min |\n| 9 | Equivalence test (TOST) and paired Wilcoxon | computed | <1 min |\n| 12 | Table 5 — frozen-backbone probe | trains heads on cached features | ~35 min |\n| 13 | Table 6 — Messidor-2, five folds, calibration, prevalence | released checkpoints | ~8 min |\n| 14 | Figure 4 — feature-to-class-logit curves | released checkpoints | ~2 min |\n| 15 | Figure 5 and the sign-consistency analysis | released checkpoints | ~6 min |\n| 16 | Tables 7–8 and Figure 6 — same-model deletion | released checkpoints | ~20 min |\n| 17 | Table 9 — removal-operator invariance | released checkpoints | ~10 min |\n| 18 | Table 10 — remove-and-retrain (ROAR) | released checkpoints | ~25 min |\n\nTable 6 in the manuscript (comparison with published methods) is drawn from the\nliterature and is not computed here.\n\n## How to run\n\nAttach these four public inputs, then **Run All**:\n\n1. `hakmesyo/dr-kan-mlp-aptos` — APTOS MobileNetV3-Large checkpoints and fold predictions\n2. `cezeriotonom/dr-kan-mlp-results` — DDR fold predictions\n3. `APTOS 2019 Blindness Detection` (Kaggle competition)\n4. `messidor-2` (Kaggle dataset)\n\nSections 1–11 and 13–16 finish in well under an hour. Sections 12, 17 and 18 are the\naudit's expensive controls; set `RUN_HEAVY = False` in Section 1 to skip them and\nreproduce everything else.\n\n## Reproducibility caveats, stated plainly\n\n* **The KAN head is not bit-reproducible.** Under a fixed seed the MLP head reproduces\n  exactly; the KAN head varies by roughly ±0.005 QWK, most plausibly from\n  non-deterministic operations in the B-spline grid. Retraining will therefore not\n  reproduce the fourth decimal of Table 3. No conclusion in the paper turns on\n  differences of that size.\n* **KernelSHAP is a sampling estimator.** Its rows in Tables 7–9 shift in the third\n  decimal between runs. Its qualitative behaviour — indistinguishable from a random\n  control at practical budgets — does not. Partial-dependence and Integrated Gradients\n  rankings are deterministic and reproduce exactly.\n* **DDR uses a 50% class-stratified subsample** (n = 6,260) with early-stopping patience\n  5, to keep the five-fold DDR grid inside Kaggle's session limits. This is stated in the\n  manuscript (Section 3.3, Table 4 caption).\n* **Only the MobileNetV3-Large fold predictions are released for APTOS.** Section 7\n  therefore reproduces that backbone's rows of Table 3 exactly; the EfficientNet-B3 and\n  ResNet-50 rows came from an earlier session whose fold predictions were not retained.\n  MobileNetV3-Large is the backbone used for every interpretability analysis, so\n  Sections 12–18 are unaffected. The frozen-backbone probe of Section 12 needs no\n  checkpoints and does cover all three backbones.\n* **Tables 3, 4 and 5 are recomputed here**, not copied from stored CSVs.\n","metadata":{}},{"cell_type":"markdown","source":"## 1. Environment and configuration","metadata":{}},{"cell_type":"code","source":"import os, random, json, math, time, copy, warnings\nfrom pathlib import Path\nfrom collections import defaultdict\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, Subset\n\nimport torchvision\nfrom torchvision import transforms\nfrom torchvision.models import (\n    mobilenet_v3_large, MobileNet_V3_Large_Weights,\n    efficientnet_b3, EfficientNet_B3_Weights,\n    resnet50, ResNet50_Weights,\n)\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import (cohen_kappa_score, f1_score, recall_score,\n                             roc_auc_score, roc_curve, confusion_matrix)\nfrom scipy import stats\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\n\ndef set_seed(seed: int = SEED):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {DEVICE}\")\nprint(f\"torch: {torch.__version__} | torchvision: {torchvision.__version__}\")\n\nOUT_DIR = Path(\"/kaggle/working/results\")\nCKPT_DIR = Path(\"/kaggle/working/checkpoints\")\nOUT_DIR.mkdir(parents=True, exist_ok=True)\nCKPT_DIR.mkdir(parents=True, exist_ok=True)\n\n\n# ============================ CONFIGURATION ============================\n# REPRODUCE_FROM_CHECKPOINTS = True  (default)\n#   Loads the released checkpoints and fold predictions from the attached public\n#   datasets. Reproduces every table and figure in ~20-30 min on one T4.\n#\n# REPRODUCE_FROM_CHECKPOINTS = False\n#   Retrains from scratch (several hours). Note the KAN head is not\n#   bit-reproducible; see the caveats at the top of this notebook.\nREPRODUCE_FROM_CHECKPOINTS = True\n\n# Backbone used for external validation and all interpretability analyses.\n# MobileNetV3-Large is the most parameter-efficient backbone and the paper's\n# interpretability showcase.\nSHOWCASE_BACKBONE = \"MobileNetV3-Large\"\n\n# DDR grid settings, reported transparently in the manuscript.\nDDR_SUBSAMPLE_FRAC = 0.5   # class-stratified subsample; 1.0 = full 12,522\nDDR_PATIENCE = 5           # DDR early-stopping patience (APTOS uses 10)\n\n# KernelSHAP sampling. Fixed seed for repeatability; the estimator remains\n# stochastic in the sense that a different seed shifts the third decimal.\nSHAP_SEED = SEED\nSHAP_BACKGROUND = 100\nSHAP_COALITIONS = 100\nSHAP_EVAL_POINTS = 200\n\ndef backbones_to_run():\n    return list(BACKBONE_SPECS)\n\nprint(f\"REPRODUCE_FROM_CHECKPOINTS = {REPRODUCE_FROM_CHECKPOINTS}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:10:01.397628Z","iopub.execute_input":"2026-07-24T17:10:01.398504Z","iopub.status.idle":"2026-07-24T17:10:16.924240Z","shell.execute_reply.started":"2026-07-24T17:10:01.398471Z","shell.execute_reply":"2026-07-24T17:10:16.923407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Audit configuration.\n#\n# RUN_HEAVY = True   reproduces the three expensive controls:\n#                      Section 12 (frozen-backbone probe, ~35 min)\n#                      Section 17 (removal-operator invariance, ~10 min)\n#                      Section 18 (remove-and-retrain, ~25 min)\n# RUN_HEAVY = False\nRUN_HEAVY = True\n\n# Equivalence margin for the TOST procedure, in QWK. Motivated by the KAN head's\n# own run-to-run variability under a fixed seed (Section 6.5 of the manuscript).\nTOST_MARGIN = 0.015\n\nprint(f\"RUN_HEAVY = {RUN_HEAVY} | TOST_MARGIN = {TOST_MARGIN}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:10:22.401421Z","iopub.execute_input":"2026-07-24T17:10:22.402198Z","iopub.status.idle":"2026-07-24T17:10:22.408282Z","shell.execute_reply.started":"2026-07-24T17:10:22.402166Z","shell.execute_reply":"2026-07-24T17:10:22.406832Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Datasets and preprocessing","metadata":{}},{"cell_type":"code","source":"IMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\n\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2),\n    transforms.RandomApply([transforms.GaussianBlur(kernel_size=(3, 7))], p=0.3),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\neval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\n\nclass DRDataset(Dataset):\n    '''Generic DR-grading image dataset (5-class, or binary for Messidor-2).'''\n\n    def __init__(self, df, image_dir, image_col, label_col, transform, ext=\"\"):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = Path(image_dir)\n        self.image_col = image_col\n        self.label_col = label_col\n        self.transform = transform\n        self.ext = ext\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = self.image_dir / f\"{row[self.image_col]}{self.ext}\"\n        image = torchvision.io.read_image(str(img_path), mode=torchvision.io.ImageReadMode.RGB)\n        image = transforms.ToPILImage()(image)\n        image = self.transform(image)\n        label = int(row[self.label_col])\n        return image, label\n\n\ndef load_aptos2019():\n    # Kaggle mounts competition data either directly or under /competitions/.\n    candidates = [\n        \"/kaggle/input/aptos2019-blindness-detection\",\n        \"/kaggle/input/competitions/aptos2019-blindness-detection\",\n    ]\n    root = next((Path(c) for c in candidates if (Path(c) / \"train.csv\").is_file()), None)\n    if root is None:\n        raise FileNotFoundError(\"APTOS train.csv not found under /kaggle/input\")\n    df = pd.read_csv(root / \"train.csv\")\n    return dict(df=df, image_dir=root / \"train_images\",\n                image_col=\"id_code\", label_col=\"diagnosis\", ext=\".png\", n_classes=5)\n\n\ndef load_ddr():\n    import glob, os\n    csvs = glob.glob(\"/kaggle/input/**/DR_grading.csv\", recursive=True)\n    if not csvs:\n        raise FileNotFoundError(\"DDR DR_grading.csv not found under /kaggle/input\")\n    csv_path = Path(csvs[0])\n    df = pd.read_csv(csv_path)\n    # Normalise column names (this DDR mirror uses APTOS-style id_code/diagnosis).\n    if \"diagnosis\" not in df.columns and \"label\" in df.columns:\n        df = df.rename(columns={\"label\": \"diagnosis\"})\n    if \"id_code\" not in df.columns and \"image\" in df.columns:\n        df = df.rename(columns={\"image\": \"id_code\"})\n    # Already-gradable mirror; drop any remaining ungradable rows (grade 5).\n    df = df[df[\"diagnosis\"] != 5].reset_index(drop=True)\n    # Optional class-stratified subsample (speed; reported in the manuscript).\n    if \"DDR_SUBSAMPLE_FRAC\" in globals() and DDR_SUBSAMPLE_FRAC < 1.0:\n        df = (df.groupby(\"diagnosis\", group_keys=False)\n                .sample(frac=DDR_SUBSAMPLE_FRAC, random_state=SEED)\n                .reset_index(drop=True))\n        print(f\"[DDR subsample] stratified frac={DDR_SUBSAMPLE_FRAC} -> {len(df)} images\")\n    # Auto-locate the image directory (filenames in id_code already include .jpg;\n    # some mirrors nest images under DR_grading/DR_grading/).\n    sample = str(df.iloc[0][\"id_code\"])\n    image_dir = None\n    for cand in glob.glob(str(csv_path.parent / \"**\"), recursive=True):\n        if os.path.isdir(cand) and os.path.isfile(os.path.join(cand, sample)):\n            image_dir = Path(cand); break\n    if image_dir is None:\n        raise FileNotFoundError(f\"DDR image folder containing {sample} not found\")\n    print(f\"DDR: {len(df)} gradable images | images at {image_dir}\")\n    print(df[\"diagnosis\"].value_counts().sort_index().to_string())\n    return dict(df=df, image_dir=image_dir, image_col=\"id_code\", label_col=\"diagnosis\",\n                ext=\"\", n_classes=5)\n\n\ndef load_messidor2():\n    candidates = [\n        \"/kaggle/input/messidor2\",\n        \"/kaggle/input/datasets/xyaustin/messidor2/messidor-2\",\n        \"/kaggle/input/messidor-2\",\n    ]\n    root = next((Path(c) for c in candidates if (Path(c) / \"messidor_data.csv\").is_file()), None)\n    if root is None:\n        raise FileNotFoundError(\"Messidor-2 messidor_data.csv not found under /kaggle/input\")\n    df = pd.read_csv(root / \"messidor_data.csv\")\n    # Keep only gradable images with a valid adjudicated grade.\n    if \"adjudicated_gradable\" in df.columns:\n        df = df[df[\"adjudicated_gradable\"] == 1]\n    df = df.dropna(subset=[\"adjudicated_dr_grade\"]).reset_index(drop=True)\n    df[\"adjudicated_dr_grade\"] = df[\"adjudicated_dr_grade\"].astype(int)\n    # Binary referral protocol: grades 0-1 -> non-referable (0), >=2 -> referable (1).\n    df[\"referable\"] = (df[\"adjudicated_dr_grade\"] >= 2).astype(int)\n    print(f\"Messidor-2: {len(df)} gradable images \"\n          f\"(referable={int(df['referable'].sum())}, non-referable={int((df['referable']==0).sum())})\")\n    # Image filenames already include their extension in `image_id`, so ext=\"\".\n    return dict(df=df, image_dir=root / \"images\", image_col=\"image_id\",\n                label_col=\"referable\", ext=\"\", n_classes=2)\n\n\n\n# ---------------------------------------------------------------------\n# One-time image cache (SPEED, result-preserving).\n# The training/eval pipeline's FIRST op is Resize((224,224)); the original\n# images (APTOS/DDR up to ~2000px) are otherwise decoded+resized every epoch,\n# which is the I/O bottleneck. We pre-resize each image once to 224x224\n# (PIL BILINEAR, identical to transforms.Resize) and cache it as a small PNG,\n# then point the dataset at the cache. Because Resize((224,224)) applied again\n# to a 224x224 image is a no-op, augmentations and results are unchanged.\n# ---------------------------------------------------------------------\nfrom PIL import Image\nfrom concurrent.futures import ThreadPoolExecutor\n\n\ndef _cache_one(task):\n    src, out, size = task\n    if not out.exists():\n        Image.open(src).convert(\"RGB\").resize((size, size), Image.BILINEAR).save(out, \"PNG\")\n\n\ndef precompute_image_cache(spec, tag, size=224, workers=4):\n    cache_dir = Path(f\"/tmp/img_cache_{tag}\")\n    cache_dir.mkdir(parents=True, exist_ok=True)\n    df = spec[\"df\"].reset_index(drop=True).copy()\n    src_dir, ext, col = Path(spec[\"image_dir\"]), spec[\"ext\"], spec[\"image_col\"]\n    tasks = [(src_dir / f\"{df.iloc[i][col]}{ext}\", cache_dir / f\"{i}.png\", size)\n             for i in range(len(df))]\n    df[\"_cache_id\"] = [str(i) for i in range(len(df))]\n    t0 = time.time()\n    with ThreadPoolExecutor(max_workers=workers) as ex:\n        for k, _ in enumerate(ex.map(_cache_one, tasks)):\n            if (k + 1) % 3000 == 0:\n                print(f\"  [{tag}] cached {k + 1}/{len(df)}\")\n    # Redirect the spec to the cache (DRDataset code is unchanged).\n    spec[\"df\"], spec[\"image_dir\"] = df, cache_dir\n    spec[\"image_col\"], spec[\"ext\"] = \"_cache_id\", \".png\"\n    print(f\"[{tag}] image cache ready: {len(df)} imgs @ {size}x{size} \"\n          f\"({time.time() - t0:.0f}s, one-time)\")\n    return spec\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:10:28.805905Z","iopub.execute_input":"2026-07-24T17:10:28.806387Z","iopub.status.idle":"2026-07-24T17:10:28.829642Z","shell.execute_reply.started":"2026-07-24T17:10:28.806360Z","shell.execute_reply":"2026-07-24T17:10:28.828988Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Table 2 — APTOS 2019 class distribution","metadata":{}},{"cell_type":"code","source":"aptos = load_aptos2019()\naptos = precompute_image_cache(aptos, \"aptos\")\nclass_names_5 = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"]\n\ndist = aptos[\"df\"][aptos[\"label_col\"]].value_counts().sort_index()\ntable2 = pd.DataFrame({\n    \"Grade\": class_names_5,\n    \"Label\": range(5),\n    \"Samples\": dist.reindex(range(5), fill_value=0).values,\n})\ntable2[\"Proportion (%)\"] = (100 * table2[\"Samples\"] / table2[\"Samples\"].sum()).round(1)\n\nprint(\"Table 2 -- APTOS 2019 class distribution (n = {})\".format(table2[\"Samples\"].sum()))\nprint(table2.to_string(index=False))\ntable2.to_csv(OUT_DIR / \"table2_aptos_distribution.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:10:36.073028Z","iopub.execute_input":"2026-07-24T17:10:36.073307Z","iopub.status.idle":"2026-07-24T17:14:20.523196Z","shell.execute_reply.started":"2026-07-24T17:10:36.073284Z","shell.execute_reply":"2026-07-24T17:14:20.522423Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Model architecture and Table 1 — parameter matching\n\nHead hidden dimensions are calibrated per backbone so that KAN and MLP heads hold approximately the same number of parameters (~100K), isolating head architecture as the sole design variable.\n\nNote the dropout asymmetry, which Section 12 quantifies: `MLPHead` drops hidden activations, `KANHead` drops the raw input vector.","metadata":{}},{"cell_type":"code","source":"# ---------------------------------------------------------------------\n# KANLinear: from-scratch implementation following Liu et al. (2024),\n# \"KAN: Kolmogorov-Arnold Networks\", with G=5 uniform grid intervals\n# and cubic (k=3) B-spline basis functions, as specified in the paper.\n# ---------------------------------------------------------------------\nclass KANLinear(nn.Module):\n    def __init__(self, in_features, out_features, grid_size=5, spline_order=3,\n                 grid_range=(-2.0, 2.0)):\n        super().__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.grid_size = grid_size\n        self.spline_order = spline_order\n\n        h = (grid_range[1] - grid_range[0]) / grid_size\n        grid = (torch.arange(-spline_order, grid_size + spline_order + 1) * h\n                + grid_range[0])\n        grid = grid.expand(in_features, -1).contiguous()\n        self.register_buffer(\"grid\", grid)\n\n        # Base (residual) weight, analogous to a standard linear layer,\n        # applied to a SiLU-activated input (as in the reference KAN impl.).\n        self.base_weight = nn.Parameter(torch.empty(out_features, in_features))\n        # Spline coefficients c_i for each of the (grid_size + spline_order)\n        # B-spline basis functions, per input-output edge.\n        self.spline_weight = nn.Parameter(\n            torch.empty(out_features, in_features, grid_size + spline_order))\n        # Per-edge scaling coefficient (\"spline_scaler\"): the structural\n        # asymmetry vs. MLP acknowledged in the manuscript (Table 1 note).\n        self.spline_scaler = nn.Parameter(torch.empty(out_features, in_features))\n\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        nn.init.kaiming_uniform_(self.base_weight, a=math.sqrt(5))\n        with torch.no_grad():\n            noise = (torch.rand(self.grid_size + 1, self.in_features, self.out_features) - 0.5) * 0.1\n            self.spline_weight.data.copy_(\n                self.curve2coeff(self.grid.T[self.spline_order:-self.spline_order], noise)\n            )\n        nn.init.kaiming_uniform_(self.spline_scaler, a=math.sqrt(5))\n\n    def b_splines(self, x):\n        # x: (batch, in_features) -> (batch, in_features, grid_size + spline_order)\n        grid = self.grid  # (in_features, grid_size + 2*spline_order + 1)\n        x = x.unsqueeze(-1)\n        bases = ((x >= grid[:, :-1]) & (x < grid[:, 1:])).to(x.dtype)\n        for k in range(1, self.spline_order + 1):\n            left = (x - grid[:, : -(k + 1)]) / (grid[:, k:-1] - grid[:, : -(k + 1)])\n            right = (grid[:, k + 1:] - x) / (grid[:, k + 1:] - grid[:, 1:-k])\n            bases = left * bases[:, :, :-1] + right * bases[:, :, 1:]\n        return bases.contiguous()\n\n    def curve2coeff(self, x, y):\n        A = self.b_splines(x).transpose(0, 1)   # x is (grid_size+1, in_features); do NOT transpose\n        B = y.transpose(0, 1)\n        sol = torch.linalg.lstsq(A, B).solution\n        return sol.permute(2, 0, 1)\n\n    @property\n    def scaled_spline_weight(self):\n        return self.spline_weight * self.spline_scaler.unsqueeze(-1)\n\n    def forward(self, x):\n        base_out = F.linear(F.silu(x), self.base_weight)\n        spline_out = F.linear(\n            self.b_splines(x).view(x.size(0), -1),\n            self.scaled_spline_weight.view(self.out_features, -1),\n        )\n        return base_out + spline_out\n\n    def edge_function(self, i, j, x_range=(-2.0, 2.0), n_points=200):\n        '''Evaluate the learned 1-D activation phi_{ij}(x) for plotting\n        (Section 4/10 -- underlies Figures 2 and 5).'''\n        xs = torch.linspace(*x_range, n_points)\n        with torch.no_grad():\n            probe = torch.zeros(n_points, self.in_features)\n            probe[:, i] = xs\n            bs = self.b_splines(probe)[:, i, :]                 # (n_points, n_basis)\n            phi = bs @ self.scaled_spline_weight[j, i, :]        # spline component only\n        return xs.numpy(), phi.numpy()\n\n\nclass MLPHead(nn.Module):\n    def __init__(self, d_in, h_mlp, n_classes=5, p_drop=0.3):\n        super().__init__()\n        self.fc1 = nn.Linear(d_in, h_mlp)\n        self.act = nn.ReLU()\n        self.drop = nn.Dropout(p_drop)\n        self.fc2 = nn.Linear(h_mlp, n_classes)\n\n    def forward(self, x):\n        return self.fc2(self.drop(self.act(self.fc1(x))))\n\n\nclass KANHead(nn.Module):\n    def __init__(self, d_in, h_kan, n_classes=5, p_drop=0.3, grid_size=5, spline_order=3):\n        super().__init__()\n        self.drop = nn.Dropout(p_drop)\n        self.kan1 = KANLinear(d_in, h_kan, grid_size=grid_size, spline_order=spline_order)\n        self.kan2 = KANLinear(h_kan, n_classes, grid_size=grid_size, spline_order=spline_order)\n\n    def forward(self, x):\n        return self.kan2(self.kan1(self.drop(x)))\n\n\nBACKBONE_SPECS = {\n    \"MobileNetV3-Large\": dict(d=1280, h_mlp=78, h_kan=8),\n    \"EfficientNet-B3\":   dict(d=1536, h_mlp=65, h_kan=7),\n    \"ResNet-50\":         dict(d=2048, h_mlp=49, h_kan=5),\n}\n\n\ndef build_backbone(name: str, pretrained: bool = True):\n    if name == \"MobileNetV3-Large\":\n        weights = MobileNet_V3_Large_Weights.DEFAULT if pretrained else None\n        m = mobilenet_v3_large(weights=weights)\n        # Keep the first classifier block (Linear 960->1280 + Hardswish) so the\n        # pooled feature vector is 1280-dim, matching Table 1 (d=1280). Drop only\n        # the final Dropout + Linear(1280->1000) ImageNet classifier head.\n        m.classifier = nn.Sequential(m.classifier[0], m.classifier[1])\n    elif name == \"EfficientNet-B3\":\n        weights = EfficientNet_B3_Weights.DEFAULT if pretrained else None\n        m = efficientnet_b3(weights=weights)\n        m.classifier = nn.Identity()\n    elif name == \"ResNet-50\":\n        weights = ResNet50_Weights.DEFAULT if pretrained else None\n        m = resnet50(weights=weights)\n        m.fc = nn.Identity()\n    else:\n        raise ValueError(name)\n    return m\n\n\nclass DRModel(nn.Module):\n    '''Backbone (global-average-pooled feature vector f in R^d) + head.'''\n\n    def __init__(self, backbone_name, head_type, n_classes=5, h_kan_override=None):\n        super().__init__()\n        assert head_type in (\"mlp\", \"kan\")\n        spec = BACKBONE_SPECS[backbone_name]\n        self.backbone_name = backbone_name\n        self.head_type = head_type\n        self.backbone = build_backbone(backbone_name)\n        self.d = spec[\"d\"]\n        if head_type == \"mlp\":\n            self.head = MLPHead(self.d, spec[\"h_mlp\"], n_classes)\n        else:\n            # h_kan_override lets the cross-fold consistency study (Fig. 6)\n            # use a wider spline basis (h_kan=64) via the SAME training path,\n            # instead of silently building a discarded wide head. Default\n            # None reproduces the parameter-matched Table 1 head exactly.\n            h_kan = spec[\"h_kan\"] if h_kan_override is None else h_kan_override\n            self.head = KANHead(self.d, h_kan, n_classes)\n\n    def extract_features(self, x):\n        '''Global-average-pooled backbone feature vector f in R^d (used\n        directly by the SHAP/faithfulness protocol in Section 12).'''\n        if \"MobileNet\" in self.backbone_name:\n            feats = self.backbone.features(x)\n            feats = self.backbone.avgpool(feats).flatten(1)   # 960 channels\n            feats = self.backbone.classifier(feats)           # -> 1280 (matches d)\n        elif \"EfficientNet\" in self.backbone_name:\n            feats = self.backbone.features(x)\n            feats = F.adaptive_avg_pool2d(feats, 1).flatten(1)  # 1536\n        else:  # ResNet-50\n            feats = self.backbone(x)                            # 2048\n        return feats\n\n    def forward(self, x):\n        f = self.extract_features(x)\n        return self.head(f)\n\n    def freeze_backbone(self, freeze: bool):\n        for p in self.backbone.parameters():\n            p.requires_grad = not freeze\n\n\ndef count_head_params(model):\n    return sum(p.numel() for p in model.head.parameters() if p.requires_grad)\n\n\n# ---- Table 1: parameter-matched classification head design ----\nrows = []\nfor name, spec in BACKBONE_SPECS.items():\n    mlp_model = DRModel(name, \"mlp\")\n    kan_model = DRModel(name, \"kan\")\n    p_mlp = count_head_params(mlp_model)\n    p_kan = count_head_params(kan_model)\n    rows.append({\n        \"Backbone\": name, \"d\": spec[\"d\"],\n        \"h_mlp\": spec[\"h_mlp\"], \"MLP params\": p_mlp,\n        \"h_kan\": spec[\"h_kan\"], \"KAN params\": p_kan,\n        \"Delta (%)\": round(100 * (p_kan - p_mlp) / p_mlp, 1),\n    })\n    del mlp_model, kan_model\n\ntable1 = pd.DataFrame(rows)\nprint(\"Table 1 -- Parameter-matched classification head design\")\nprint(table1.to_string(index=False))\ntable1.to_csv(OUT_DIR / \"table1_param_matching.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:14:43.955681Z","iopub.execute_input":"2026-07-24T17:14:43.956304Z","iopub.status.idle":"2026-07-24T17:14:47.181041Z","shell.execute_reply.started":"2026-07-24T17:14:43.956276Z","shell.execute_reply":"2026-07-24T17:14:47.180226Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training protocol\n\nDefinitions only. Training runs only when `REPRODUCE_FROM_CHECKPOINTS = False`.","metadata":{}},{"cell_type":"code","source":"def class_weights(labels, n_classes):\n    n = np.bincount(labels, minlength=n_classes).astype(np.float64)\n    n[n == 0] = 1.0\n    w = len(labels) / (n_classes * n)\n    return torch.tensor(w, dtype=torch.float32)\n\n\ndef make_loader(df, image_dir, image_col, label_col, ext, indices, transform, batch_size=16, shuffle=True):\n    ds = DRDataset(df.iloc[indices], image_dir, image_col, label_col, transform, ext)\n    return DataLoader(ds, batch_size=batch_size, shuffle=shuffle, num_workers=4,\n                      pin_memory=True, drop_last=shuffle)\n\n\ndef train_one_fold(backbone_name, head_type, dataset_spec, train_idx, val_idx,\n                   fold_id, n_classes=5, max_epochs=50, patience=10, tag=\"\",\n                   h_kan_override=None):\n    \"\"\"Train a single (backbone, head) fold.\n\n    Checkpoint selection: the manuscript's Training Protocol states that early\n    stopping is \"monitored on validation QWK\", so we keep the epoch with the\n    highest validation QWK and drive the patience counter with QWK. (The paper's\n    separate sentence about retaining the \"lowest validation loss\" checkpoint\n    contradicts this; the text must be aligned to the QWK criterion actually\n    used here.) We additionally log validation loss and per-class recall at the\n    selected epoch so either rule can be audited WITHOUT retraining.\n\n    h_kan_override: if given, the KAN head uses this hidden width instead of the\n    parameter-matched Table 1 value. Used only by the Fig. 6 consistency study\n    (h_kan=64). Default None reproduces Tables 3-5 exactly.\n    \"\"\"\n    set_seed(SEED + fold_id)\n\n    df, image_dir = dataset_spec[\"df\"], dataset_spec[\"image_dir\"]\n    image_col, label_col, ext = dataset_spec[\"image_col\"], dataset_spec[\"label_col\"], dataset_spec[\"ext\"]\n\n    train_loader = make_loader(df, image_dir, image_col, label_col, ext, train_idx, train_transform, shuffle=True)\n    val_loader   = make_loader(df, image_dir, image_col, label_col, ext, val_idx, eval_transform, shuffle=False)\n\n    model = DRModel(backbone_name, head_type, n_classes=n_classes,\n                    h_kan_override=h_kan_override).to(DEVICE)\n    weights = class_weights(df.iloc[train_idx][label_col].values, n_classes).to(DEVICE)\n    criterion = nn.CrossEntropyLoss(weight=weights)\n\n    optimizer = torch.optim.AdamW([\n        {\"params\": model.backbone.parameters(), \"lr\": 5e-5},\n        {\"params\": model.head.parameters(),     \"lr\": 5e-4},\n    ], weight_decay=1e-2)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max_epochs, eta_min=1e-6)\n\n    best_qwk, best_state, epochs_no_improve = -np.inf, None, 0\n    best_f1 = best_val_loss = best_recall = None\n    best_preds = best_labels = None\n    history = []\n\n    for epoch in range(max_epochs):\n        model.freeze_backbone(epoch < 3)   # freeze backbone for first 3 epochs\n        model.train()\n        running_loss = 0.0\n        for x, y in train_loader:\n            x, y = x.to(DEVICE), y.to(DEVICE)\n            optimizer.zero_grad()\n            logits = model(x)\n            loss = criterion(logits, y)\n            loss.backward()\n            optimizer.step()\n            running_loss += loss.item() * x.size(0)\n        scheduler.step()\n\n        model.eval()\n        val_loss_sum = 0.0\n        all_preds, all_labels = [], []\n        with torch.no_grad():\n            for x, y in val_loader:\n                x = x.to(DEVICE)\n                logits = model(x)\n                val_loss_sum += criterion(logits, y.to(DEVICE)).item() * x.size(0)\n                all_preds.append(logits.argmax(1).cpu().numpy())\n                all_labels.append(y.numpy())\n        all_preds = np.concatenate(all_preds)\n        all_labels = np.concatenate(all_labels)\n        val_loss = val_loss_sum / len(val_idx)\n        val_qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n        val_f1 = f1_score(all_labels, all_preds, average=\"macro\")\n        val_recall = recall_score(all_labels, all_preds, average=None,\n                                  labels=list(range(n_classes)), zero_division=0)\n\n        history.append({\"epoch\": epoch, \"train_loss\": running_loss / len(train_idx),\n                        \"val_loss\": val_loss, \"val_qwk\": val_qwk, \"val_f1\": val_f1})\n\n        if val_qwk > best_qwk:\n            best_qwk, best_f1, best_val_loss = val_qwk, val_f1, val_loss\n            best_recall = val_recall\n            best_state = copy.deepcopy(model.state_dict())\n            best_preds, best_labels = all_preds, all_labels\n            epochs_no_improve = 0\n        else:\n            epochs_no_improve += 1\n            if epochs_no_improve >= patience:\n                break\n\n    suffix = \"\" if h_kan_override is None else f\"_hk{h_kan_override}\"\n    run_id = f\"{tag}_{backbone_name}_{head_type}{suffix}_fold{fold_id}\"\n    ckpt_name = f\"{run_id}.pt\"\n    torch.save(best_state, CKPT_DIR / ckpt_name)\n    pd.DataFrame(history).to_csv(OUT_DIR / f\"history_{run_id}.csv\", index=False)\n    np.savez(OUT_DIR / f\"val_preds_{run_id}.npz\", preds=best_preds, labels=best_labels, val_idx=val_idx)\n\n    del model\n    torch.cuda.empty_cache()\n    return {\"backbone\": backbone_name, \"head\": head_type, \"fold\": fold_id,\n            \"qwk\": best_qwk, \"f1\": best_f1, \"val_loss\": best_val_loss,\n            \"recall\": best_recall.tolist(), \"checkpoint\": str(CKPT_DIR / ckpt_name)}\n\n\ndef run_full_grid(dataset_spec, tag, n_splits=5, backbones=None, patience=10):\n    \"\"\"Full 3-backbone x 2-head x n_splits-fold grid on one dataset.\n    tag=\"aptos\" -> Table 3; tag=\"ddr\" -> Table 4. Folds are shared across\n    backbones and heads (same StratifiedKFold split) so MLP vs KAN is paired.\"\"\"\n    df, label_col, n_classes = dataset_spec[\"df\"], dataset_spec[\"label_col\"], dataset_spec[\"n_classes\"]\n    skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=SEED)\n    fold_indices = list(skf.split(df, df[label_col]))\n\n    results = []\n    for backbone_name in (backbones or list(BACKBONE_SPECS)):\n        for head_type in (\"mlp\", \"kan\"):\n            for fold_id, (train_idx, val_idx) in enumerate(fold_indices):\n                r = train_one_fold(backbone_name, head_type, dataset_spec,\n                                   train_idx, val_idx, fold_id,\n                                   n_classes=n_classes, tag=tag, patience=patience)\n                print(f\"[{tag}] {backbone_name:18s} {head_type:3s} fold{fold_id} \"\n                      f\"-> QWK={r['qwk']:.4f}  F1={r['f1']:.4f}\")\n                results.append(r)\n    return pd.DataFrame(results)\n\n\n# Shared helper (used by Fig. 5, Fig. 6 and the Table 7 faithfulness protocol):\n# extract the global-average-pooled backbone feature matrix f in R^d.\ndef extract_backbone_features(model, loader):\n    feats, labels = [], []\n    model.eval()\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(DEVICE)\n            f = model.extract_features(x).cpu().numpy()\n            feats.append(f)\n            labels.append(y.numpy())\n    return np.concatenate(feats), np.concatenate(labels)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:14:54.594930Z","iopub.execute_input":"2026-07-24T17:14:54.595568Z","iopub.status.idle":"2026-07-24T17:14:54.618057Z","shell.execute_reply.started":"2026-07-24T17:14:54.595539Z","shell.execute_reply":"2026-07-24T17:14:54.617151Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Full training driver (used only when REPRODUCE_FROM_CHECKPOINTS = False).\nif not REPRODUCE_FROM_CHECKPOINTS:\n    print(\"Retraining from scratch. This takes several hours and, because the KAN \"\n          \"head is not bit-reproducible, will not reproduce the fourth decimal of \"\n          \"Table 3. See the caveats at the top of this notebook.\")\n    aptos_results = run_full_grid(aptos, \"aptos\", backbones_to_run(), patience=10)\n    print(aptos_results)\nelse:\n    print(\"Using released checkpoints (default). Skipping training.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:15:01.622277Z","iopub.execute_input":"2026-07-24T17:15:01.622840Z","iopub.status.idle":"2026-07-24T17:15:01.627642Z","shell.execute_reply.started":"2026-07-24T17:15:01.622811Z","shell.execute_reply":"2026-07-24T17:15:01.626978Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Released artifacts\n\nResolves the public checkpoint and fold-prediction datasets, and defines the feature-deletion routine used from Section 16 onward.","metadata":{}},{"cell_type":"code","source":"# Resolve the released artifacts (checkpoints + fold predictions) from the\n# attached public datasets, and mirror the fold predictions into OUT_DIR.\nimport glob\n\ndef _find(pattern, what):\n    hits = sorted(glob.glob(pattern, recursive=True))\n    if not hits:\n        raise FileNotFoundError(\n            f\"Could not locate {what}. Attach the public datasets listed at the \"\n            f\"top of this notebook (dr-kan-mlp-aptos, dr-kan-mlp-results).\"\n        )\n    return hits\n\nAPTOS_CKPT_DIR = str(Path(_find(\"/kaggle/input/**/aptos_MobileNetV3-Large_kan_fold0.pt\",\n                                \"APTOS checkpoints\")[0]).parent)\nAPTOS_PRED_DIR = str(Path(_find(\"/kaggle/input/**/val_preds_aptos_MobileNetV3-Large_kan_fold0.npz\",\n                                \"APTOS fold predictions\")[0]).parent)\nDDR_PRED_FILES = _find(\"/kaggle/input/**/val_preds_ddr_*.npz\", \"DDR fold predictions\")\n\nfor f in glob.glob(APTOS_PRED_DIR + \"/val_preds_aptos_*.npz\"):\n    dst = OUT_DIR / Path(f).name\n    if not dst.exists():\n        dst.write_bytes(Path(f).read_bytes())\n\nprint(f\"APTOS checkpoints : {APTOS_CKPT_DIR}\")\nprint(f\"APTOS predictions : {APTOS_PRED_DIR}\")\nprint(f\"DDR predictions   : {len(DDR_PRED_FILES)} fold files\")\n\n\ndef aptos_checkpoint(head, fold):\n    return f\"{APTOS_CKPT_DIR}/aptos_{SHOWCASE_BACKBONE}_{head}_fold{fold}.pt\"\n\n\ndef load_head_model(head, fold):\n    m = DRModel(SHOWCASE_BACKBONE, head, n_classes=5).to(DEVICE)\n    m.load_state_dict(torch.load(aptos_checkpoint(head, fold), map_location=DEVICE))\n    m.eval()\n    return m\n\n\ndef deletion_curve(head, features, labels, ranking, steps, mean_vector):\n    \"\"\"QWK as top-ranked features are replaced by their training-set mean.\"\"\"\n    order = np.argsort(-ranking)\n    d_in = features.shape[1]\n    dev = next(head.parameters()).device\n    qwks = []\n    for frac in steps:\n        n_remove = int(round(frac * d_in))\n        modified = features.copy()\n        modified[:, order[:n_remove]] = mean_vector[order[:n_remove]]\n        with torch.no_grad():\n            preds = head(torch.tensor(modified, dtype=torch.float32, device=dev)).argmax(1).cpu().numpy()\n        qwks.append(cohen_kappa_score(labels, preds, weights=\"quadratic\"))\n    return qwks\n\n\ndef bootstrap_ci(values, n_resamples=2000, ci=95, seed=SEED):\n    \"\"\"Percentile bootstrap CI of the mean (Efron & Tibshirani).\"\"\"\n    values = np.asarray(values, dtype=float)\n    rng = np.random.default_rng(seed)\n    means = [np.mean(rng.choice(values, size=len(values), replace=True))\n             for _ in range(n_resamples)]\n    lo = np.percentile(means, (100 - ci) / 2)\n    hi = np.percentile(means, 100 - (100 - ci) / 2)\n    return float(lo), float(hi)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:15:05.398179Z","iopub.execute_input":"2026-07-24T17:15:05.398692Z","iopub.status.idle":"2026-07-24T17:16:00.617482Z","shell.execute_reply.started":"2026-07-24T17:15:05.398662Z","shell.execute_reply":"2026-07-24T17:16:00.616738Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Table 3 — APTOS 2019 five-fold cross-validation\n\nRecomputed from the released fold predictions. Only MobileNetV3-Large fold predictions are part of the released artifacts; see the note at the top.","metadata":{}},{"cell_type":"code","source":"# Table 3 -- APTOS 2019 five-fold cross-validation.\n# Recomputed from the released fold predictions (not copied from a stored CSV).\nimport re\n\ndef results_from_predictions(prefix, files):\n    rows = []\n    for f in files:\n        m = re.search(prefix + r\"_(.+?)_(mlp|kan)_fold(\\d+)\\.npz\", f)\n        if not m:\n            continue\n        bk, hd, fo = m.group(1), m.group(2), int(m.group(3))\n        d = np.load(f)\n        rows.append({\"backbone\": bk, \"head\": hd.upper(), \"fold\": fo,\n                     \"qwk\": cohen_kappa_score(d[\"labels\"], d[\"preds\"], weights=\"quadratic\"),\n                     \"macro_f1\": f1_score(d[\"labels\"], d[\"preds\"], average=\"macro\")})\n    return pd.DataFrame(rows)\n\n\ndef summarise(df):\n    out = []\n    for (bk, hd), g in df.groupby([\"backbone\", \"head\"]):\n        lo, hi = bootstrap_ci(g[\"qwk\"].values)\n        out.append({\"Backbone\": bk, \"Head\": hd,\n                    \"QWK_mean\": g[\"qwk\"].mean(), \"QWK_std\": g[\"qwk\"].std(ddof=1),\n                    \"CI_low\": lo, \"CI_high\": hi, \"Macro_F1\": g[\"macro_f1\"].mean()})\n    return pd.DataFrame(out).sort_values([\"Backbone\", \"Head\"]).reset_index(drop=True)\n\n\naptos_fold_results = results_from_predictions(\n    \"val_preds_aptos\", sorted(glob.glob(APTOS_PRED_DIR + \"/val_preds_aptos_*.npz\")))\ntable3 = summarise(aptos_fold_results)\nprint(\"Table 3 -- APTOS 2019, five-fold CV\")\nprint(table3.round(4).to_string(index=False))\ntable3.to_csv(OUT_DIR / \"table3_aptos_results.csv\", index=False)\n\nmissing = {\"EfficientNet-B3\", \"ResNet-50\"} - set(table3[\"Backbone\"])\nif missing:\n    print(f\"\\nNote: fold predictions for {sorted(missing)} on APTOS are not part of the \"\n          f\"released artifacts, so those rows of the manuscript's Table 3 are not \"\n          f\"recomputed here. Set REPRODUCE_FROM_CHECKPOINTS = False to retrain them \"\n          f\"(several hours; the KAN head is not bit-reproducible).\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:16:06.488315Z","iopub.execute_input":"2026-07-24T17:16:06.488774Z","iopub.status.idle":"2026-07-24T17:16:06.644733Z","shell.execute_reply.started":"2026-07-24T17:16:06.488738Z","shell.execute_reply":"2026-07-24T17:16:06.643932Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Table 4 — DDR five-fold cross-validation\n\nDDR was trained on a 50% class-stratified subsample (n = 6,260) with early-stopping patience 5, as reported in the manuscript.","metadata":{}},{"cell_type":"code","source":"# Table 4 -- DDR five-fold cross-validation (50% class-stratified subsample).\n# Recomputed from the released fold predictions.\nddr_fold_results = results_from_predictions(\"val_preds_ddr\", DDR_PRED_FILES)\nn_val = sum(len(np.load(f)[\"labels\"]) for f in DDR_PRED_FILES\n            if f\"{SHOWCASE_BACKBONE}_kan\" in f)\nprint(f\"DDR subsample size (5 folds, one backbone/head): n = {n_val}\")\n\ntable4 = summarise(ddr_fold_results)\nprint(\"\\nTable 4 -- DDR, five-fold CV\")\nprint(table4.round(4).to_string(index=False))\ntable4.to_csv(OUT_DIR / \"table4_ddr_results.csv\", index=False)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:16:11.614040Z","iopub.execute_input":"2026-07-24T17:16:11.614550Z","iopub.status.idle":"2026-07-24T17:16:12.288139Z","shell.execute_reply.started":"2026-07-24T17:16:11.614523Z","shell.execute_reply":"2026-07-24T17:16:12.287199Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Equivalence testing (TOST) and the paired Wilcoxon check\n\nThe claim is equivalence, not difference. A failure to reject equality is not evidence of equality, so the TOST procedure carries the claim and Wilcoxon is reported descriptively.","metadata":{}},{"cell_type":"code","source":"# Statistical comparison behind Tables 3-4.\n#\n# KAN and MLP share folds, so the comparison is paired. The claim under test is\n# EQUIVALENCE, and a null-hypothesis test cannot support it: failing to reject\n# equality is not evidence of equality. We therefore report a two one-sided tests\n# (TOST) procedure at a margin of TOST_MARGIN QWK, and give the Wilcoxon\n# signed-rank test only as a descriptive check.\n#\n# Five-fold training sets overlap, so fold estimates are not independent and any\n# test built on them is anti-conservative (Dietterich 1998; Nadeau & Bengio 2003).\nfrom scipy import stats\n\n\ndef tost(d, margin):\n    \"\"\"Two one-sided tests for equivalence of paired differences d within +/- margin.\"\"\"\n    n = len(d)\n    se = d.std(ddof=1) / np.sqrt(n)\n    p_lo = stats.t.sf((d.mean() + margin) / se, n - 1)   # H0: mu <= -margin\n    p_hi = stats.t.cdf((d.mean() - margin) / se, n - 1)  # H0: mu >= +margin\n    return max(p_lo, p_hi), se\n\n\nprint(\"Paired analysis of KAN - MLP per-fold QWK\\n\")\nfor name, df in ((\"APTOS\", aptos_fold_results), (\"DDR\", ddr_fold_results)):\n    for bk, g in df.groupby(\"backbone\"):\n        kan = g[g[\"head\"] == \"KAN\"].sort_values(\"fold\")[\"qwk\"].values\n        mlp = g[g[\"head\"] == \"MLP\"].sort_values(\"fold\")[\"qwk\"].values\n        if len(kan) != len(mlp) or len(kan) < 3:\n            continue\n        d = kan - mlp\n        p_tost, se = tost(d, TOST_MARGIN)\n        try:\n            p_w = stats.wilcoxon(kan, mlp).pvalue\n        except ValueError:\n            p_w = float(\"nan\")\n        lo = d.mean() - stats.t.ppf(0.95, len(d) - 1) * se\n        hi = d.mean() + stats.t.ppf(0.95, len(d) - 1) * se\n        verdict = \"EQUIVALENT\" if p_tost < 0.05 else \"not established\"\n        print(f\"{name:6s} {bk:20s} dQWK {d.mean():+.4f}  90% CI [{lo:+.4f}, {hi:+.4f}]  \"\n              f\"TOST p={p_tost:.3f} ({verdict})   Wilcoxon p={p_w:.3f}\")\n\nprint(f\"\\nEquivalence margin: {TOST_MARGIN} QWK. With five folds the procedure has little\")\nprint(\"power at a tighter margin; the manuscript claims equivalence only at this one.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:16:15.457944Z","iopub.execute_input":"2026-07-24T17:16:15.459167Z","iopub.status.idle":"2026-07-24T17:16:15.489151Z","shell.execute_reply.started":"2026-07-24T17:16:15.459133Z","shell.execute_reply":"2026-07-24T17:16:15.488443Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Spline visualisation utilities\n\nFeature-to-class-logit curves are partial dependences through the composed two-layer head — the same procedure, at the same cost, for either head. Defined here because the audit utilities of Section 11 build on them.","metadata":{}},{"cell_type":"code","source":"# ---------------------------------------------------------------------\n# FAITHFUL feature -> class-logit curves for a full 2-layer head.\n#\n# Why this replaces `KANLinear.edge_function(i, c)` for interpretation:\n# in the head KANLinear(d -> h_kan) -> KANLinear(h_kan -> 5), the FIRST layer's\n# edges map an input feature to a HIDDEN unit, NOT to a class. Plotting\n# edge_function(feature, c) and labelling column c as DR-class c therefore\n# mislabels hidden-unit activations as class contributions. The correct,\n# architecture-honest object is a partial-dependence curve through the WHOLE\n# head: vary one backbone feature across its empirical range, hold the others\n# at their dataset mean, and read the resulting class-c logit. This works\n# identically for MLP and KAN heads (so the two can be compared directly) and\n# does not alter any accuracy number in Tables 1-5.\n# ---------------------------------------------------------------------\ndef head_class_curves(head, feat_idx, feature_means, x_lo, x_hi, n_points=200):\n    \"\"\"Return (xs, logits[n_points, C]) as feature `feat_idx` sweeps [x_lo, x_hi].\"\"\"\n    dev = next(head.parameters()).device\n    xs = np.linspace(float(x_lo), float(x_hi), n_points).astype(np.float32)\n    X = np.tile(feature_means.astype(np.float32), (n_points, 1))\n    X[:, feat_idx] = xs\n    with torch.no_grad():\n        logits = head(torch.tensor(X, device=dev)).cpu().numpy()\n    return xs, logits\n\n\ndef feature_class_ranges(head, feature_means, x_lo, x_hi, n_classes=5, n_points=60):\n    \"\"\"Per-feature, per-class discriminative range of the true feature->class\n    curve: r[f, c] = max_x phi_fc(x) - min_x phi_fc(x). One forward per feature.\"\"\"\n    d = feature_means.shape[0]\n    ranges = np.zeros((d, n_classes), dtype=np.float32)\n    for f in range(d):\n        _, logits = head_class_curves(head, f, feature_means, x_lo[f], x_hi[f], n_points)\n        ranges[f] = logits.max(0) - logits.min(0)\n    return ranges\n\n\n\n\ndef select_representative_features(head, feat_means, feat_lo, feat_hi, n_classes=5, top_k=8):\n    ranges = feature_class_ranges(head, feat_means, feat_lo, feat_hi, n_classes, n_points=60)\n    mean_mag = ranges.mean(axis=1)\n    top8 = np.argsort(-mean_mag)[:top_k].tolist()\n\n    curves = {}   # feat -> (xs, logits[n_points, C])\n    for f in top8:\n        curves[f] = head_class_curves(head, f, feat_means, feat_lo[f], feat_hi[f], n_points=200)\n\n    def n_turns(y):\n        d = np.diff(y)\n        s = np.sign(d); s = s[s != 0]\n        return int(np.sum(np.diff(s) != 0)) if len(s) else 0\n\n    non_monotone = max(top8, key=lambda f: n_turns(curves[f][1][:, 0]))\n    strongest_dec_nodr = min(top8, key=lambda f: curves[f][1][-1, 0] - curves[f][1][0, 0])\n    strongest_inc_pdr = max(top8, key=lambda f: curves[f][1][-1, 4] - curves[f][1][0, 4])\n    largest_mag = top8[int(np.argmax(mean_mag[top8]))]\n\n    selected = list(dict.fromkeys([non_monotone, strongest_dec_nodr, strongest_inc_pdr, largest_mag]))\n    return selected, top8, curves\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:16:23.415149Z","iopub.execute_input":"2026-07-24T17:16:23.415609Z","iopub.status.idle":"2026-07-24T17:16:23.427873Z","shell.execute_reply.started":"2026-07-24T17:16:23.415576Z","shell.execute_reply":"2026-07-24T17:16:23.426944Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Audit utilities\n\nShared machinery for the controls: head constructors under three dropout regimes, a head trainer on cached features, the three attribution rankings, and the three removal operators.","metadata":{}},{"cell_type":"code","source":"# Shared utilities for the audit controls (Sections 12, 15, 16, 17, 18).\nimport copy, time\nfrom math import comb\nfrom scipy import stats\nfrom sklearn.model_selection import StratifiedKFold\n\ntry:\n    import shap\nexcept ImportError:\n    import subprocess, sys\n    subprocess.run([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \"shap\"], check=True)\n    import shap\n\n\nclass MLPHeadInputDrop(nn.Module):\n    \"\"\"MLP head with dropout on the input vector, mirroring KANHead's placement.\n\n    The reference designs are not symmetric: MLPHead drops hidden activations,\n    KANHead drops the raw backbone features. Section 12 quantifies what that costs.\n    \"\"\"\n    def __init__(self, d_in, h, n_classes=5, p_drop=0.3):\n        super().__init__()\n        self.drop = nn.Dropout(p_drop)\n        self.fc1 = nn.Linear(d_in, h)\n        self.act = nn.ReLU()\n        self.fc2 = nn.Linear(h, n_classes)\n\n    def forward(self, x):\n        return self.fc2(self.act(self.fc1(self.drop(x))))\n\n\nclass KANHeadHiddenDrop(nn.Module):\n    \"\"\"KAN head with dropout on the hidden activations, mirroring MLPHead's placement.\"\"\"\n    def __init__(self, d_in, h_kan, n_classes=5, p_drop=0.3, grid_size=5, spline_order=3):\n        super().__init__()\n        self.kan1 = KANLinear(d_in, h_kan, grid_size=grid_size, spline_order=spline_order)\n        self.drop = nn.Dropout(p_drop)\n        self.kan2 = KANLinear(h_kan, n_classes, grid_size=grid_size, spline_order=spline_order)\n    def forward(self, x):\n        return self.kan2(self.drop(self.kan1(x)))\n\n\ndef make_head(backbone, head_type, regime=\"as-published\"):\n    spec = BACKBONE_SPECS[backbone]\n    d = spec[\"d\"]\n    if regime == \"as-published\":\n        return MLPHead(d, spec[\"h_mlp\"], 5, p_drop=0.3) if head_type == \"mlp\" \\\n            else KANHead(d, spec[\"h_kan\"], 5, p_drop=0.3)\n    if regime == \"no-dropout\":\n        return MLPHead(d, spec[\"h_mlp\"], 5, p_drop=0.0) if head_type == \"mlp\" \\\n            else KANHead(d, spec[\"h_kan\"], 5, p_drop=0.0)\n    if regime == \"both-input\":\n        return MLPHeadInputDrop(d, spec[\"h_mlp\"], 5, p_drop=0.3) if head_type == \"mlp\" \\\n            else KANHead(d, spec[\"h_kan\"], 5, p_drop=0.3)\n    if regime == \"both-hidden\":\n        return MLPHead(d, spec[\"h_mlp\"], 5, p_drop=0.3) if head_type == \"mlp\" \\\n            else KANHeadHiddenDrop(d, spec[\"h_kan\"], 5, p_drop=0.3)\n    raise ValueError(regime)\n\n\ndef fit_head(proto, Xtr, ytr, lr=5e-3, epochs=60, seed=SEED):\n    \"\"\"Train a fresh copy of `proto` on cached features. Seconds on a T4.\"\"\"\n    torch.manual_seed(seed)\n    head = copy.deepcopy(proto).to(DEVICE)\n    Xt = torch.tensor(Xtr, dtype=torch.float32, device=DEVICE)\n    yt = torch.tensor(ytr, dtype=torch.long, device=DEVICE)\n    cnt = np.bincount(ytr, minlength=5).astype(float); cnt[cnt == 0] = 1\n    w = torch.tensor(len(ytr) / (5 * cnt), dtype=torch.float32, device=DEVICE)\n    opt = torch.optim.AdamW(head.parameters(), lr=lr, weight_decay=1e-2)\n    sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs)\n    n, bs = len(yt), 256\n    head.train()\n    for _ in range(epochs):\n        perm = torch.randperm(n, device=DEVICE)\n        for i in range(0, n, bs):\n            b = perm[i:i + bs]\n            opt.zero_grad()\n            F.cross_entropy(head(Xt[b]), yt[b], weight=w).backward()\n            opt.step()\n        sch.step()\n    head.eval()\n    return head\n\n\ndef qwk_of(head, X, y):\n    dev = next(head.parameters()).device\n    with torch.no_grad():\n        p = head(torch.tensor(X, dtype=torch.float32, device=dev)).argmax(1).cpu().numpy()\n    return cohen_kappa_score(y, p, weights=\"quadratic\")\n\n\ndef fold_features(head_type, fold):\n    \"\"\"Cached (train_feats, train_y, val_feats, val_y, head) from that fold's own backbone.\"\"\"\n    m = load_head_model(head_type, fold)\n    npz = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold{fold}.npz\")\n    tr = np.setdiff1d(np.arange(len(aptos[\"df\"])), npz[\"val_idx\"])\n    vl = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                     aptos[\"ext\"], npz[\"val_idx\"], eval_transform, 16, shuffle=False)\n    tl = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                     aptos[\"ext\"], tr, eval_transform, 16, shuffle=False)\n    vf, vy = extract_backbone_features(m, vl)\n    tf, ty = extract_backbone_features(m, tl)\n    head = copy.deepcopy(m.head)\n    del m; torch.cuda.empty_cache()\n    return tf, ty, vf, vy, head\n\n\n# ---- attribution rankings -------------------------------------------------\ndef rank_pd(head, vf, lo, hi):\n    \"\"\"Partial-dependence range. Model-agnostic: identical procedure for both heads.\"\"\"\n    return feature_class_ranges(head, vf.mean(0), lo, hi, 5, n_points=60).mean(1)\n\n\ndef rank_ig(head, ev, baseline, n_steps=32):\n    \"\"\"Integrated Gradients from the training-set mean, the vector deletion imputes.\"\"\"\n    dev = next(head.parameters()).device\n    X = torch.tensor(ev, dtype=torch.float32, device=dev)\n    b = torch.tensor(baseline, dtype=torch.float32, device=dev).unsqueeze(0)\n    attr = torch.zeros(X.shape[1], device=dev)\n    for c in range(5):\n        g = torch.zeros_like(X)\n        for a in torch.linspace(0, 1, n_steps, device=dev):\n            xi = (b + a * (X - b)).detach().requires_grad_(True)\n            out = F.softmax(head(xi), 1)[:, c].sum()\n            gr, = torch.autograd.grad(out, xi)\n            g += gr.detach()\n        attr += ((X - b) * g / n_steps).abs().mean(0)\n    return (attr / 5).detach().cpu().numpy()\n\n\ndef rank_shap(head, background, ev, nsamples=100):\n    np.random.seed(SEED)\n    dev = next(head.parameters()).device\n    f = lambda x: F.softmax(head(torch.tensor(x, dtype=torch.float32, device=dev)), 1).detach().cpu().numpy()\n    d = background.shape[1]\n    sv = shap.KernelExplainer(f, background).shap_values(ev, nsamples=nsamples)\n    arr = np.stack([np.abs(a) for a in sv]) if isinstance(sv, list) else np.abs(np.asarray(sv))\n    ax = tuple(a for a in range(arr.ndim) if arr.shape[a] != d)\n    return np.asarray(arr.mean(axis=ax) if ax else arr).reshape(-1)\n\n\n# ---- deletion machinery ---------------------------------------------------\nDELETION_STEPS = [0.0, 0.05, 0.10, 0.20, 0.30, 0.40, 0.50]\n\n\ndef ablate(X, idx, operator, train_ref, rng):\n    \"\"\"Remove columns `idx` under one of three operators (Section 3.5).\"\"\"\n    Z = X.copy()\n    if operator == \"mean\":\n        Z[:, idx] = train_ref.mean(0)[idx]\n    elif operator == \"marginal\":          # keep the marginal, destroy the association\n        for j in idx:\n            Z[:, j] = rng.choice(train_ref[:, j], size=Z.shape[0], replace=True)\n    elif operator == \"zero\":\n        Z[:, idx] = 0.0\n    else:\n        raise ValueError(operator)\n    return Z\n\n\ndef aor(curve, random_curve, steps=DELETION_STEPS):\n    \"\"\"Area over the random control, Eq. (3) of the manuscript.\"\"\"\n    return float(np.trapz(np.asarray(random_curve) - np.asarray(curve), steps))\n\n\nprint(\"audit utilities ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:16:28.090241Z","iopub.execute_input":"2026-07-24T17:16:28.091197Z","iopub.status.idle":"2026-07-24T17:16:32.755415Z","shell.execute_reply.started":"2026-07-24T17:16:28.091165Z","shell.execute_reply":"2026-07-24T17:16:32.754776Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Table 5 — frozen-backbone probe *(heavy, ~35 min)*\n\nBoth heads are trained on an identical, frozen ImageNet representation, each at its own best learning rate, under three dropout regimes. This is the only regime in which the head is, without qualification, the sole variable.","metadata":{}},{"cell_type":"code","source":"# Table 5 -- frozen-backbone probe.\n#\n# Under joint fine-tuning each head reshapes its own representation, so the\n# equivalence of Table 3 could be a compensation effect. Here the ImageNet\n# backbone is frozen, its features are cached once, and BOTH heads are trained on\n# the identical feature matrix. Each head gets its own best learning rate, which\n# favours the two symmetrically.\n#\n# Three dropout regimes are compared, because the reference designs are not\n# symmetric: MLPHead drops hidden activations (h_mlp <= 78 units), KANHead drops\n# the raw input (d >= 1280 units).\nif not RUN_HEAVY:\n    print(\"RUN_HEAVY = False -> skipping the frozen-backbone probe (Table 5).\")\nelse:\n    BACKBONES = list(BACKBONE_SPECS)\n    LRS = [1e-3, 5e-3, 1e-2]\n    #REGIMES = [\"as-published\", \"no-dropout\", \"both-input\"]\n    REGIMES = [\"as-published\", \"no-dropout\", \"both-input\", \"both-hidden\"]\n\n    print(\"Head parameter counts stay matched under every regime:\")\n    for bk in BACKBONES:\n        parts = []\n        for r in REGIMES:\n            pm = sum(p.numel() for p in make_head(bk, \"mlp\", r).parameters())\n            pk = sum(p.numel() for p in make_head(bk, \"kan\", r).parameters())\n            parts.append(f\"{r}: MLP {pm:,} KAN {pk:,}\")\n        print(f\"  {bk:20s} \" + \" | \".join(parts))\n\n    rows = []\n    for bk in BACKBONES:\n        base = DRModel(bk, \"mlp\", n_classes=5).to(DEVICE).eval()   # ImageNet weights, frozen\n        ld = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                         aptos[\"ext\"], np.arange(len(aptos[\"df\"])), eval_transform, 16, shuffle=False)\n        Fx, Fy = extract_backbone_features(base, ld)\n        del base; torch.cuda.empty_cache()\n        print(f\"\\n{bk}: frozen features {Fx.shape}\")\n\n        protos = {(t, r): make_head(bk, t, r) for t in (\"mlp\", \"kan\") for r in REGIMES}\n        skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=SEED)\n        for fold, (tr, va) in enumerate(skf.split(Fx, Fy)):\n            for (t, r), proto in protos.items():\n                for lr in LRS:\n                    h = fit_head(proto, Fx[tr], Fy[tr], lr=lr)\n                    rows.append({\"backbone\": bk, \"regime\": r, \"head\": t.upper(),\n                                 \"lr\": lr, \"fold\": fold, \"QWK\": qwk_of(h, Fx[va], Fy[va])})\n                    del h; torch.cuda.empty_cache()\n            print(f\"  fold {fold} done\")\n\n    probe = pd.DataFrame(rows)\n    probe.to_csv(OUT_DIR / \"table5_frozen_probe.csv\", index=False)\n\n    print(\"\\nTable 5 -- each head at its own best learning rate, per dropout regime\")\n    print(\"=\" * 76)\n    for bk in BACKBONES:\n        for r in REGIMES:\n            best = {}\n            for t in (\"MLP\", \"KAN\"):\n                sub = probe[(probe[\"backbone\"] == bk) & (probe[\"regime\"] == r) & (probe[\"head\"] == t)]\n                blr = sub.groupby(\"lr\")[\"QWK\"].mean().idxmax()\n                best[t] = (blr, sub[sub[\"lr\"] == blr].sort_values(\"fold\")[\"QWK\"].values)\n            (m_lr, m), (k_lr, k) = best[\"MLP\"], best[\"KAN\"]\n            d = k - m\n            se = d.std(ddof=1) / np.sqrt(len(d))\n            p = 2 * stats.t.sf(abs(d.mean() / se), len(d) - 1)\n            p_tost = max(stats.t.sf((d.mean() + TOST_MARGIN) / se, len(d) - 1),\n                         stats.t.cdf((d.mean() - TOST_MARGIN) / se, len(d) - 1))\n            flag = \"EQUIVALENT\" if p_tost < 0.05 else \"not equivalent\"\n            print(f\"  {bk:20s} {r:14s} MLP {m.mean():.4f}  KAN {k.mean():.4f}  \"\n                  f\"d={d.mean():+.4f} (p={p:.3f})  TOST p={p_tost:.3f} {flag}\")\n    print(\"\\nEquivalence holds on every backbone when dropout is placed identically,\")\n    print(\"and fails whenever it is not. The asymmetry alone costs KAN ~0.03 QWK.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13. Table 6 — Messidor-2, five folds, calibration and prevalence","metadata":{}},{"cell_type":"code","source":"# Table 6 -- external validation on Messidor-2, all five folds.\n#\n# The manuscript's earlier version used the single best fold per head; the AUC\n# advantage it reported did not survive evaluation of all five. Two candidate\n# explanations for the operating-point failure are tested here: temperature\n# scaling (calibration) and prevalence shift.\nfrom sklearn.metrics import roc_auc_score, roc_curve, confusion_matrix\n\nmessidor = load_messidor2(); messidor = precompute_image_cache(messidor, \"messidor\")\n\n\ndef logits_of(model, ds, label_col, idx):\n    ld = make_loader(ds[\"df\"], ds[\"image_dir\"], ds[\"image_col\"], label_col, ds[\"ext\"],\n                     idx, eval_transform, 16, shuffle=False)\n    L, Y = [], []\n    with torch.no_grad():\n        for x, y in ld:\n            L.append(model(x.to(DEVICE)).cpu()); Y.append(y)\n    return torch.cat(L), torch.cat(Y)\n\n\ndef referral(logits, T=1.0):\n    return F.softmax(logits / T, dim=1)[:, 2:].sum(1).numpy()   # P(grade >= 2)\n\n\ndef fit_temperature(logits, labels):\n    logT = torch.zeros(1, requires_grad=True)\n    opt = torch.optim.LBFGS([logT], lr=0.1, max_iter=100)\n    def closure():\n        opt.zero_grad()\n        loss = F.cross_entropy(logits / torch.exp(logT), labels)\n        loss.backward(); return loss\n    opt.step(closure)\n    return float(torch.exp(logT).item())\n\n\ndef youden(y, s):\n    fpr, tpr, thr = roc_curve(y, s); return thr[np.argmax(tpr - fpr)]\n\n\ndef sens_spec(y, s, t):\n    tn, fp, fn, tp = confusion_matrix(y, (s >= t).astype(int), labels=[0, 1]).ravel()\n    return tp / (tp + fn), tn / (tn + fp)\n\n\nrows, temps = [], {}\nfor fold in range(5):\n    for tag in (\"mlp\", \"kan\"):\n        m = load_head_model(tag, fold)\n        Lm, Ym = logits_of(m, messidor, \"referable\", messidor[\"df\"].index.values)\n        va = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold{fold}.npz\")[\"val_idx\"]\n        La, Ya = logits_of(m, aptos, aptos[\"label_col\"], va)\n        del m; torch.cuda.empty_cache()\n\n        ya = (Ya.numpy() >= 2).astype(int); ym = Ym.numpy()\n        sa, sm = referral(La), referral(Lm)\n        thr = youden(ya, sa)                                    # fixed on APTOS\n        se_n, sp_n = sens_spec(ym, sm, thr)\n        q = np.quantile(sm, 1 - ym.mean())                      # prevalence-matched\n        se_p, sp_p = sens_spec(ym, sm, q)\n\n        T = fit_temperature(La, Ya)                             # calibration on APTOS val\n        temps.setdefault(tag, []).append(T)\n        sa_t, sm_t = referral(La, T), referral(Lm, T)\n        thr_t = youden(ya, sa_t)\n        se_t, sp_t = sens_spec(ym, sm_t, thr_t)\n\n        rows.append({\"fold\": fold, \"head\": tag.upper(), \"AUC\": roc_auc_score(ym, sm),\n                     \"pi_src\": ya.mean(), \"pi_tgt\": ym.mean(),\n                     \"Sens_transfer\": se_n, \"Spec_transfer\": sp_n,\n                     \"Sens_prevmatch\": se_p, \"Spec_prevmatch\": sp_p,\n                     \"Sens_temp\": se_t, \"Spec_temp\": sp_t})\n    print(f\"  fold {fold} done\")\n\nmess = pd.DataFrame(rows)\nmess.to_csv(OUT_DIR / \"table6_messidor.csv\", index=False)\n\nprint(f\"\\nReferable prevalence: APTOS validation {mess['pi_src'].mean():.3f} -> \"\n      f\"Messidor-2 {mess['pi_tgt'].mean():.3f}\")\nprint(\"\\nTable 6 -- five-fold mean +/- sd\")\ncols = [\"AUC\", \"Sens_transfer\", \"Spec_transfer\", \"Sens_prevmatch\", \"Spec_prevmatch\"]\nprint(mess.groupby(\"head\")[cols].agg([\"mean\", \"std\"]).round(3).to_string())\n\nk = mess[mess[\"head\"] == \"KAN\"].sort_values(\"fold\")[\"AUC\"].values\nm_ = mess[mess[\"head\"] == \"MLP\"].sort_values(\"fold\")[\"AUC\"].values\nt, p = stats.ttest_rel(k, m_)\nprint(f\"\\npaired per-fold dAUC (KAN - MLP): {(k - m_).mean():+.4f} +/- {(k - m_).std(ddof=1):.4f}, \"\n      f\"t({len(k)-1})={t:.2f}, p={p:.3f}\")\n\nprint(\"\\nCalibration check (temperature fitted on the APTOS validation fold):\")\nprint(mess.groupby(\"head\")[[\"Sens_temp\", \"Spec_temp\"]].agg([\"mean\", \"std\"]).round(3).to_string())\nfor tag, ts in temps.items():\n    print(f\"  mean temperature {tag.upper()}: {np.mean(ts):.3f}\")\nprint(\"\\nTemperature scaling does not restore sensitivity: the failure is not global\")\nprint(\"over-confidence. A threshold placed at the target prevalence stabilises both\")\nprint(\"heads, which points to prior shift. Note that rescaling posteriors by the prior\")\nprint(\"odds cannot by itself change decisions: it is a monotone map, so if the threshold\")\nprint(\"is transported through it the confusion matrix is unchanged.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 14. Figure 4 — feature-to-class-logit curves","metadata":{}},{"cell_type":"code","source":"# Figure 4 -- feature-to-class-logit curves of the deployed KAN head.\n# Curves are partial dependences through the composed two-layer head: one\n# feature sweeps its empirical range, all others stay at their validation mean.\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nplt.rcParams.update({\n    \"font.family\": \"serif\", \"font.size\": 9,\n    \"axes.labelsize\": 9, \"axes.titlesize\": 9.5,\n    \"xtick.labelsize\": 8, \"ytick.labelsize\": 8,\n    \"axes.linewidth\": 0.8, \"pdf.fonttype\": 42,\n})\n\nFOLD = 1   # checkpoint used for Figure 4 in the manuscript\n\nkan_fig = load_head_model(\"kan\", FOLD)\nval_idx = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold{FOLD}.npz\")[\"val_idx\"]\nloader = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                     aptos[\"ext\"], val_idx, eval_transform, batch_size=16, shuffle=False)\nvf, _ = extract_backbone_features(kan_fig, loader)\nf_mean = vf.mean(0)\nf_lo, f_hi = np.percentile(vf, 1, 0), np.percentile(vf, 99, 0)\n\nselected, top8, curves = select_representative_features(kan_fig.head, f_mean, f_lo, f_hi)\nfeats = selected[:2]                       # the complementary benign / referable pair\nprint(\"top-8 by mean PD range:\", top8)\nprint(\"features shown in Figure 4:\", feats)\n\nfig, axes = plt.subplots(len(feats), 5, figsize=(11, 2.5 * len(feats)), sharex=\"row\")\nfor r, f in enumerate(feats):\n    xs, logits = curves[f]\n    for c in range(5):\n        ax = axes[r, c]\n        ax.plot(xs, logits[:, c], lw=1.6, color=\"#1f4e79\")\n        ax.axhline(0, color=\"0.6\", lw=0.7, ls=\"--\")\n        ax.grid(alpha=0.22, lw=0.5)\n        ax.spines[[\"top\", \"right\"]].set_visible(False)\n        if r == 0:\n            ax.set_title(class_names_5[c], pad=5)\n        if c == 0:\n            ax.set_ylabel(f\"Feature {f}\\nclass logit\")\n        ax.set_xlabel(\"Feature activation\")\n\nplt.tight_layout()\nPath(\"figs\").mkdir(exist_ok=True)\nplt.savefig(\"figs/figure_4.pdf\", dpi=300)\nplt.savefig(OUT_DIR / \"figure_4.pdf\", dpi=300)\nplt.show()\n\ndel kan_fig\ntorch.cuda.empty_cache()\nprint(\"saved figs/figure_4.pdf\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 15. Figure 5 and the sign-consistency analysis","metadata":{}},{"cell_type":"code","source":"# Figure 5 -- consistency of the deployed KAN head across the five APTOS folds.\n# For the top-4 features (by mean partial-dependence range), the feature-to-\n# class-logit curve of each fold is compared with every other fold.\nkan_models = [load_head_model(\"kan\", fold) for fold in range(5)]\n\nref_val = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold0.npz\")[\"val_idx\"]\nref_loader = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                         aptos[\"ext\"], ref_val, eval_transform, batch_size=16, shuffle=False)\nref_feats, _ = extract_backbone_features(kan_models[0], ref_loader)\nfeat_means = ref_feats.mean(0)\nfeat_lo, feat_hi = np.percentile(ref_feats, 1, 0), np.percentile(ref_feats, 99, 0)\n\nmean_range = np.zeros(feat_means.shape[0])\nfor m in kan_models:\n    mean_range += feature_class_ranges(m.head, feat_means, feat_lo, feat_hi, 5, n_points=40).mean(1) / len(kan_models)\ntop_features = np.argsort(-mean_range)[:4]\nprint(\"Top-4 features by mean PD range:\", top_features.tolist())\n\nrows = []\nfor fi in top_features:\n    for c in range(5):\n        curves = [head_class_curves(m.head, int(fi), feat_means, feat_lo[fi], feat_hi[fi], 200)[1][:, c]\n                  for m in kan_models]\n        r = [stats.pearsonr(curves[i], curves[j])[0] for i in range(5) for j in range(i + 1, 5)]\n        rows.append({\"feature\": int(fi), \"class\": class_names_5[c], \"mean_pairwise_r\": float(np.nanmean(r))})\n\nconsistency = pd.DataFrame(rows)\nn_high = int((consistency[\"mean_pairwise_r\"] >= 0.95).sum())\nn_neg = int((consistency[\"mean_pairwise_r\"] < 0).sum())\nprint(f\"\\nr >= 0.95: {n_high}/20   |   r < 0: {n_neg}/20   |   mean r: {consistency['mean_pairwise_r'].mean():.3f}\")\nprint(consistency.round(3).to_string(index=False))\nconsistency.to_csv(OUT_DIR / \"figure5_consistency.csv\", index=False)\n\nlabels = [f\"F{r.feature}x{r['class']}\" for _, r in consistency.iterrows()]\nvals = consistency[\"mean_pairwise_r\"].values\nfig, ax = plt.subplots(figsize=(10, 5))\nax.bar(range(len(vals)), vals, color=[\"crimson\" if v < 0.85 else \"steelblue\" for v in vals])\nax.set_xticks(range(len(vals))); ax.set_xticklabels(labels, rotation=90, fontsize=7)\nax.axhline(0.85, color=\"black\", ls=\"--\", lw=0.8); ax.axhline(0.0, color=\"grey\", lw=0.6)\nax.set_ylabel(\"Mean pairwise Pearson r (10 fold pairs)\"); ax.set_ylim(-0.35, 1.05)\nax.set_title(f\"Deployed head consistency (h$_{{kan}}$=8, APTOS): {n_high}/20 combinations r$\\\\geq$0.95\", fontsize=10)\nplt.tight_layout()\nPath(\"figs\").mkdir(exist_ok=True)\nplt.savefig(\"figs/figure_5.pdf\", dpi=300); plt.savefig(OUT_DIR / \"figure_5.pdf\", dpi=300)\nplt.show()\n\nfor m in kan_models: del m\ntorch.cuda.empty_cache()\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sign-consistency analysis behind Section 5.2.\n#\n# With five folds there are ten pairs. If k folds learn a curve with an inverted\n# sign and the rest agree, the mean pairwise r is\n#     [C(k,2) + C(5-k,2) - k(5-k)] / 10\n# which is +0.2 for k=1 and -0.2 for k=2. Observed unstable values sit at +/-0.2.\n# The question is whether this is sign inversion or noise on flat curves.\nfor k in range(3):\n    print(f\"  reference: {k} inverted fold(s) -> mean pairwise r = \"\n          f\"{(comb(k,2)+comb(5-k,2)-k*(5-k))/10:+.2f}\")\n\nTOP_N = 32\nmodels = [load_head_model(\"kan\", f) for f in range(5)]\nref = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold0.npz\")[\"val_idx\"]\nld = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                 aptos[\"ext\"], ref, eval_transform, 16, shuffle=False)\nrf, _ = extract_backbone_features(models[0], ld)\nfmean, flo, fhi = rf.mean(0), np.percentile(rf, 1, 0), np.percentile(rf, 99, 0)\n\nmr = np.zeros(fmean.shape[0])\nfor m in models:\n    mr += feature_class_ranges(m.head, fmean, flo, fhi, 5, n_points=40).mean(1) / 5\ntop = np.argsort(-mr)[:TOP_N]\nprint(f\"\\nanalysing the top {TOP_N} features x 5 classes = {TOP_N*5} combinations\")\n\nout = []\nfor fi in top:\n    for c in range(5):\n        cv = [head_class_curves(m.head, int(fi), fmean, flo[fi], fhi[fi], 200)[1][:, c] for m in models]\n        amp = float(np.mean([x.max() - x.min() for x in cv]))\n        pw = np.array([stats.pearsonr(cv[i], cv[j])[0] for i in range(5) for j in range(i + 1, 5)])\n        out.append({\"feature\": int(fi), \"class\": class_names_5[c], \"amplitude\": amp,\n                    \"mean_r\": pw.mean(),\n                    \"pairs_strong\": int((np.abs(pw) > 0.9).sum()),\n                    \"pairs_weak\": int((np.abs(pw) <= 0.9).sum())})\nsc = pd.DataFrame(out)\nsc[\"unstable\"] = sc[\"mean_r\"] < 0.85\nsc.to_csv(OUT_DIR / \"sign_consistency.csv\", index=False)\n\nprint(\"\\nunstable fraction per class (mean_r < 0.85):\")\nprint(sc.groupby(\"class\")[\"unstable\"].mean().round(3).to_string())\n\nu = sc[sc[\"unstable\"]]\nprint(f\"\\namong unstable rows: {u['pairs_strong'].sum()} of {10*len(u)} fold pairs have |r| > 0.9\")\nprint(\"  -> the correlations pile at +/-1, so the instability is SIGN INVERSION, not noise\")\n\nrho, p = stats.spearmanr(sc[\"amplitude\"], sc[\"mean_r\"])\nprint(f\"\\nSpearman(amplitude, mean_r) = {rho:.3f} (p = {p:.4f})\")\nprint(\"  -> the curves whose direction flips are the curves whose amplitude is small.\")\nprint(\"     Reporting Pearson r without amplitude therefore overstates unreliability.\")\n\nfor m in models: del m\ntorch.cuda.empty_cache()\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 16. Tables 7–8 and Figure 6 — same-model deletion *(~20 min)*\n\nFour rankings are evaluated on each head and features are deleted on the head that produced them, so the attribution method is isolated from the model.","metadata":{}},{"cell_type":"code","source":"!pip install -q shap\nimport shap\n\n# Tables 7-8 and Figure 7 -- same-model faithfulness.\n# Four rankings are computed on EACH head and deleted on the SAME head that\n# produced them, so the attribution method is isolated from the model it\n# explains. A steeper QWK decline means a ranking more faithful to that head.\nDELETION_STEPS = [0.0, 0.05, 0.10, 0.20, 0.30, 0.40, 0.50]\nBUDGET_SWEEP = [100, 300, 800]\n\n\ndef pd_range_ranking(head, feat_mean, lo, hi):\n    \"\"\"Partial-dependence range. Model-agnostic: applied to both heads.\"\"\"\n    return feature_class_ranges(head, feat_mean, lo, hi, 5, n_points=60).mean(1)\n\n\ndef kernelshap_ranking(head, background, eval_feats, n_coalitions, n_background):\n    np.random.seed(SHAP_SEED)\n    dev = next(head.parameters()).device\n    predict = lambda x: F.softmax(head(torch.tensor(x, dtype=torch.float32, device=dev)), 1).detach().cpu().numpy()\n    d = background.shape[1]\n    explainer = shap.KernelExplainer(predict, background[:n_background])\n    sv = explainer.shap_values(eval_feats, nsamples=n_coalitions)\n    arr = np.stack([np.abs(s) for s in sv]) if isinstance(sv, list) else np.abs(np.asarray(sv))\n    axes = tuple(a for a in range(arr.ndim) if arr.shape[a] != d)\n    return np.asarray(arr.mean(axis=axes) if axes else arr).reshape(-1)\n\n\ndef integrated_gradients_ranking(head, eval_feats, baseline, n_steps=32, n_classes=5):\n    \"\"\"Baseline = training-set mean, i.e. the vector the deletion test imputes.\"\"\"\n    dev = next(head.parameters()).device\n    X = torch.tensor(eval_feats, dtype=torch.float32, device=dev)\n    b = torch.tensor(baseline, dtype=torch.float32, device=dev).unsqueeze(0)\n    attr = torch.zeros(X.shape[1], device=dev)\n    for c in range(n_classes):\n        grads = torch.zeros_like(X)\n        for a in torch.linspace(0, 1, n_steps, device=dev):\n            xi = (b + a * (X - b)).detach().requires_grad_(True)\n            out = F.softmax(head(xi), 1)[:, c].sum()\n            g, = torch.autograd.grad(out, xi)\n            grads += g.detach()\n        attr += ((X - b) * grads / n_steps).abs().mean(0)\n    return (attr / n_classes).detach().cpu().numpy()\n\n\ndef area_over_random(curve, random_curve, steps):\n    return float(np.trapz(np.asarray(random_curve) - np.asarray(curve), steps))\n\n\ncurve_rows, area_rows, sweep_rows = [], [], []\nfor fold in range(5):\n    models = {\"KAN\": load_head_model(\"kan\", fold), \"MLP\": load_head_model(\"mlp\", fold)}\n    val_idx = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold{fold}.npz\")[\"val_idx\"]\n    train_idx = np.setdiff1d(np.arange(len(aptos[\"df\"])), val_idx)\n    val_loader = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                             aptos[\"ext\"], val_idx, eval_transform, batch_size=16, shuffle=False)\n    train_loader = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"], aptos[\"label_col\"],\n                               aptos[\"ext\"], train_idx, eval_transform, batch_size=16, shuffle=False)\n\n    for tag, model in models.items():\n        # Each head is probed on its own backbone features (independent fine-tuning).\n        vf, vy = extract_backbone_features(model, val_loader)\n        tf, _ = extract_backbone_features(model, train_loader)\n        mean_v = tf.mean(0)\n        lo, hi = np.percentile(vf, 1, 0), np.percentile(vf, 99, 0)\n        head = model.head\n        rng = np.random.default_rng(SEED + fold)\n\n        rankings = {\n            \"PD-range\":            pd_range_ranking(head, vf.mean(0), lo, hi),\n            \"IntegratedGradients\": integrated_gradients_ranking(head, vf[:SHAP_EVAL_POINTS], mean_v),\n            \"KernelSHAP\":          kernelshap_ranking(head, tf, vf[:SHAP_EVAL_POINTS], SHAP_COALITIONS, SHAP_BACKGROUND),\n            \"Random\":              rng.permutation(vf.shape[1]).astype(float),\n        }\n        curves = {k: deletion_curve(head, vf, vy, r, DELETION_STEPS, mean_v) for k, r in rankings.items()}\n        for k, cur in curves.items():\n            for step, q in zip(DELETION_STEPS, cur):\n                curve_rows.append({\"fold\": fold, \"model\": tag, \"method\": k, \"pct_removed\": step, \"QWK\": q})\n            area_rows.append({\"fold\": fold, \"model\": tag, \"method\": k,\n                              \"area_over_random\": area_over_random(cur, curves[\"Random\"], DELETION_STEPS)})\n\n        if fold == 0:   # budget sensitivity, one fold\n            for n_coal in BUDGET_SWEEP:\n                r = kernelshap_ranking(head, tf[:40], vf[:120], n_coal, 40)\n                cur = deletion_curve(head, vf, vy, r, DELETION_STEPS, mean_v)\n                sweep_rows.append({\"model\": tag, \"n_coalitions\": n_coal,\n                                   \"area_over_random\": area_over_random(cur, curves[\"Random\"], DELETION_STEPS)})\n\n    for m in models.values(): del m\n    torch.cuda.empty_cache()\n    print(f\"fold {fold} done\")\n\ncurves_df = pd.DataFrame(curve_rows)\nareas_df = pd.DataFrame(area_rows)\nsweep_df = pd.DataFrame(sweep_rows)\n\nprint(\"\\nTable 7 -- area over random (five-fold mean +/- std). Higher = more faithful.\")\nprint(areas_df.groupby([\"model\", \"method\"])[\"area_over_random\"].agg([\"mean\", \"std\"]).round(4).to_string())\n\nprint(\"\\nDeletion curves (five-fold mean QWK)\")\nprint(curves_df.groupby([\"model\", \"method\", \"pct_removed\"])[\"QWK\"].mean().round(3)\n      .reset_index().pivot_table(index=[\"model\", \"method\"], columns=\"pct_removed\", values=\"QWK\").to_string())\n\nprint(\"\\nTable 8 -- KernelSHAP budget sensitivity (fold 0)\")\nprint(sweep_df.pivot_table(index=\"model\", columns=\"n_coalitions\", values=\"area_over_random\").round(4).to_string())\n\ncurves_df.to_csv(OUT_DIR / \"table7_deletion_curves.csv\", index=False)\nareas_df.to_csv(OUT_DIR / \"table7_areas.csv\", index=False)\nsweep_df.to_csv(OUT_DIR / \"table8_shap_budget.csv\", index=False)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Figure 7 -- deletion curves for the same-model faithfulness result.\nmean_curves = curves_df.groupby([\"model\", \"method\", \"pct_removed\"])[\"QWK\"].mean().reset_index()\norder = [\"PD-range\", \"IntegratedGradients\", \"KernelSHAP\", \"Random\"]\npretty = {\"PD-range\": \"PD-range\", \"IntegratedGradients\": \"Integrated Gradients\",\n          \"KernelSHAP\": \"KernelSHAP\", \"Random\": \"Random\"}\nstyle = {\"PD-range\": (\"-\", \"o\"), \"IntegratedGradients\": (\"-\", \"s\"),\n         \"KernelSHAP\": (\"--\", \"^\"), \"Random\": (\":\", \"x\")}\n\nfig, axes = plt.subplots(1, 2, figsize=(10, 4.2), sharey=True)\nfor ax, model in zip(axes, [\"KAN\", \"MLP\"]):\n    for method in order:\n        d = mean_curves[(mean_curves[\"model\"] == model) & (mean_curves[\"method\"] == method)].sort_values(\"pct_removed\")\n        ls, mk = style[method]\n        ax.plot(d[\"pct_removed\"] * 100, d[\"QWK\"], ls, marker=mk, ms=4, label=pretty[method])\n    ax.set_title(f\"{model} head\")\n    ax.set_xlabel(\"Features removed (%)\")\n    ax.grid(alpha=0.3)\naxes[0].set_ylabel(\"Quadratic weighted kappa\")\naxes[0].legend(fontsize=8, loc=\"lower left\")\nplt.tight_layout()\nPath(\"figs\").mkdir(exist_ok=True)\nplt.savefig(\"figs/figure_7_deletion_curves.pdf\", dpi=300)\nplt.savefig(OUT_DIR / \"figure_7_deletion_curves.pdf\", dpi=300)\nplt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 17. Table 9 — removal-operator invariance *(heavy, ~10 min)*","metadata":{}},{"cell_type":"code","source":"# Table 9 -- removal-operator invariance.\n#\n# Mean-imputation deletion shares its reference vector with PD-range and with the\n# IG baseline, which could favour them by construction (Sturmfels et al. 2020).\n# If PD-range leads only under `mean`, the earlier result was a metric artefact.\nif not RUN_HEAVY:\n    print(\"RUN_HEAVY = False -> skipping the removal-operator control (Table 9).\")\nelse:\n    rows = []\n    for fold in range(5):\n        for tag in (\"kan\", \"mlp\"):\n            tf, ty, vf, vy, head = fold_features(tag, fold)\n            head = head.to(DEVICE).eval()\n            lo, hi = np.percentile(vf, 1, 0), np.percentile(vf, 99, 0)\n            mean_v = tf.mean(0)\n            R = {\"PD-range\":  rank_pd(head, vf, lo, hi),\n                 \"IG\":        rank_ig(head, vf[:200], mean_v),\n                 \"SHAP-rand\": rank_shap(head, tf[:100], vf[:200]),\n                 \"SHAP-mean\": rank_shap(head, mean_v.reshape(1, -1), vf[:200]),\n                 \"Random\":    np.random.default_rng(SEED + fold).permutation(vf.shape[1]).astype(float)}\n            for op in (\"marginal\", \"mean\", \"zero\"):\n                rng = np.random.default_rng(SEED + fold)\n                curves = {}\n                for name, r in R.items():\n                    order = np.argsort(-r)\n                    curves[name] = [qwk_of(head, ablate(vf, order[:int(round(k * vf.shape[1]))],\n                                                        op, tf, rng), vy) for k in DELETION_STEPS]\n                for name, cur in curves.items():\n                    rows.append({\"fold\": fold, \"head\": tag.upper(), \"operator\": op,\n                                 \"ranking\": name, \"AOR\": aor(cur, curves[\"Random\"])})\n            del head; torch.cuda.empty_cache()\n        print(f\"  fold {fold} done\")\n\n    ops = pd.DataFrame(rows)\n    ops.to_csv(OUT_DIR / \"table9_operators.csv\", index=False)\n    print(\"\\nTable 9 -- AOR (five-fold mean) by removal operator\")\n    print(ops.pivot_table(index=[\"head\", \"ranking\"], columns=\"operator\",\n                          values=\"AOR\", aggfunc=\"mean\").round(4).to_string())\n    print(\"\\nPD-range and IG lead under all three operators, and by the widest margin\")\n    print(\"under `marginal`, which discards the mean reference entirely. The matched\")\n    print(\"mean baseline does improve KernelSHAP, but leaves it an order of magnitude behind.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 18. Table 10 — remove and retrain *(heavy, ~25 min)*\n\nThe control that reverses the picture: after retraining, no ranking beats random.","metadata":{}},{"cell_type":"code","source":"# Table 10 -- remove and retrain (Hooker et al., NeurIPS 2019).\n#\n# Deleting the features a model relies on pushes its input off the manifold it was\n# fitted on, so the degradation can measure distribution shift rather than lost\n# information. The corrective is to retrain. Here the backbone is fixed, so the\n# features are cached and a ~100K-parameter head retrains in seconds.\nif not RUN_HEAVY:\n    print(\"RUN_HEAVY = False -> skipping the remove-and-retrain control (Table 10).\")\nelse:\n    ROAR_STEPS = [0.0, 0.10, 0.30, 0.50]\n    rows = []\n    for fold in range(5):\n        for tag in (\"kan\", \"mlp\"):\n            tf, ty, vf, vy, head = fold_features(tag, fold)\n            head = head.to(DEVICE).eval()\n            lo, hi = np.percentile(vf, 1, 0), np.percentile(vf, 99, 0)\n            mean_v = tf.mean(0)\n            R = {\"PD-range\":  rank_pd(head, vf, lo, hi),\n                 \"IG\":        rank_ig(head, vf[:200], mean_v),\n                 \"SHAP-mean\": rank_shap(head, mean_v.reshape(1, -1), vf[:200]),\n                 \"Random\":    np.random.default_rng(SEED + fold).permutation(vf.shape[1]).astype(float)}\n            del head; torch.cuda.empty_cache()\n\n            proto = make_head(SHOWCASE_BACKBONE, tag, \"as-published\")\n            rng = np.random.default_rng(SEED + fold)\n            for name, r in R.items():\n                order = np.argsort(-r)\n                for k in ROAR_STEPS:\n                    n_rm = int(round(k * vf.shape[1]))\n                    Xtr = ablate(tf, order[:n_rm], \"mean\", tf, rng)\n                    Xva = ablate(vf, order[:n_rm], \"mean\", tf, rng)\n                    h = fit_head(proto, Xtr, ty)                 # fresh head, from scratch\n                    rows.append({\"fold\": fold, \"head\": tag.upper(), \"ranking\": name,\n                                 \"removed\": k, \"QWK\": qwk_of(h, Xva, vy)})\n                    del h; torch.cuda.empty_cache()\n        print(f\"  fold {fold} done\")\n\n    roar = pd.DataFrame(rows)\n    roar.to_csv(OUT_DIR / \"table10_roar.csv\", index=False)\n    print(\"\\nTable 10 -- QWK after retraining (five-fold mean)\")\n    print(roar.pivot_table(index=[\"head\", \"ranking\"], columns=\"removed\",\n                           values=\"QWK\", aggfunc=\"mean\").round(3).to_string())\n    print(\"\\nAfter retraining no ranking degrades faster than a random ordering. The\")\n    print(\"features any ranking selects are what the fitted head uses, not what the task\")\n    print(\"requires: the pooled backbone representation is redundant. A deletion curve\")\n    print(\"computed without retraining cannot distinguish the two.\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================================================================\n#  FIGURE 4  --  REPLACEMENT FOR CELL A ONLY\n#\n#  My earlier version put the legend on top of the x-axis labels. The\n#  cause was placing the legend at a fixed negative offset while\n#  bbox_inches=\"tight\" was re-cropping the canvas, so the reserved space\n#  collapsed and the two collided.\n#\n#  Fixed by reserving the vertical strip first and only then placing the\n#  legend into it:\n#\n#      fig.tight_layout(rect=[0, 0.13, 1, 1])   # keep the bottom 13% free\n#      fig.legend(..., loc=\"lower center\", bbox_to_anchor=(0.5, 0.005))\n#      fig.savefig(..., bbox_inches=None)       # do NOT re-crop\n#\n#  I rendered this layout and checked it before sending it. The output is\n#  exactly 17.4 x 7.4 cm, which is the journal's full text width, so\n#  LaTeX places it at native size with no rescaling.\n#\n#  Paste over your existing CELL A and run just that cell.\n# =====================================================================\n\nimport numpy as np, torch, matplotlib.pyplot as plt\nfrom pathlib import Path\n\nCM = 1 / 2.54\nplt.rcParams.update({\n    \"font.size\": 9, \"axes.labelsize\": 9, \"axes.titlesize\": 9,\n    \"xtick.labelsize\": 8, \"ytick.labelsize\": 8, \"legend.fontsize\": 8,\n    \"axes.spines.top\": False, \"axes.spines.right\": False,\n})\nPath(\"figs\").mkdir(exist_ok=True)\n\nFEATURES = [1141, 1000]\nFOLD     = 1\n\nm = load_head_model(\"kan\", FOLD)\nref_val = np.load(OUT_DIR / f\"val_preds_aptos_{SHOWCASE_BACKBONE}_kan_fold{FOLD}.npz\")[\"val_idx\"]\nref_loader = make_loader(aptos[\"df\"], aptos[\"image_dir\"], aptos[\"image_col\"],\n                         aptos[\"label_col\"], aptos[\"ext\"], ref_val,\n                         eval_transform, batch_size=16, shuffle=False)\nref_feats, _ = extract_backbone_features(m, ref_loader)\nfeat_means = ref_feats.mean(0)\nfeat_lo, feat_hi = np.percentile(ref_feats, 1, 0), np.percentile(ref_feats, 99, 0)\n\ncolours = plt.cm.viridis(np.linspace(0.05, 0.9, 5))\nfig, axes = plt.subplots(1, 2, figsize=(17.4 * CM, 7.4 * CM), sharey=True)\n\nfor ax, fi in zip(axes, FEATURES):\n    grid, curves = head_class_curves(m.head, int(fi), feat_means,\n                                     feat_lo[fi], feat_hi[fi], 200)\n    for c in range(5):\n        ax.plot(grid, curves[:, c], lw=1.6, color=colours[c],\n                label=class_names_5[c])\n    ax.axhline(0, color=\"0.75\", lw=0.7, zorder=0)\n    ax.set_title(f\"Feature {fi}\")\n    ax.set_xlabel(\"Feature value (1st\\u201399th percentile)\")\n\naxes[0].set_ylabel(\"Class logit\")\n\n# reserve the strip BEFORE placing the legend, then do not re-crop\nfig.tight_layout(rect=[0, 0.13, 1, 1])\nh, l = axes[0].get_legend_handles_labels()\nfig.legend(h, l, loc=\"lower center\", ncol=5, frameon=False,\n           bbox_to_anchor=(0.5, 0.005), columnspacing=2.0, handletextpad=0.5)\n\nfig.savefig(\"figs/figure_4.pdf\", bbox_inches=None)\nfig.savefig(OUT_DIR / \"figure_4.pdf\", bbox_inches=None)\nplt.show()\nprint(\"wrote figs/figure_4.pdf   17.4 x 7.4 cm, vector\")\n\ndel m\ntorch.cuda.empty_cache()\n\n# =====================================================================\n#  CHECK IT YOURSELF before moving on: the legend row must sit clearly\n#  BELOW both \"Feature value\" labels, with visible white space between\n#  them. If they touch, raise 0.13 to 0.16 and re-run this cell.\n# =====================================================================","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================================================================\n#  FIGURE 5  --  REPLACEMENT FOR CELL B ONLY\n#\n#  My previous version of this cell had two collisions I did not check\n#  before sending it. I have now rendered this layout and measured every\n#  label's bounding box against every other one and against all 160 data\n#  points. All clear:\n#\n#      legend vs x-axis label ....... clear\n#      legend vs band label ......... clear\n#      band label vs rho label ...... clear\n#      data points on band label .... 0\n#      data points on rho label ..... 0\n#      legend inside the canvas ..... yes\n#      page size .................... 8.4 x 9.2 cm (single column)\n#\n#  What changed from the broken version:\n#    1. The legend was landing on the x-axis label, the same bug as in\n#       Figure 4. Fixed the same way: reserve the strip with\n#       tight_layout(rect=...) FIRST, then place a fig.legend into it,\n#       and save with bbox_inches=None so the canvas is not re-cropped.\n#    2. \"Spearman rho\" sat on top of 37 data points in the r = 1.0\n#       cluster. Both annotations now sit below the shaded band, in the\n#       empty strip where no point can fall.\n#    3. The panel is taller (9.2 cm) so the two clusters separate.\n#\n#  Paste over your existing CELL B. It reads sign_consistency.csv from\n#  disk: no GPU, no model, about two seconds.\n# =====================================================================\n\nimport numpy as np, pandas as pd, matplotlib.pyplot as plt\nfrom scipy import stats\nfrom pathlib import Path\n\nCM = 1 / 2.54\nplt.rcParams.update({\n    \"font.size\": 9, \"axes.labelsize\": 9, \"axes.titlesize\": 9,\n    \"xtick.labelsize\": 8, \"ytick.labelsize\": 8, \"legend.fontsize\": 8,\n    \"axes.spines.top\": False, \"axes.spines.right\": False,\n})\nPath(\"figs\").mkdir(exist_ok=True)\n\nsc = pd.read_csv(OUT_DIR / \"sign_consistency.csv\")\nrho, pval = stats.spearmanr(sc[\"amplitude\"], sc[\"mean_r\"])\ncolours = plt.cm.viridis([0.05, 0.26, 0.47, 0.68, 0.90])\nmarkers = [\"o\", \"s\", \"^\", \"D\", \"v\"]\n\nfig, ax = plt.subplots(figsize=(8.4 * CM, 9.2 * CM))\nax.axhspan(-0.25, 0.25, color=\"0.93\", zorder=0, lw=0)\nax.axhline(0.85, color=\"0.55\", lw=0.8, ls=\"--\", zorder=1)\n\nfor c, name in enumerate(class_names_5):\n    sub = sc[sc[\"class\"] == name]\n    ax.scatter(sub[\"amplitude\"], sub[\"mean_r\"], s=20, marker=markers[c],\n               facecolor=colours[c], edgecolor=\"white\", linewidth=0.5,\n               alpha=0.85, label=name, zorder=3)\n\nax.set_xscale(\"log\")\nax.set_xlabel(\"Curve amplitude (logit range)\")\nax.set_ylabel(\"Mean pairwise Pearson $r$\")\nax.set_ylim(-0.62, 1.12)\n\n# both annotations live below the band, in the strip that holds no data\nax.text(sc[\"amplitude\"].max() * 0.97, -0.45, \"sign-inversion region\",\n        ha=\"right\", va=\"center\", fontsize=7, color=\"0.40\")\nax.text(sc[\"amplitude\"].min() * 1.02, -0.45, rf\"Spearman $\\rho$ = {rho:.2f}\",\n        ha=\"left\", va=\"center\", fontsize=8)\n\n# reserve the strip BEFORE placing the legend, then do not re-crop\nfig.tight_layout(rect=[0, 0.20, 1, 1])\nh, l = ax.get_legend_handles_labels()\nfig.legend(h, l, loc=\"lower center\", bbox_to_anchor=(0.5, 0.005), ncol=3,\n           frameon=False, handletextpad=0.3, columnspacing=1.3,\n           borderpad=0.2, labelspacing=0.35)\n\nfig.savefig(\"figs/figure_5.pdf\", bbox_inches=None)\nfig.savefig(OUT_DIR / \"figure_5.pdf\", bbox_inches=None)\nplt.show()\nprint(\"wrote figs/figure_5.pdf   8.4 x 9.2 cm, vector\")\nprint(f\"n = {len(sc)} combinations   rho = {rho:.3f}   p = {pval:.2e}\")\n\n# ---- self-check: prints a warning if anything still collides ----\nfig.canvas.draw(); _r = fig.canvas.get_renderer()\ndef _bb(a):\n    b = a.get_window_extent(renderer=_r); return (b.x0, b.y0, b.x1, b.y1)\ndef _ov(a, b):\n    return not (a[2] <= b[0] or b[2] <= a[0] or a[3] <= b[1] or b[3] <= a[1])\n_leg, _xl = _bb(fig.legends[0]), _bb(ax.xaxis.label)\n_pts = ax.transData.transform(np.c_[sc[\"amplitude\"], sc[\"mean_r\"]])\n_hits = sum(1 for t in ax.texts\n            for x, y in _pts\n            if _bb(t)[0]-3 <= x <= _bb(t)[2]+3 and _bb(t)[1]-3 <= y <= _bb(t)[3]+3)\nprint(\"layout check -> legend/xlabel:\", \"OVERLAP\" if _ov(_leg, _xl) else \"clear\",\n      \"| points under labels:\", _hits)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================================================================\n#  TWO CELLS TO APPEND AT THE END OF YOUR KAGGLE NOTEBOOK\n#\n#  Do not edit any existing cell. Paste each block below into a new\n#  cell at the bottom.\n#\n#  WHAT TO RUN FIRST (all fast -- definitions and checkpoint loading):\n#      cells 0-24\n#\n#  WHAT TO SKIP (the old heavy ones, not needed by either cell below):\n#      cell 26  Table 5,  frozen-backbone probe   ~35 min\n#      cell 35  Tables 7-8, deletion protocol     ~20 min\n#      cell 38  Table 9,  removal operators       ~10 min\n#\n#  Both cells below depend only on cell 24's utilities\n#  (fold_features, rank_pd, rank_ig, rank_shap, ablate, make_head,\n#  fit_head, qwk_of, aor, DELETION_STEPS) plus config from cells 2-3.\n#\n#  Make sure RUN_HEAVY = True in cell 3.\n# =====================================================================\n \n \n# =====================================================================\n#  CELL 1  --  B2: KernelSHAP beyond M = d          (~30-60 min)\n#\n#  Section 4.5 argues that the flat KernelSHAP deletion curve is an\n#  underdetermination artefact: the estimator solves for d = 1280\n#  coefficients from M sampled coalitions, and every budget in the\n#  literature has M < d. The account predicts the deficit closes once\n#  M >= d. The manuscript currently cannot test that, because no budget\n#  we ran reached it. This cell runs M = 1600 > d and settles it.\n#\n#  Fold 0 only, both heads. rank_shap already takes nsamples, so\n#  nothing new is defined here.\n# =====================================================================\n \n# ---- SMOKE TEST SWITCH -------------------------------------------\n# True  = tiny run, ~3 min, just proves the code works\n# False = the real run\nSMOKE_TEST = False\n# ------------------------------------------------------------------\n \nBUDGETS = [100, 1600] if SMOKE_TEST else [100, 300, 800, 1600]\nFOLD    = 0\n \nsweep_rows = []\nfor tag in (\"kan\", \"mlp\"):\n    tf, ty, vf, vy, head = fold_features(tag, FOLD)\n    head = head.to(DEVICE).eval()\n    mean_v = tf.mean(0)\n    rng = np.random.default_rng(SEED + FOLD)\n \n    # random control, shared across budgets\n    rand_order = rng.permutation(vf.shape[1])\n    rand_curve = [qwk_of(head, ablate(vf, rand_order[:int(round(k * vf.shape[1]))],\n                                      \"mean\", tf, rng), vy) for k in DELETION_STEPS]\n \n    for M in BUDGETS:\n        r = rank_shap(head, mean_v.reshape(1, -1), vf[:200], nsamples=M)\n        order = np.argsort(-r)\n        curve = [qwk_of(head, ablate(vf, order[:int(round(k * vf.shape[1]))],\n                                     \"mean\", tf, rng), vy) for k in DELETION_STEPS]\n        a = aor(curve, rand_curve)\n        sweep_rows.append({\"head\": tag.upper(), \"M\": M, \"M_over_d\": M / vf.shape[1],\n                           \"AOR\": a})\n        print(f\"  {tag.upper()}  M={M:5d}  (M/d={M/vf.shape[1]:.2f})  AOR={a:+.4f}\")\n \n    del head; torch.cuda.empty_cache()\n \nshap_sweep = pd.DataFrame(sweep_rows)\nshap_sweep.to_csv(OUT_DIR / \"shap_budget_beyond_d.csv\", index=False)\n \nprint(\"\\nKernelSHAP budget sweep, fold 0\")\nprint(shap_sweep.pivot_table(index=\"head\", columns=\"M\", values=\"AOR\").round(4).to_string())\nprint(\"\\nRead it like this. If AOR keeps climbing at M = 1600 and lands near\")\nprint(\"PD-range's value (KAN 0.056, MLP 0.045), the underdetermination account\")\nprint(\"is confirmed and Section 4.5 can drop its 'untested at M >= d' caveat.\")\nprint(\"If AOR stays flat at M = 1600, the account is WRONG: the deficit is not\")\nprint(\"a sampling artefact, and Correction 3 must be rewritten, not just softened.\")\n \n ","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================================================================\n#  CELL 2  --  B1: multi-seed remove-and-retrain    (~4 h at 10 seeds)\n#\n#  The single-seed version cannot separate a 0.003-0.007 QWK\n#  ranking-versus-random gap from the head's own +/-0.005 run-to-run\n#  variability. This is, by the manuscript's own admission, the thinnest\n#  evidence it offers for its strongest negative claim -- and the claim\n#  is now in the title.\n#\n#  Rankings are deterministic given the fitted head, so they are computed\n#  once per fold and head; only the retraining is reseeded. That is about\n#  four times faster than reseeding the whole loop.\n#\n#  Start with N_ROAR_SEEDS = 10. Kaggle allows 12 h GPU sessions, so this\n#  finishes in one overnight run with room to spare.\n# =====================================================================\n \n# Uses the same SMOKE_TEST switch set in the cell above.\nROAR_STEPS   = [0.0, 0.10, 0.30, 0.50]\nN_ROAR_SEEDS = 2 if SMOKE_TEST else 10\nN_FOLDS      = 1 if SMOKE_TEST else 5\nROAR_SEEDS   = [SEED + 100 * j for j in range(N_ROAR_SEEDS)]\n \nif SMOKE_TEST:\n    print(\">>> SMOKE TEST: 1 fold, 2 seeds. Numbers are NOT usable.\")\n    print(\">>> If this finishes without an error, set SMOKE_TEST = False and commit.\\n\")\n \nrows = []\nfor fold in range(N_FOLDS):\n    for tag in (\"kan\", \"mlp\"):\n        tf, ty, vf, vy, head = fold_features(tag, fold)\n        head = head.to(DEVICE).eval()\n        lo, hi = np.percentile(vf, 1, 0), np.percentile(vf, 99, 0)\n        mean_v = tf.mean(0)\n        R = {\"PD-range\":  rank_pd(head, vf, lo, hi),\n             \"IG\":        rank_ig(head, vf[:200], mean_v),\n             \"SHAP-mean\": rank_shap(head, mean_v.reshape(1, -1), vf[:200]),\n             \"Random\":    np.random.default_rng(SEED + fold).permutation(vf.shape[1]).astype(float)}\n        del head; torch.cuda.empty_cache()\n \n        #proto = make_head(SHOWCASE_BACKBONE, tag, \"as-published\")\n        proto = make_head(SHOWCASE_BACKBONE, tag, \"no-dropout\")\n        for name, r in R.items():\n            order = np.argsort(-r)\n            for k in ROAR_STEPS:\n                n_rm = int(round(k * vf.shape[1]))\n                for sd in ROAR_SEEDS:\n                    rng = np.random.default_rng(sd + fold)\n                    Xtr = ablate(tf, order[:n_rm], \"mean\", tf, rng)\n                    Xva = ablate(vf, order[:n_rm], \"mean\", tf, rng)\n                    h = fit_head(proto, Xtr, ty, seed=sd)      # fresh head, reseeded\n                    rows.append({\"fold\": fold, \"seed\": sd, \"head\": tag.upper(),\n                                 \"ranking\": name, \"removed\": k,\n                                 \"QWK\": qwk_of(h, Xva, vy)})\n                    del h; torch.cuda.empty_cache()\n    print(f\"  fold {fold} done ({N_ROAR_SEEDS} seeds)\")\n \nroar = pd.DataFrame(rows)\n#roar.to_csv(OUT_DIR / \"table10_roar_multiseed.csv\", index=False)\nroar.to_csv(OUT_DIR / \"table10_roar_multiseed_nodropout.csv\", index=False)\n \nprint(f\"\\nQWK after retraining (mean over 5 folds x {N_ROAR_SEEDS} seeds)\")\nprint(roar.pivot_table(index=[\"head\", \"ranking\"], columns=\"removed\",\n                       values=\"QWK\", aggfunc=\"mean\").round(4).to_string())\nprint(\"\\nSeed-level standard deviation\")\nprint(roar.pivot_table(index=[\"head\", \"ranking\"], columns=\"removed\",\n                       values=\"QWK\", aggfunc=\"std\").round(4).to_string())\n \n# Paired test at 50% removal: does any ranking degrade faster than random?\nfrom scipy.stats import wilcoxon\nprint(\"\\nRanking vs Random at 50% removal, paired over fold x seed\")\nprint(\"(negative delta = ranking degrades MORE than random, i.e. it locates something)\")\ntop = roar[roar[\"removed\"] == 0.50]\nfor hd in (\"KAN\", \"MLP\"):\n    base = top[(top[\"head\"] == hd) & (top[\"ranking\"] == \"Random\")] \\\n              .set_index([\"fold\", \"seed\"])[\"QWK\"]\n    for name in (\"PD-range\", \"IG\", \"SHAP-mean\"):\n        m = top[(top[\"head\"] == hd) & (top[\"ranking\"] == name)] \\\n               .set_index([\"fold\", \"seed\"])[\"QWK\"]\n        d = (m - base).dropna()\n        try:\n            _, p = wilcoxon(d)\n        except ValueError:\n            p = float(\"nan\")\n        print(f\"  {hd:3s} {name:10s} delta = {d.mean():+.4f}  sd = {d.std():.4f} \"\n              f\" n = {len(d):3d}  Wilcoxon p = {p:.3f}\")\n \nif SMOKE_TEST:\n    print(\"\\n>>> SMOKE TEST ONLY. Set SMOKE_TEST = False in the cell above,\")\n    print(\">>> then Save Version -> Save & Run All (Commit).\")\nelse:\n    print(\"\\nIf every delta is indistinguishable from zero, the manuscript's claim\")\n    print(\"stands and is now properly powered.\")\nprint(\"If any delta is reliably negative, then reliance IS necessity for that\")\nprint(\"ranking, and Section 4.5, the abstract and the title all need revising.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Feature range diagnostic: does the pooled representation fit the spline grid?\n# KANLinear grid_range=(-2,2); basis support dies out past |x|=4.4.\nrows = []\nfor fold in range(5):\n    tf, ty, vf, vy, head = fold_features(\"kan\", fold)\n    F = np.vstack([tf, vf])\n    tot = F.size\n    rows.append({\n        \"fold\": fold,\n        \"in_grid_|x|<=2\":   (np.abs(F) <= 2.0).sum() / tot,\n        \"taper_2<|x|<=4.4\": ((np.abs(F) > 2.0) & (np.abs(F) <= 4.4)).sum() / tot,\n        \"dead_|x|>4.4\":     (np.abs(F) > 4.4).sum() / tot,\n        \"p50\": np.percentile(F, 50), \"p99\": np.percentile(F, 99),\n        \"min\": F.min(), \"max\": F.max(),\n    })\n    print(rows[-1])\n    del head, tf, vf, F; torch.cuda.empty_cache()\n\ndiag = pd.DataFrame(rows)\ndiag.to_csv(OUT_DIR / \"feature_grid_diagnostic.csv\", index=False)\nprint(\"\\n=== MEAN OVER FOLDS ===\")\nprint(diag.drop(columns=\"fold\").mean().round(4).to_string())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T17:16:41.622344Z","iopub.execute_input":"2026-07-24T17:16:41.623485Z","iopub.status.idle":"2026-07-24T17:17:17.931220Z","shell.execute_reply.started":"2026-07-24T17:16:41.623442Z","shell.execute_reply":"2026-07-24T17:17:17.930243Z"}},"outputs":[],"execution_count":null}]}