{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\n\nimport matplotlib.pyplot as plt\nfrom scipy.stats import mannwhitneyu\n\nfrom sklearn.metrics import (\n    roc_curve,\n    roc_auc_score,\n    confusion_matrix,\n    ConfusionMatrixDisplay,\n    f1_score,\n    precision_score,\n    recall_score,\n    accuracy_score,\n    average_precision_score\n)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:29.743843Z","iopub.execute_input":"2026-08-02T11:36:29.744805Z","iopub.status.idle":"2026-08-02T11:36:34.407669Z","shell.execute_reply.started":"2026-08-02T11:36:29.744702Z","shell.execute_reply":"2026-08-02T11:36:34.406698Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load All Fold Files","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\n\nfold_files = {\n\n    0: \"/kaggle/input/notebooks/ashiasultana/isic-md-missing-aw-ga-fi-asym-image-trainfold0-1/fold_0_results.pth\",\n\n    1: \"/kaggle/input/notebooks/ashiasultana/isic-md-missing-aw-ga-fi-asym-image-trainfold0-1/fold_1_results.pth\",\n\n    2: \"/kaggle/input/notebooks/ashiasultana/isic-md-missing-aw-ga-fi-asym-image-trainfold2-3/fold_2_results.pth\",\n\n    3: \"/kaggle/input/notebooks/ashiasultana/isic-md-missing-aw-ga-fi-asym-image-trainfold2-3/fold_3_results.pth\",\n\n    4: \"/kaggle/input/notebooks/ashiasultana/isic-md-missing-aw-ga-fi-asym-image-trainfold-4/fold_4_results.pth\"\n}\n\n\n\nFOLD_RESULTS = {}\n\nfor fold, path in fold_files.items():\n\n    FOLD_RESULTS[fold] = torch.load(\n        path,\n        weights_only=False\n    )\n\nprint(\"Loaded folds:\", list(FOLD_RESULTS.keys()))\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:34.409122Z","iopub.execute_input":"2026-08-02T11:36:34.409636Z","iopub.status.idle":"2026-08-02T11:36:34.568707Z","shell.execute_reply.started":"2026-08-02T11:36:34.409595Z","shell.execute_reply":"2026-08-02T11:36:34.567745Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ALL_GATES = {}\n\nfor fold in FOLD_RESULTS:\n\n    gate_stats = FOLD_RESULTS[fold].get(\"gate_stats\", {})\n\n    ALL_GATES[fold] = gate_stats\n\nprint(\"Gate stats loaded for folds:\", list(ALL_GATES.keys()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:34.569848Z","iopub.execute_input":"2026-08-02T11:36:34.570155Z","iopub.status.idle":"2026-08-02T11:36:34.576659Z","shell.execute_reply.started":"2026-08-02T11:36:34.570129Z","shell.execute_reply":"2026-08-02T11:36:34.575497Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Reconstruct OOF Predictions\nReconstruct Out-of-Fold (OOF) predictions\nDuring 5-fold cross-validation, each fold predicts only\nits own validation set. This code combines the predictions\nfrom all folds back into their original dataset order so\nthat every sample has exactly one prediction from the model\nthat did not train on it.","metadata":{}},{"cell_type":"code","source":"max_idx = max(\n    FOLD_RESULTS[f][\"val_idx\"].max()\n    for f in FOLD_RESULTS\n)\n\nOOF_PREDS = np.zeros(max_idx + 1)\n\nOOF_TARGETS = np.zeros(max_idx + 1)\n\nfor fold in FOLD_RESULTS:\n\n    idx = FOLD_RESULTS[fold][\"val_idx\"]\n\n    OOF_PREDS[idx] = (\n        FOLD_RESULTS[fold][\"preds\"]\n    )\n\n    OOF_TARGETS[idx] = (\n        FOLD_RESULTS[fold][\"targets\"]\n    )\n\nprint(\"OOF reconstruction complete\")\nprint(\"Samples:\", len(OOF_PREDS))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:34.578579Z","iopub.execute_input":"2026-08-02T11:36:34.578885Z","iopub.status.idle":"2026-08-02T11:36:34.596528Z","shell.execute_reply.started":"2026-08-02T11:36:34.578858Z","shell.execute_reply":"2026-08-02T11:36:34.595544Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\n    \"\\n\"\n    + \"=\"*20\n    + \" FINAL CV RESULTS \"\n    + \"=\"*20\n)\n\nfor fold in sorted(FOLD_RESULTS.keys()):\n\n    print(\n        f\"Fold {fold}: \"\n        f\"{FOLD_RESULTS[fold]['best_auc']:.4f}\"\n    )\n\nfold_scores = [\n\n    FOLD_RESULTS[f][\"best_auc\"]\n\n    for f in sorted(FOLD_RESULTS.keys())\n]\n\nmean_auc = np.mean(\n    fold_scores\n)\n\nstd_auc = np.std(\n    fold_scores\n)\n\nprint(\n    f\"\\nMean AUC : {mean_auc:.4f}\"\n)\n\nprint(\n    f\"Std AUC  : {std_auc:.4f}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:34.597840Z","iopub.execute_input":"2026-08-02T11:36:34.598216Z","iopub.status.idle":"2026-08-02T11:36:34.615823Z","shell.execute_reply.started":"2026-08-02T11:36:34.598185Z","shell.execute_reply":"2026-08-02T11:36:34.614873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"OOF_GATES = np.zeros_like(OOF_PREDS, dtype=float)\n\nfor fold in FOLD_RESULTS:\n\n    idx = FOLD_RESULTS[fold][\"val_idx\"]\n\n    # IMPORTANT: only works if you saved per-sample gate values\n    # if not available, fallback handled below\n\n    if \"gate_values\" in FOLD_RESULTS[fold]:\n\n        OOF_GATES[idx] = FOLD_RESULTS[fold][\"gate_values\"]\n\nprint(\"OOF gate reconstruction complete\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:34.616836Z","iopub.execute_input":"2026-08-02T11:36:34.617210Z","iopub.status.idle":"2026-08-02T11:36:34.639467Z","shell.execute_reply.started":"2026-08-02T11:36:34.617172Z","shell.execute_reply":"2026-08-02T11:36:34.638507Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ROC Curve","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,6))\n\nfor fold in sorted(FOLD_RESULTS.keys()):\n\n    roc_history = FOLD_RESULTS[fold][\"roc_history\"]\n\n    fpr, tpr, auc_fold = roc_history[-1]\n\n    plt.plot(\n        fpr,\n        tpr,\n        alpha=0.7,\n        label=f\"Fold {fold} (AUC={auc_fold:.3f})\"\n    )\n\nfpr_oof, tpr_oof, _ = roc_curve(\n    OOF_TARGETS,\n    OOF_PREDS\n)\n\nauc_oof = roc_auc_score(\n    OOF_TARGETS,\n    OOF_PREDS\n)\n\nplt.plot(\n    fpr_oof,\n    tpr_oof,\n    linewidth=3,\n    label=f\"OOF Mean (AUC={auc_oof:.3f})\"\n)\n\nplt.plot([0,1],[0,1],\"k--\")\n\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.title(\"ROC Curves Across Folds\")\nplt.legend()\nplt.grid()\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingROC.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:34.640593Z","iopub.execute_input":"2026-08-02T11:36:34.641026Z","iopub.status.idle":"2026-08-02T11:36:35.471085Z","shell.execute_reply.started":"2026-08-02T11:36:34.640987Z","shell.execute_reply":"2026-08-02T11:36:35.469899Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Fold-Level Mean Gate Activation Distribution\n\nThis histogram summarizes the mean gate activation statistics obtained from each cross-validation fold. Rather than displaying gate activations for individual samples, the figure presents the average gate values recorded for different groups (e.g., overall, melanoma, benign, samples with missing metadata, and samples with complete metadata).\n\nThe distribution provides a high-level overview of the consistency of the missingness-aware gating mechanism across the five folds. Similar mean values across folds indicate stable gating behavior and suggest that the model learns a consistent strategy for modulating image features based on metadata availability.\n\n**Note:** This figure represents fold-level summary statistics. A histogram generated from the reconstructed out-of-fold (`OOF_GATES`) values provides a more detailed visualization of gate activations at the individual sample level.","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8, 5))\n\nall_gates = []\n\nfor fold in FOLD_RESULTS:\n\n    if \"gate_stats\" in FOLD_RESULTS[fold]:\n\n        gate_stats = FOLD_RESULTS[fold][\"gate_stats\"]\n\n        # Collect all available mean gate statistics\n        for key, value in gate_stats.items():\n\n            if \"mean\" in key and value is not None:\n                all_gates.append(value)\n\nplt.hist(\n    all_gates,\n    bins=20,\n    edgecolor=\"black\"\n)\n\nplt.title(\"Distribution of Fold-Level Mean Gate Activations\")\nplt.xlabel(\"Mean Gate Activation\")\nplt.ylabel(\"Frequency\")\n\n\nplt.grid(alpha=0.3)\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingFold-LevelGateAnalysis.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:35.472331Z","iopub.execute_input":"2026-08-02T11:36:35.472695Z","iopub.status.idle":"2026-08-02T11:36:35.972998Z","shell.execute_reply.started":"2026-08-02T11:36:35.472660Z","shell.execute_reply":"2026-08-02T11:36:35.971865Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Distribution of Missingness-Aware Gate Activations (Out-of-Fold)\n\nThis histogram shows the distribution of gate activation values computed for all out-of-fold (OOF) validation samples across the entire dataset (58,457 samples).\n\nEach value represents the mean activation of the missingness-aware gating mechanism for a single lesion. The gate controls how much image feature information is retained after being modulated by metadata and missingness signals.\n\nHigher values indicate stronger preservation of image features, while lower values indicate stronger suppression of image features based on the learned missingness representation.\n\nBecause this plot is constructed from OOF predictions across all cross-validation folds, it reflects the **true model behavior on unseen data** and provides a reliable estimate of how the gating mechanism behaves across the full dataset.","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nplt.hist(\n    OOF_GATES,\n    bins=30,\n    edgecolor=\"black\"\n)\n\nplt.xlabel(\"Mean Gate Activation\")\nplt.ylabel(\"Number of Samples\")\nplt.title(\"Distribution of Missingness-Aware Gate Activations\")\n\n\nplt.grid(alpha=0.3)\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingDistribution-LevelGateAnalysis.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:35.974279Z","iopub.execute_input":"2026-08-02T11:36:35.974544Z","iopub.status.idle":"2026-08-02T11:36:36.468207Z","shell.execute_reply.started":"2026-08-02T11:36:35.974515Z","shell.execute_reply":"2026-08-02T11:36:36.467046Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Class-wise Distribution of Missingness-Aware Gate Activations\n\nThis boxplot compares the distribution of missingness-aware gate activation values between benign and melanoma lesions across all out-of-fold predictions.\n\nEach value represents the mean activation of the learned gating mechanism for a single sample. Higher values correspond to greater preservation of image features, while lower values indicate stronger suppression based on metadata-driven missingness signals.\n\nThis analysis evaluates whether the model applies different feature modulation strategies depending on the lesion class. A separation between the two distributions suggests that the gating mechanism is learning class-dependent representations, potentially contributing to improved discrimination between benign and malignant cases.\n\nThis supports the hypothesis that metadata-aware gating is not uniform but adapts based on the underlying clinical characteristics of the input samples.","metadata":{}},{"cell_type":"code","source":"benign_alpha = OOF_GATES[OOF_TARGETS == 0]\nmelanoma_alpha = OOF_GATES[OOF_TARGETS == 1]\n\nprint(len(benign_alpha), len(melanoma_alpha))\n\nplt.figure(figsize=(6,5))\n\nplt.boxplot(\n    [benign_alpha, melanoma_alpha],\n    labels=[\"Benign\", \"Melanoma\"],\n    showfliers=False\n)\n\nplt.ylabel(\"Gate Activation (Alpha)\")\nplt.title(\"Distribution of Missingness-Aware Gate Values by Class\")\n\nplt.grid(alpha=0.5)\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingAlphaDistributionGateAnalysis.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:36.471396Z","iopub.execute_input":"2026-08-02T11:36:36.471893Z","iopub.status.idle":"2026-08-02T11:36:36.821910Z","shell.execute_reply.started":"2026-08-02T11:36:36.471864Z","shell.execute_reply":"2026-08-02T11:36:36.821041Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Missingness-Stratified Gate Analysis\n\nThis analysis compares gate activation values between samples with missing metadata (age, sex, or site missing) and samples with complete metadata using the out-of-fold predictions.\n\nA Mann–Whitney U test is used to check whether the difference between the two groups is statistically significant.\n\nMean gate values are also reported for:\n- Benign vs melanoma cases  \n- Missing vs present metadata cases  \n\nThis helps evaluate whether the model’s gating mechanism adapts based on metadata availability and lesion class.","metadata":{}},{"cell_type":"code","source":"print(\"\\n=== MISSINGNESS STRATIFIED GATE ANALYSIS ===\")\n\nOOF_MISSING = np.zeros_like(OOF_PREDS, dtype=bool)\n\nfor fold in FOLD_RESULTS:\n    idx = FOLD_RESULTS[fold][\"val_idx\"]\n\n    age = FOLD_RESULTS[fold][\"age_missing\"]\n    sex = FOLD_RESULTS[fold][\"sex_missing\"]\n    site = FOLD_RESULTS[fold][\"site_missing\"]\n\n    OOF_MISSING[idx] = (age + sex + site) > 0\n\nmissing_alpha = OOF_GATES[OOF_MISSING]\npresent_alpha = OOF_GATES[~OOF_MISSING]\n\nstat, p = mannwhitneyu(\n    missing_alpha,\n    present_alpha,\n    alternative=\"two-sided\"\n)\n\nprint(f\"Missing vs Present p = {p:.6e}\")\n\nprint(\"Benign mean gate :\", benign_alpha.mean())\nprint(\"Melanoma mean gate :\", melanoma_alpha.mean())\n\nprint(\"Missing mean gate :\", missing_alpha.mean())\nprint(\"Present mean gate :\", present_alpha.mean())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:36.823031Z","iopub.execute_input":"2026-08-02T11:36:36.823398Z","iopub.status.idle":"2026-08-02T11:36:36.840270Z","shell.execute_reply.started":"2026-08-02T11:36:36.823371Z","shell.execute_reply":"2026-08-02T11:36:36.839070Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Class-wise Gate Comparison (Statistical Test)\n\nThis analysis compares missingness-aware gate activation values between benign and melanoma lesions using a Mann–Whitney U test.\n\nThe goal is to determine whether the gating mechanism behaves differently across the two clinical classes.\n\nA significant p-value suggests that gate activations differ between benign and melanoma samples, indicating that the model may apply different levels of feature modulation depending on lesion type.","metadata":{}},{"cell_type":"code","source":"from scipy.stats import mannwhitneyu\n\n# Gate values by class\nbenign_alpha = OOF_GATES[OOF_TARGETS == 0]\nmelanoma_alpha = OOF_GATES[OOF_TARGETS == 1]\n\nprint(\"=== STATISTICAL TESTS ===\")\n\nstat, p = mannwhitneyu(\n    benign_alpha,\n    melanoma_alpha,\n    alternative=\"two-sided\"\n)\n\nprint(f\"Mann-Whitney U statistic: {stat:.2f}\")\nprint(f\"P-value: {p:.6e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:36.841538Z","iopub.execute_input":"2026-08-02T11:36:36.841836Z","iopub.status.idle":"2026-08-02T11:36:36.862118Z","shell.execute_reply.started":"2026-08-02T11:36:36.841808Z","shell.execute_reply":"2026-08-02T11:36:36.861279Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Final Gate Summary","metadata":{}},{"cell_type":"code","source":"print(\"\\n==============================\")\nprint(\"GATE ANALYSIS SUMMARY\")\nprint(\"==============================\")\n\nif len(OOF_GATES) > 0:\n\n    print(f\"Mean gate: {OOF_GATES.mean():.4f}\")\n    print(f\"Std gate : {OOF_GATES.std():.4f}\")\n\n    print(f\"Image Dominant (Gate > 0.5): {(OOF_GATES > 0.5).mean():.4f}\")\n    print(f\"Metadata Dominant (Gate < 0.5): {(OOF_GATES < 0.5).mean():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:36.863435Z","iopub.execute_input":"2026-08-02T11:36:36.864041Z","iopub.status.idle":"2026-08-02T11:36:36.871596Z","shell.execute_reply.started":"2026-08-02T11:36:36.864009Z","shell.execute_reply":"2026-08-02T11:36:36.870513Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Threshold Search","metadata":{}},{"cell_type":"code","source":"thresholds = np.arange(\n    0.05,\n    0.96,\n    0.01\n)\n\nscores = []\n\nfor threshold in thresholds:\n\n    preds = (\n        OOF_PREDS > threshold\n    ).astype(int)\n\n    f1 = f1_score(\n        OOF_TARGETS,\n        preds\n    )\n\n    scores.append(f1)\n\nbest_idx = np.argmax(scores)\n\nbest_threshold = thresholds[\n    best_idx\n]\n\nbest_f1 = scores[\n    best_idx\n]\n\nplt.figure(figsize=(8,5))\n\nplt.plot(\n    thresholds,\n    scores\n)\n\nplt.xlabel(\n    \"Threshold\"\n)\n\nplt.ylabel(\n    \"F1 Score\"\n)\n\nplt.title(\n    \"Threshold Search\"\n)\n\nplt.grid()\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingThresholdSearch.png\", dpi=300, bbox_inches=\"tight\")\n\nplt.show()\n\nprint(\n    f\"Best Threshold = {best_threshold:.2f}\"\n)\n\nprint(\n    f\"Best F1 = {best_f1:.4f}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:36.872791Z","iopub.execute_input":"2026-08-02T11:36:36.873187Z","iopub.status.idle":"2026-08-02T11:36:38.096214Z","shell.execute_reply.started":"2026-08-02T11:36:36.873136Z","shell.execute_reply":"2026-08-02T11:36:38.095258Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Final Metrics at Best Threshold","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score,\n    confusion_matrix\n)\n\n# Overall OOF ROC-AUC\nauc = roc_auc_score(\n    OOF_TARGETS,\n    OOF_PREDS\n)\n\n# Overall OOF PR-AUC\npr_auc = average_precision_score(\n    OOF_TARGETS,\n    OOF_PREDS\n)\n\nbinary_preds = (OOF_PREDS > best_threshold).astype(int)\n\naccuracy = accuracy_score(\n    OOF_TARGETS,\n    binary_preds\n)\n\nprecision = precision_score(\n    OOF_TARGETS,\n    binary_preds,\n    zero_division=0\n)\n\nrecall = recall_score(\n    OOF_TARGETS,\n    binary_preds,\n    zero_division=0\n)\n\nf1 = f1_score(\n    OOF_TARGETS,\n    binary_preds,\n    zero_division=0\n)\n\ncm = confusion_matrix(\n    OOF_TARGETS,\n    binary_preds\n)\n\ntn, fp, fn, tp = cm.ravel()\n\nspecificity = tn / (tn + fp)\n\nprint(f\"AUC         : {auc:.4f}\")\nprint(f\"PR AUC      : {pr_auc:.4f}\")\nprint(f\"Accuracy    : {accuracy:.4f}\")\nprint(f\"Precision   : {precision:.4f}\")\nprint(f\"Recall      : {recall:.4f}\")\nprint(f\"F1          : {f1:.4f}\")\nprint(f\"Specificity : {specificity:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:38.097233Z","iopub.execute_input":"2026-08-02T11:36:38.097543Z","iopub.status.idle":"2026-08-02T11:36:38.202575Z","shell.execute_reply.started":"2026-08-02T11:36:38.097508Z","shell.execute_reply":"2026-08-02T11:36:38.201670Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\n    \"\\n\"\n    + \"=\"*20\n    + \" FINAL CV RESULTS \"\n    + \"=\"*20\n)\n\nfor fold in sorted(FOLD_RESULTS.keys()):\n\n    print(\n        f\"Fold {fold}: \"\n        f\"AUC={FOLD_RESULTS[fold]['best_auc']:.4f}\"\n    )\n\nauc_scores = [\n    FOLD_RESULTS[f][\"best_auc\"]\n    for f in sorted(FOLD_RESULTS.keys())\n]\n\npr_auc_scores = [\n    max(FOLD_RESULTS[f][\"history\"][\"pr_auc\"])\n    for f in sorted(FOLD_RESULTS.keys())\n]\n\nprint(\"\\n===== AUC =====\")\n\nprint(\n    f\"Mean AUC : {np.mean(auc_scores):.4f}\"\n)\n\nprint(\n    f\"Std AUC  : {np.std(auc_scores):.4f}\"\n)\n\nprint(\"\\n===== PR AUC =====\")\n\nprint(\n    f\"Mean PR AUC : {np.mean(pr_auc_scores):.4f}\"\n)\n\nprint(\n    f\"Std PR AUC  : {np.std(pr_auc_scores):.4f}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:38.203715Z","iopub.execute_input":"2026-08-02T11:36:38.204130Z","iopub.status.idle":"2026-08-02T11:36:38.212625Z","shell.execute_reply.started":"2026-08-02T11:36:38.204092Z","shell.execute_reply":"2026-08-02T11:36:38.211773Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Confusion Matrix","metadata":{}},{"cell_type":"code","source":"disp = ConfusionMatrixDisplay(\n    confusion_matrix=cm\n)\n\ndisp.plot()\n\nplt.title(\n    f\"Confusion Matrix @ {best_threshold:.2f}\"\n)\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingConfusionMatrix.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:38.213666Z","iopub.execute_input":"2026-08-02T11:36:38.213925Z","iopub.status.idle":"2026-08-02T11:36:38.720266Z","shell.execute_reply.started":"2026-08-02T11:36:38.213902Z","shell.execute_reply":"2026-08-02T11:36:38.719145Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Epoch Metric Curves","metadata":{}},{"cell_type":"code","source":"fold_histories = [\n\n    FOLD_RESULTS[f][\"history\"]\n\n    for f in sorted(\n        FOLD_RESULTS.keys()\n    )\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:38.721594Z","iopub.execute_input":"2026-08-02T11:36:38.721923Z","iopub.status.idle":"2026-08-02T11:36:38.727310Z","shell.execute_reply.started":"2026-08-02T11:36:38.721887Z","shell.execute_reply":"2026-08-02T11:36:38.726383Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train Loss","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nfor h in fold_histories:\n\n    plt.plot(\n        h[\"train_loss\"],\n        alpha=0.4\n    )\n\nmean_curve = np.mean(\n\n    [h[\"train_loss\"]\n     for h in fold_histories],\n\n    axis=0\n)\n\nplt.plot(\n    mean_curve,\n    linewidth=3\n)\n\nplt.title(\n    \"Train Loss\"\n)\n\nplt.xlabel(\n    \"Epoch\"\n)\n\nplt.ylabel(\n    \"Loss\"\n)\n\nplt.grid()\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingPer-EpochTrainLoss.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:38.728506Z","iopub.execute_input":"2026-08-02T11:36:38.729070Z","iopub.status.idle":"2026-08-02T11:36:39.252452Z","shell.execute_reply.started":"2026-08-02T11:36:38.729037Z","shell.execute_reply":"2026-08-02T11:36:39.251615Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## AUC","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nfor fold, h in enumerate(fold_histories):\n\n    plt.plot(\n        h[\"auc\"],\n        alpha=0.6,\n        label=f\"Fold {fold}\"\n    )\n\nmean_curve = np.mean(\n    [h[\"auc\"] for h in fold_histories],\n    axis=0\n)\n\nplt.plot(\n    mean_curve,\n    linewidth=3,\n    label=\"Mean\"\n)\n\nplt.title(\"Validation AUC\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"AUC\")\nplt.legend()\nplt.grid()\n\n# SAVE FIGURE\nplt.savefig(\"MissingnessAwareGatingPer-EpochValidationAUC.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:39.253565Z","iopub.execute_input":"2026-08-02T11:36:39.253990Z","iopub.status.idle":"2026-08-02T11:36:39.825814Z","shell.execute_reply.started":"2026-08-02T11:36:39.253951Z","shell.execute_reply":"2026-08-02T11:36:39.824523Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## PR-AUC","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nfor fold, h in enumerate(fold_histories):\n\n    plt.plot(\n        h[\"pr_auc\"],\n        alpha=0.6,\n        label=f\"Fold {fold}\"\n    )\n\nmean_curve = np.mean(\n    [h[\"pr_auc\"] for h in fold_histories],\n    axis=0\n)\n\nplt.plot(\n    mean_curve,\n    linewidth=3,\n    label=\"Mean\"\n)\n\nplt.title(\"Validation PR-AUC\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"PR-AUC\")\nplt.legend()\nplt.grid()\n\nplt.savefig(\"MissingnessAwareGatingPer-EpochValidationPR-AUC.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:39.827048Z","iopub.execute_input":"2026-08-02T11:36:39.827437Z","iopub.status.idle":"2026-08-02T11:36:40.447551Z","shell.execute_reply.started":"2026-08-02T11:36:39.827402Z","shell.execute_reply":"2026-08-02T11:36:40.446769Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## F1","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nfor fold, h in enumerate(fold_histories):\n\n    plt.plot(\n        h[\"f1\"],\n        alpha=0.6,\n        label=f\"Fold {fold}\"\n    )\n\nmean_curve = np.mean(\n    [h[\"f1\"] for h in fold_histories],\n    axis=0\n)\n\nplt.plot(\n    mean_curve,\n    linewidth=3,\n    label=\"Mean\"\n)\n\nplt.title(\"Validation F1\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"F1\")\nplt.legend()\nplt.grid()\n\nplt.savefig(\"MissingnessAwareGatingPer-EpochValidationF1.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:40.448592Z","iopub.execute_input":"2026-08-02T11:36:40.448898Z","iopub.status.idle":"2026-08-02T11:36:41.027962Z","shell.execute_reply.started":"2026-08-02T11:36:40.448863Z","shell.execute_reply":"2026-08-02T11:36:41.026920Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Recall","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nfor fold, h in enumerate(fold_histories):\n\n    plt.plot(\n        h[\"recall\"],\n        alpha=0.6,\n        label=f\"Fold {fold}\"\n    )\n\nmean_curve = np.mean(\n    [h[\"recall\"] for h in fold_histories],\n    axis=0\n)\n\nplt.plot(\n    mean_curve,\n    linewidth=3,\n    label=\"Mean\"\n)\n\nplt.title(\"Validation Recall\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Recall\")\nplt.legend()\nplt.grid()\n\nplt.savefig(\"MissingnessAwareGatingPer-EpochValidationRecall.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:41.029097Z","iopub.execute_input":"2026-08-02T11:36:41.029410Z","iopub.status.idle":"2026-08-02T11:36:41.636017Z","shell.execute_reply.started":"2026-08-02T11:36:41.029375Z","shell.execute_reply":"2026-08-02T11:36:41.634998Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Specificity","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\n\nfor fold, h in enumerate(fold_histories):\n\n    plt.plot(\n        h[\"specificity\"],\n        alpha=0.6,\n        label=f\"Fold {fold}\"\n    )\n\nmean_curve = np.mean(\n    [h[\"specificity\"] for h in fold_histories],\n    axis=0\n)\n\nplt.plot(\n    mean_curve,\n    linewidth=3,\n    label=\"Mean\"\n)\n\nplt.title(\"Validation Specificity\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Specificity\")\nplt.legend()\nplt.grid()\n\nplt.savefig(\"MissingnessAwareGatingPer-EpochValidationSpecificity.png\", dpi=300, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-02T11:36:41.637214Z","iopub.execute_input":"2026-08-02T11:36:41.637576Z","iopub.status.idle":"2026-08-02T11:36:42.279152Z","shell.execute_reply.started":"2026-08-02T11:36:41.637538Z","shell.execute_reply":"2026-08-02T11:36:42.278096Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}