{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":29653,"databundleVersionId":2420395,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"> ## ⚠️ Remarque importante concernant l'exécution du notebook\n>\n> Ce notebook met en œuvre un pipeline de Deep Learning relativement exigeant en mémoire, notamment en raison de :\n>\n> - l'utilisation d'images IRM haute résolution (`224 × 224`) ;\n> - la représentation de chaque patient par plusieurs coupes (`NUM_SLICES`) ;\n> - l'extraction de caractéristiques avec **EfficientNetB0** appliqué à chaque coupe (`TimeDistributed`) ;\n> - l'entraînement en **validation croisée à 5 plis (5-Fold Cross-Validation)**.\n>\n> En conséquence, sur certaines configurations Kaggle, le notebook peut afficher le message :\n>\n> ```text\n> Your notebook tried to allocate more memory than is available.\n> ```\n>\n> ou redémarrer automatiquement le kernel.\n>\n> **Ce comportement est lié aux limites de mémoire de l'environnement Kaggle et ne signifie pas nécessairement qu'une erreur est présente dans le code.**\n>\n> Pour limiter ce problème, plusieurs optimisations ont été mises en place :\n>\n> - réduction de la taille du batch (`BATCH_SIZE = 2`) ;\n> - nettoyage explicite de la mémoire (`gc.collect()` et `tf.keras.backend.clear_session()`) entre les plis de validation ;\n> - limitation du nombre de couches dégelées lors du fine-tuning ;\n> - mise en cache des données prétraitées afin d'éviter le rechargement des fichiers DICOM.\n>\n> Si le kernel redémarre malgré ces optimisations, il suffit généralement de relancer le notebook. Les fichiers intermédiaires sauvegardés permettent de reprendre l'expérience sans devoir reconstruire l'ensemble du pipeline.","metadata":{}},{"cell_type":"markdown","source":"# 🧠 Prédiction du statut de méthylation de MGMT à partir d’IRM cérébrales multimodales\n\n## Présentation générale\n\nCe notebook présente un pipeline de **Deep Learning multimodal** destiné à prédire le **statut de méthylation du promoteur du gène MGMT** à partir d’images IRM cérébrales issues du dataset **RSNA-MICCAI Brain Tumor Radiogenomic Classification**.\n\nLe problème est formulé comme une tâche de **classification binaire au niveau du patient** :\n\n- **MGMT = 0** : promoteur non méthylé ;\n- **MGMT = 1** : promoteur méthylé.\n\nContrairement à une approche basée sur une seule coupe IRM, ce pipeline exploite plusieurs coupes informatives pour chaque patient et combine plusieurs modalités d’imagerie afin d’obtenir une représentation plus complète des caractéristiques anatomiques et tumorales.\n\nL’objectif est de construire un modèle capable d’extraire automatiquement des caractéristiques pertinentes à partir des images IRM, puis de produire une probabilité de méthylation de MGMT pour chaque patient.\n\n---\n\n# Pipeline global\n\n```text\nVolumes DICOM multimodaux\n          │\n          ▼\nChargement et tri anatomique des coupes\n          │\n          ▼\nNormalisation des intensités IRM\n          │\n          ▼\nSélection automatique des coupes informatives\n          │\n          ▼\nConstruction d’images multimodales\nFLAIR - T1wCE - T2w\n          │\n          ▼\nMise en cache des tenseurs en uint8\n          │\n          ▼\nExtraction des caractéristiques avec EfficientNetB0\n          │\n          ▼\nAgrégation des coupes par mécanisme d’attention\n          │\n          ▼\nClassification binaire du statut MGMT\n```\n\n---\n\n# 1. Chargement des données DICOM\n\nChaque patient contient plusieurs séries d’images IRM au format DICOM.\n\nLes trois modalités utilisées dans ce notebook sont :\n\n- **FLAIR** ;\n- **T1wCE** ;\n- **T2w**.\n\nLes fichiers DICOM sont chargés et triés selon leur position anatomique à l’aide des métadonnées disponibles :\n\n- `ImagePositionPatient` ;\n- `SliceLocation` ;\n- `InstanceNumber`.\n\nLes valeurs des pixels sont également corrigées à partir des paramètres DICOM :\n\n- `RescaleSlope` ;\n- `RescaleIntercept`.\n\nCette étape permet de reconstruire correctement chaque volume IRM en respectant l’ordre anatomique des coupes.\n\n---\n\n# 2. Prétraitement et normalisation des volumes IRM\n\nLes intensités IRM ne sont pas standardisées et peuvent varier fortement d’un patient à l’autre, notamment en fonction du scanner et du protocole d’acquisition.\n\nUne normalisation est donc appliquée séparément à chaque modalité.\n\nLe prétraitement comprend :\n\n- la création d’un masque approximatif du cerveau ;\n- la suppression des valeurs extrêmes par seuillage percentile ;\n- une normalisation par **Z-score** ;\n- la conservation d’un arrière-plan égal à zéro.\n\nCette étape permet de réduire les variations liées à l’acquisition et d’améliorer la comparabilité des patients.\n\n---\n\n# 3. Sélection automatique des coupes informatives\n\nLes différentes modalités ne possèdent pas toujours le même nombre de coupes.\n\nLa modalité **FLAIR** est utilisée comme volume de référence. Les positions correspondantes dans T1wCE et T2w sont estimées à partir de leur position relative dans le volume.\n\nChaque coupe reçoit ensuite un **score d’information** calculé à partir de plusieurs critères :\n\n- proportion de tissu cérébral ;\n- présence de régions de forte intensité ;\n- variation des intensités ;\n- richesse des contours et des textures.\n\nLes `NUM_SLICES` coupes ayant les scores les plus élevés sont retenues.\n\nCette stratégie réduit le nombre d’images traitées par le modèle tout en privilégiant les zones potentiellement les plus pertinentes.\n\n---\n\n# 4. Construction des images multimodales\n\nPour chaque position anatomique sélectionnée, une image à trois canaux est construite :\n\n| Canal | Modalité IRM |\n|---|---|\n| Rouge | FLAIR |\n| Vert | T1wCE |\n| Bleu | T2w |\n\nChaque coupe est redimensionnée en :\n\n```text\n224 × 224 pixels\n```\n\nLe tenseur final représentant un patient possède donc la forme :\n\n```text\n(NUM_SLICES, 224, 224, 3)\n```\n\nCette représentation est compatible avec un réseau convolutionnel pré-entraîné sur ImageNet.\n\n---\n\n# 5. Mise en cache des tenseurs patients\n\nAprès le prétraitement, chaque tenseur patient est sauvegardé dans un fichier NumPy au format :\n\n```text\n.npy\n```\n\nLes tenseurs sont stockés en type :\n\n```text\nuint8\n```\n\nafin de réduire l’espace disque et la consommation de mémoire.\n\nIls sont convertis en `float32` uniquement lorsqu’ils sont chargés pour l’entraînement.\n\nLa mise en cache permet :\n\n- d’éviter de relire les fichiers DICOM à chaque époque ;\n- d’accélérer considérablement l’entraînement ;\n- de réduire la charge du processeur ;\n- de limiter l’utilisation de la RAM ;\n- d’identifier les patients invalides ou incomplets.\n\n---\n\n# 6. Augmentation des données\n\nAfin de réduire le surapprentissage et d’améliorer la capacité de généralisation du modèle, plusieurs transformations légères sont appliquées pendant l’entraînement :\n\n- retournement horizontal ;\n- variation modérée de la luminosité ;\n- variation modérée du contraste ;\n- suppression aléatoire d’une modalité.\n\nLa même transformation géométrique est appliquée à toutes les coupes d’un même patient afin de conserver leur cohérence spatiale.\n\nLes valeurs des images restent comprises entre :\n\n```text\n0 et 255\n```\n\nce qui est compatible avec le prétraitement interne d’EfficientNetB0.\n\n---\n\n# 7. Extraction des caractéristiques avec EfficientNetB0\n\nChaque coupe multimodale est traitée par un réseau **EfficientNetB0** pré-entraîné sur ImageNet.\n\nConfiguration utilisée :\n\n```python\ninclude_top=False\nweights=\"imagenet\"\npooling=\"avg\"\n```\n\nLe réseau n’est pas utilisé directement comme classifieur. Il joue le rôle d’**extracteur de caractéristiques**.\n\nPour chaque coupe, EfficientNetB0 produit un vecteur décrivant notamment :\n\n- les contours ;\n- les textures ;\n- les contrastes ;\n- les formes ;\n- les motifs visuels locaux.\n\nLe même réseau est appliqué à toutes les coupes grâce à la couche :\n\n```python\nTimeDistributed()\n```\n\nAinsi, toutes les coupes sont analysées par le même encodeur et partagent les mêmes poids.\n\n---\n\n# 8. Projection des caractéristiques\n\nLes caractéristiques extraites par EfficientNet sont ensuite projetées dans un espace de dimension plus faible grâce à une couche dense.\n\nCette étape permet :\n\n- de réduire la dimension des vecteurs ;\n- de limiter le nombre de paramètres ;\n- de faciliter l’agrégation des coupes ;\n- de réduire le risque de surapprentissage.\n\nUne régularisation par `Dropout` est également appliquée.\n\n---\n\n# 9. Agrégation par mécanisme d’attention\n\nToutes les coupes ne contiennent pas la même quantité d’information.\n\nCertaines peuvent montrer clairement la tumeur, alors que d’autres contiennent principalement du cerveau sain.\n\nUn mécanisme d’**attention** apprend automatiquement à attribuer un poids différent à chaque coupe.\n\nExemple conceptuel :\n\n```text\nCoupe 1  → poids faible\nCoupe 2  → poids faible\nCoupe 3  → poids élevé\nCoupe 4  → poids très élevé\n...\n```\n\nLes vecteurs de caractéristiques sont ensuite combinés selon ces poids afin d’obtenir une représentation unique du patient.\n\nCette représentation résume les informations les plus utiles provenant des différentes coupes IRM.\n\n---\n\n# 10. Classification du statut MGMT\n\nLa représentation patient obtenue après l’attention est envoyée vers des couches denses de classification.\n\nLa dernière couche utilise une fonction d’activation :\n\n```python\nsigmoid\n```\n\nElle produit une probabilité comprise entre 0 et 1.\n\nExemple :\n\n```text\n0.85 → forte probabilité de MGMT méthylé\n0.20 → faible probabilité de MGMT méthylé\n```\n\nLa prédiction finale est obtenue en comparant cette probabilité à un seuil de classification.\n\n---\n\n# 11. Stratégie d’entraînement en deux phases\n\nL’apprentissage est réalisé en deux étapes.\n\n## Phase 1 — Backbone gelé\n\nDans un premier temps, le backbone EfficientNetB0 est gelé.\n\nSeules les couches ajoutées sont entraînées :\n\n- projection des caractéristiques ;\n- mécanisme d’attention ;\n- couches denses ;\n- couche de classification.\n\nCette phase permet au classifieur d’apprendre à utiliser les caractéristiques pré-entraînées sans modifier immédiatement EfficientNet.\n\n## Phase 2 — Fine-tuning\n\nDans un second temps, les dernières couches d’EfficientNet sont dégelées.\n\nUn faible taux d’apprentissage est utilisé afin d’adapter progressivement les représentations apprises sur ImageNet aux images IRM cérébrales.\n\nLes couches de type **Batch Normalization** restent gelées afin d’améliorer la stabilité avec une petite taille de batch.\n\n---\n\n# 12. Validation croisée stratifiée\n\nLes performances sont estimées avec une **validation croisée stratifiée à `N_FOLDS` plis**.\n\nÀ chaque pli :\n\n- une partie des patients est utilisée pour l’entraînement ;\n- une autre partie est utilisée pour la validation ;\n- la proportion des deux classes est préservée ;\n- aucun patient n’est présent simultanément dans les deux ensembles.\n\nChaque pli est entraîné séparément afin de réduire la consommation mémoire et de sécuriser les résultats.\n\nAprès chaque pli, les éléments suivants sont sauvegardés :\n\n- poids du modèle ;\n- probabilités prédites ;\n- métriques du pli ;\n- identifiants des patients de validation.\n\n---\n\n# 13. Prédictions Out-of-Fold\n\nLes prédictions des différents plis sont fusionnées afin de produire des prédictions dites :\n\n```text\nOut-of-Fold\n```\n\nChaque patient reçoit une prédiction produite par un modèle qui ne l’a jamais utilisé pendant l’entraînement.\n\nCette stratégie fournit une estimation plus fiable de la capacité de généralisation du modèle qu’une seule séparation entraînement-validation.\n\n---\n\n# 14. Évaluation finale\n\nLes performances du modèle sont évaluées à l’aide des métriques suivantes :\n\n- ROC-AUC ;\n- PR-AUC ;\n- Accuracy ;\n- Balanced Accuracy ;\n- F1-score ;\n- Matthews Correlation Coefficient ;\n- Sensibilité ;\n- Spécificité ;\n- Précision ;\n- Valeur prédictive négative.\n\nLes visualisations suivantes sont également produites :\n\n- matrice de confusion ;\n- courbe ROC ;\n- courbe Precision-Recall ;\n- comparaison des ROC-AUC entre les plis.\n\nLa matrice de confusion finale est calculée à partir des prédictions Out-of-Fold.\n\n---\n\n# Résumé de l’architecture\n\n```text\nPatient\n  │\n  ▼\nVolumes DICOM FLAIR, T1wCE et T2w\n  │\n  ▼\nTri anatomique et normalisation\n  │\n  ▼\nSélection de NUM_SLICES coupes informatives\n  │\n  ▼\nConstruction d’images multimodales 224 × 224 × 3\n  │\n  ▼\nEfficientNetB0 appliqué à chaque coupe\n  │\n  ▼\nVecteurs de caractéristiques par coupe\n  │\n  ▼\nProjection des caractéristiques\n  │\n  ▼\nMécanisme d’attention\n  │\n  ▼\nReprésentation unique du patient\n  │\n  ▼\nDense + Sigmoid\n  │\n  ▼\nProbabilité de méthylation de MGMT\n```\n\n---\n\n# Gestion de la mémoire\n\nCe pipeline est relativement exigeant en mémoire en raison du nombre de coupes et de l’utilisation d’un réseau convolutionnel pré-entraîné.\n\nPlusieurs optimisations sont appliquées :\n\n- cache patient enregistré en `uint8` ;\n- conversion en `float32` uniquement au chargement ;\n- `BATCH_SIZE` réduit ;\n- préchargement limité avec `prefetch(1)` ;\n- traitement séquentiel des patients ;\n- validation croisée exécutée un pli à la fois ;\n- nettoyage explicite de la mémoire avec `gc.collect()` ;\n- nettoyage de TensorFlow avec `tf.keras.backend.clear_session()`.\n\nCes optimisations rendent le notebook plus stable dans l’environnement Kaggle.\n\n---\n\n# Limites actuelles\n\nCette approche présente plusieurs limites :\n\n- les modalités IRM ne sont pas recalées précisément par une méthode de registration médicale ;\n- la sélection des coupes est fondée sur un score d’information et non sur une segmentation exacte de la tumeur ;\n- EfficientNetB0 a été pré-entraîné sur des images naturelles et non sur des IRM ;\n- le nombre de patients reste relativement limité pour un modèle profond ;\n- la relation entre l’apparence radiologique et la méthylation de MGMT reste faible et difficile à apprendre.\n\n---\n\n# Perspectives d’amélioration\n\nPlusieurs améliorations peuvent être envisagées :\n\n- registration précise entre les modalités ;\n- segmentation automatique de la tumeur ;\n- sélection des coupes selon la surface tumorale ;\n- recadrage autour de la tumeur ;\n- utilisation de modèles pré-entraînés sur des images médicales ;\n- extraction de caractéristiques radiomiques ;\n- fusion de caractéristiques profondes et radiomiques ;\n- intégration de variables cliniques ou génomiques ;\n- pré-entraînement auto-supervisé sur de grandes bases IRM.\n\nUne amélioration particulièrement importante serait l’intégration d’un modèle de segmentation pré-entraîné sur **BraTS** afin de localiser automatiquement la tumeur et de concentrer l’apprentissage sur les régions tumorales et péritumorales.","metadata":{}},{"cell_type":"code","source":"\n# ============================================================\n# 1. Imports and reproducibility\n# ============================================================\n\nimport os\nimport gc\nimport re\nimport random\nimport warnings\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import (\n    roc_auc_score,\n    roc_curve,\n    confusion_matrix,\n    classification_report,\n    accuracy_score,\n    balanced_accuracy_score,\n    f1_score,\n    matthews_corrcoef,\n)\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers, Model\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras.callbacks import (\n    EarlyStopping,\n    ReduceLROnPlateau,\n    ModelCheckpoint,\n)\nfrom tensorflow.keras.optimizers import Adam\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\n\ndef seed_everything(seed=SEED):\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n\nseed_everything(SEED)\n\nprint(\"TensorFlow:\", tf.__version__)\nprint(\"GPU devices:\", tf.config.list_physical_devices(\"GPU\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:05:50.475379Z","iopub.execute_input":"2026-07-11T17:05:50.475840Z","iopub.status.idle":"2026-07-11T17:06:04.771063Z","shell.execute_reply.started":"2026-07-11T17:05:50.475809Z","shell.execute_reply":"2026-07-11T17:06:04.770202Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ============================================================\n# 2. Configuration\n# ============================================================\n\nDATA_ROOT = Path ('/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification/')\n\nTRAIN_DIR = DATA_ROOT / \"train\"\nLABELS_PATH = DATA_ROOT / \"train_labels.csv\"\n\nMODALITIES = [\"FLAIR\", \"T1wCE\", \"T2w\"]\n\nIMG_SIZE = 224\nNUM_SLICES = 12\nBATCH_SIZE = 4\nEPOCHS_HEAD = 8\nEPOCHS_FINE = 20\nN_FOLDS = 5\n\nLEARNING_RATE_HEAD = 1e-3\nLEARNING_RATE_FINE = 1e-5\n\nCACHE_DIR = Path(\"/kaggle/working/rsna_cache\")\nCACHE_DIR.mkdir(parents=True, exist_ok=True)\n\nlabels_df = pd.read_csv(LABELS_PATH)\nlabels_df[\"BraTS21ID\"] = labels_df[\"BraTS21ID\"].astype(int)\n\nprint(labels_df.head())\nprint(\"\\nClass distribution:\")\nprint(labels_df[\"MGMT_value\"].value_counts().sort_index())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:04.772290Z","iopub.execute_input":"2026-07-11T17:06:04.772806Z","iopub.status.idle":"2026-07-11T17:06:04.805450Z","shell.execute_reply.started":"2026-07-11T17:06:04.772786Z","shell.execute_reply":"2026-07-11T17:06:04.804612Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 3. DICOM loading utilities\n# ============================================================\n\ndef natural_key(path):\n    numbers = re.findall(r\"\\d+\", Path(path).stem)\n    return int(numbers[-1]) if numbers else 0\n\n\ndef dicom_position(ds):\n    try:\n        return float(ds.ImagePositionPatient[2])\n    except Exception:\n        try:\n            return float(ds.SliceLocation)\n        except Exception:\n            try:\n                return float(ds.InstanceNumber)\n            except Exception:\n                return 0.0\n\n\ndef read_dicom_slice(path):\n    ds = pydicom.dcmread(\n        str(path),\n        force=True\n    )\n\n    image = ds.pixel_array.astype(np.float32)\n\n    slope = float(\n        getattr(ds, \"RescaleSlope\", 1.0)\n    )\n\n    intercept = float(\n        getattr(ds, \"RescaleIntercept\", 0.0)\n    )\n\n    image = image * slope + intercept\n\n    return ds, image\n\n\ndef load_dicom_volume(folder):\n    folder = Path(folder)\n\n    files = list(folder.glob(\"*.dcm\"))\n\n    if not files:\n        return None\n\n    records = []\n\n    for path in files:\n        try:\n            ds = pydicom.dcmread(\n                str(path),\n                stop_before_pixels=True,\n                force=True\n            )\n\n            records.append(\n                (dicom_position(ds), path)\n            )\n\n        except Exception:\n            records.append(\n                (natural_key(path), path)\n            )\n\n    records = sorted(\n        records,\n        key=lambda x: x[0]\n    )\n\n    slices = []\n\n    for _, path in records:\n        try:\n            _, image = read_dicom_slice(path)\n\n            if image.ndim == 2:\n                slices.append(image)\n\n        except Exception:\n            continue\n\n    if not slices:\n        return None\n\n    target_h = min(\n        image.shape[0]\n        for image in slices\n    )\n\n    target_w = min(\n        image.shape[1]\n        for image in slices\n    )\n\n    resized = []\n\n    for image in slices:\n        if image.shape != (target_h, target_w):\n            image = cv2.resize(\n                image,\n                (target_w, target_h),\n                interpolation=cv2.INTER_LINEAR\n            )\n\n        resized.append(image)\n\n    volume = np.stack(\n        resized,\n        axis=0\n    ).astype(np.float32)\n\n    return volume","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:04.806236Z","iopub.execute_input":"2026-07-11T17:06:04.806554Z","iopub.status.idle":"2026-07-11T17:06:04.820226Z","shell.execute_reply.started":"2026-07-11T17:06:04.806511Z","shell.execute_reply":"2026-07-11T17:06:04.819304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ============================================================\n# 4. MRI normalization and brain cropping\n# ============================================================\n\ndef normalize_mri(volume):\n    volume = volume.astype(np.float32)\n\n    # Brain mask\n    mask = volume > np.percentile(volume, 2)\n\n    if np.sum(mask) == 0:\n        return np.zeros_like(volume, dtype=np.float32)\n\n    brain = volume[mask]\n\n    # Remove outliers\n    low = np.percentile(brain, 1)\n    high = np.percentile(brain, 99)\n\n    volume = np.clip(volume, low, high)\n\n    # Recompute brain after clipping\n    brain = volume[mask]\n\n    mean = brain.mean()\n    std = brain.std()\n\n    if std < 1e-6:\n        std = 1.0\n\n    # Normalize only brain voxels\n    volume[mask] = (volume[mask] - mean) / std\n\n    # Background stays zero\n    volume[~mask] = 0.0\n\n    return volume.astype(np.float32)\n\ndef find_brain_bbox(volumes, margin=10):\n    combined = np.zeros_like(volumes[0], dtype=np.float32)\n\n    for volume in volumes:\n        combined += (np.abs(volume) > 1e-6).astype(np.float32)\n\n    brain_mask = np.sum(combined > 0, axis=0)\n\n    mask = brain_mask > (combined.shape[0] * 0.05)\n\n    coords = np.argwhere(mask)\n\n    if coords.size == 0:\n        h, w = projection.shape\n        return 0, h, 0, w\n\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n\n    y0 = max(0, y0 - margin)\n    x0 = max(0, x0 - margin)\n    y1 = min(projection.shape[0], y1 + margin)\n    x1 = min(projection.shape[1], x1 + margin)\n\n    return y0, y1, x0, x1\n\n\ndef crop_volume(volume, bbox):\n    y0, y1, x0, x1 = bbox\n    return volume[:, y0:y1, x0:x1]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:04.821978Z","iopub.execute_input":"2026-07-11T17:06:04.822334Z","iopub.status.idle":"2026-07-11T17:06:04.835044Z","shell.execute_reply.started":"2026-07-11T17:06:04.822315Z","shell.execute_reply":"2026-07-11T17:06:04.834392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 5. Slice alignment and information-based selection\n#    Memory-efficient version\n# ============================================================\n\nimport numpy as np\nimport cv2\n\n\ndef map_slice_index(\n    reference_index,\n    reference_depth,\n    target_depth,\n):\n    \"\"\"\n    Map one slice position from a reference volume to the\n    corresponding relative position in another volume.\n    \"\"\"\n\n    if reference_depth <= 1 or target_depth <= 1:\n        return 0\n\n    relative_position = (\n        reference_index / float(reference_depth - 1)\n    )\n\n    target_index = int(\n        round(\n            relative_position\n            * (target_depth - 1)\n        )\n    )\n\n    return int(\n        np.clip(\n            target_index,\n            0,\n            target_depth - 1,\n        )\n    )\n\n\ndef robust_rescale_slice(\n    image,\n    output_size=128,\n):\n    \"\"\"\n    Downsample and rescale one MRI slice to [0, 1]\n    only for information scoring.\n\n    The final model images remain 224 x 224.\n    Using a smaller size here reduces RAM and CPU usage.\n    \"\"\"\n\n    image = np.asarray(\n        image,\n        dtype=np.float32,\n    )\n\n    # Downsample only for scoring.\n    if (\n        image.shape[0] != output_size\n        or image.shape[1] != output_size\n    ):\n        image = cv2.resize(\n            image,\n            (output_size, output_size),\n            interpolation=cv2.INTER_AREA,\n        )\n\n    mask = np.abs(image) > 1e-6\n\n    if np.count_nonzero(mask) < 10:\n        return np.zeros(\n            image.shape,\n            dtype=np.float32,\n        )\n\n    values = image[mask]\n\n    low = float(\n        np.percentile(values, 1)\n    )\n\n    high = float(\n        np.percentile(values, 99)\n    )\n\n    if high <= low:\n        return np.zeros(\n            image.shape,\n            dtype=np.float32,\n        )\n\n    # Reuse one output array.\n    output = image.copy()\n\n    np.clip(\n        output,\n        low,\n        high,\n        out=output,\n    )\n\n    output -= low\n    output /= (\n        high - low + 1e-8\n    )\n\n    output[~mask] = 0.0\n\n    return output\n\n\ndef slice_information_score(\n    flair_slice,\n    t1ce_slice,\n    t2_slice,\n):\n    \"\"\"\n    Calculate an approximate information score for one\n    anatomical position.\n\n    This is not a true tumor segmentation score.\n    \"\"\"\n\n    flair = robust_rescale_slice(\n        flair_slice,\n        output_size=128,\n    )\n\n    t1ce = robust_rescale_slice(\n        t1ce_slice,\n        output_size=128,\n    )\n\n    t2 = robust_rescale_slice(\n        t2_slice,\n        output_size=128,\n    )\n\n    brain_mask = (\n        (flair > 0.0)\n        | (t1ce > 0.0)\n        | (t2 > 0.0)\n    )\n\n    brain_count = np.count_nonzero(\n        brain_mask\n    )\n\n    if brain_count < 10:\n        return 0.0\n\n    brain_fraction = (\n        brain_count\n        / float(brain_mask.size)\n    )\n\n    if brain_fraction < 0.05:\n        return 0.0\n\n    flair_values = flair[brain_mask]\n    t1ce_values = t1ce[brain_mask]\n    t2_values = t2[brain_mask]\n\n    if (\n        flair_values.size < 10\n        or t1ce_values.size < 10\n        or t2_values.size < 10\n    ):\n        return 0.0\n\n    flair_threshold = float(\n        np.percentile(\n            flair_values,\n            85,\n        )\n    )\n\n    t1ce_threshold = float(\n        np.percentile(\n            t1ce_values,\n            85,\n        )\n    )\n\n    flair_activity = float(\n        np.mean(\n            flair_values >= flair_threshold\n        )\n    )\n\n    t1ce_activity = float(\n        np.mean(\n            t1ce_values >= t1ce_threshold\n        )\n    )\n\n    variation = float(\n        (\n            np.std(flair_values)\n            + np.std(t1ce_values)\n            + np.std(t2_values)\n        )\n        / 3.0\n    )\n\n    lap_flair = cv2.Laplacian(\n        flair,\n        cv2.CV_32F,\n    )\n\n    lap_t1ce = cv2.Laplacian(\n        t1ce,\n        cv2.CV_32F,\n    )\n\n    edge_score = float(\n        np.log1p(\n            lap_flair.var()\n            + lap_t1ce.var()\n        )\n    )\n\n    edge_score = float(\n        np.clip(\n            edge_score / 5.0,\n            0.0,\n            1.0,\n        )\n    )\n\n    score = (\n        0.25 * brain_fraction\n        + 0.20 * flair_activity\n        + 0.20 * t1ce_activity\n        + 0.25 * variation\n        + 0.10 * edge_score\n    )\n\n    return float(score)\n\n\ndef select_diverse_top_slices(\n    candidate_indices,\n    candidate_scores,\n    num_slices,\n    minimum_distance=3,\n):\n    \"\"\"\n    Select high-scoring slices while avoiding too many\n    neighboring slices.\n    \"\"\"\n\n    order = np.argsort(\n        candidate_scores\n    )[::-1]\n\n    selected = []\n\n    for position in order:\n        slice_index = int(\n            candidate_indices[position]\n        )\n\n        if all(\n            abs(\n                slice_index - previous\n            ) >= minimum_distance\n            for previous in selected\n        ):\n            selected.append(\n                slice_index\n            )\n\n        if len(selected) == num_slices:\n            break\n\n    if len(selected) < num_slices:\n        for position in order:\n            slice_index = int(\n                candidate_indices[position]\n            )\n\n            if slice_index not in selected:\n                selected.append(\n                    slice_index\n                )\n\n            if len(selected) == num_slices:\n                break\n\n    return np.sort(\n        np.asarray(\n            selected,\n            dtype=np.int32,\n        )\n    )\n\n\ndef select_slice_indices(\n    volumes,\n    num_slices=NUM_SLICES,\n):\n    \"\"\"\n    Select informative FLAIR slice positions and map them\n    approximately to T1wCE and T2w.\n\n    Parameters\n    ----------\n    volumes\n        [FLAIR, T1wCE, T2w]\n\n    Returns\n    -------\n    selected_indices\n        Selected indices in the FLAIR reference volume.\n\n    scores\n        Information score for every FLAIR slice.\n    \"\"\"\n\n    if len(volumes) != 3:\n        raise ValueError(\n            \"volumes must contain \"\n            \"[FLAIR, T1wCE, T2w].\"\n        )\n\n    flair, t1ce, t2 = volumes\n\n    if (\n        flair is None\n        or t1ce is None\n        or t2 is None\n    ):\n        raise ValueError(\n            \"One or more MRI volumes are None.\"\n        )\n\n    if (\n        flair.ndim != 3\n        or t1ce.ndim != 3\n        or t2.ndim != 3\n    ):\n        raise ValueError(\n            \"Each MRI volume must have shape \"\n            \"(depth, height, width).\"\n        )\n\n    flair_depth = int(\n        flair.shape[0]\n    )\n\n    t1ce_depth = int(\n        t1ce.shape[0]\n    )\n\n    t2_depth = int(\n        t2.shape[0]\n    )\n\n    if flair_depth == 0:\n        raise ValueError(\n            \"FLAIR volume contains no slices.\"\n        )\n\n    scores = np.zeros(\n        flair_depth,\n        dtype=np.float32,\n    )\n\n    # Ignore extreme superior/inferior regions from the start.\n    central_start = int(\n        round(\n            flair_depth * 0.10\n        )\n    )\n\n    central_end = int(\n        round(\n            flair_depth * 0.90\n        )\n    )\n\n    if central_end <= central_start:\n        candidate_indices = np.arange(\n            flair_depth,\n            dtype=np.int32,\n        )\n    else:\n        candidate_indices = np.arange(\n            central_start,\n            central_end,\n            dtype=np.int32,\n        )\n\n    if len(candidate_indices) < num_slices:\n        candidate_indices = np.arange(\n            flair_depth,\n            dtype=np.int32,\n        )\n\n    # Score only candidate slices, not the complete volume.\n    for flair_index in candidate_indices:\n        t1ce_index = map_slice_index(\n            reference_index=int(\n                flair_index\n            ),\n            reference_depth=flair_depth,\n            target_depth=t1ce_depth,\n        )\n\n        t2_index = map_slice_index(\n            reference_index=int(\n                flair_index\n            ),\n            reference_depth=flair_depth,\n            target_depth=t2_depth,\n        )\n\n        scores[flair_index] = (\n            slice_information_score(\n                flair[int(flair_index)],\n                t1ce[t1ce_index],\n                t2[t2_index],\n            )\n        )\n\n    candidate_scores = scores[\n        candidate_indices\n    ]\n\n    minimum_distance = max(\n        1,\n        flair_depth // 100,\n    )\n\n    selected = select_diverse_top_slices(\n        candidate_indices=candidate_indices,\n        candidate_scores=candidate_scores,\n        num_slices=min(\n            num_slices,\n            len(candidate_indices),\n        ),\n        minimum_distance=minimum_distance,\n    )\n\n    if len(selected) < num_slices:\n        selected = np.linspace(\n            0,\n            flair_depth - 1,\n            num_slices,\n        ).round().astype(\n            np.int32\n        )\n\n    return (\n        selected[:num_slices],\n        scores,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:04.835983Z","iopub.execute_input":"2026-07-11T17:06:04.836303Z","iopub.status.idle":"2026-07-11T17:06:04.858973Z","shell.execute_reply.started":"2026-07-11T17:06:04.836281Z","shell.execute_reply":"2026-07-11T17:06:04.858327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 6. Build one patient tensor\n#    Memory-efficient and uint8 cache version\n# ============================================================\n\ndef patient_folder(patient_id):\n    return TRAIN_DIR / f\"{int(patient_id):05d}\"\n\n\ndef slice_to_255(image):\n    \"\"\"\n    Convert one z-score-normalized MRI slice to [0, 255].\n\n    The returned array is float32 during preprocessing.\n    It will later be saved as uint8 in the cache.\n    \"\"\"\n\n    image = np.asarray(\n        image,\n        dtype=np.float32,\n    )\n\n    mask = np.abs(image) > 1e-6\n\n    if np.count_nonzero(mask) < 10:\n        return np.zeros(\n            image.shape,\n            dtype=np.float32,\n        )\n\n    values = image[mask]\n\n    low = float(\n        np.percentile(values, 1)\n    )\n\n    high = float(\n        np.percentile(values, 99)\n    )\n\n    if high <= low:\n        return np.zeros(\n            image.shape,\n            dtype=np.float32,\n        )\n\n    output = image.copy()\n\n    np.clip(\n        output,\n        low,\n        high,\n        out=output,\n    )\n\n    output -= low\n    output /= (\n        high - low + 1e-8\n    )\n\n    output[~mask] = 0.0\n\n    output *= 255.0\n\n    return output.astype(\n        np.float32,\n        copy=False,\n    )\n\n\ndef build_patient_tensor(\n    patient_id,\n    save_cache=True,\n    verbose=False,\n):\n    \"\"\"\n    Build one patient tensor with shape:\n\n        (NUM_SLICES, IMG_SIZE, IMG_SIZE, 3)\n\n    Channels:\n        0 = FLAIR\n        1 = T1wCE\n        2 = T2w\n\n    Cache format:\n        uint8 on disk to reduce storage and RAM pressure.\n\n    Returned tensor:\n        float32 in [0, 255] for compatibility with EfficientNet.\n    \"\"\"\n\n    patient_id = int(patient_id)\n\n    cache_path = (\n        CACHE_DIR\n        / f\"{patient_id:05d}.npy\"\n    )\n\n    expected_shape = (\n        NUM_SLICES,\n        IMG_SIZE,\n        IMG_SIZE,\n        3,\n    )\n\n    # --------------------------------------------------------\n    # Load existing cache\n    # --------------------------------------------------------\n\n    if cache_path.exists():\n        try:\n            cached_tensor = np.load(\n                cache_path,\n                mmap_mode=\"r\",\n            )\n\n            if cached_tensor.shape == expected_shape:\n                if verbose:\n                    print(\n                        \"Loaded cached tensor:\",\n                        cache_path,\n                    )\n                    print(\n                        \"Cache dtype:\",\n                        cached_tensor.dtype,\n                    )\n\n                # Convert only the current patient to float32.\n                return np.asarray(\n                    cached_tensor,\n                    dtype=np.float32,\n                )\n\n            if verbose:\n                print(\n                    \"Deleting invalid cache:\",\n                    cache_path,\n                    cached_tensor.shape,\n                )\n\n            del cached_tensor\n            cache_path.unlink()\n\n        except Exception as exc:\n            print(\n                f\"Patient {patient_id}: cache error:\",\n                repr(exc),\n            )\n\n            if cache_path.exists():\n                cache_path.unlink()\n\n    # --------------------------------------------------------\n    # Locate patient folder\n    # --------------------------------------------------------\n\n    folder = patient_folder(patient_id)\n\n    if not folder.exists():\n        print(\n            \"Missing patient folder:\",\n            folder,\n        )\n        return None\n\n    if verbose:\n        print(\n            \"Patient folder:\",\n            folder,\n        )\n\n    # --------------------------------------------------------\n    # Required modalities\n    # --------------------------------------------------------\n\n    required_modalities = [\n        \"FLAIR\",\n        \"T1wCE\",\n        \"T2w\",\n    ]\n\n    volumes = {}\n\n    # --------------------------------------------------------\n    # Load and normalize modalities\n    # --------------------------------------------------------\n\n    for modality in required_modalities:\n        modality_path = (\n            folder\n            / modality\n        )\n\n        if verbose:\n            print(\n                \"\\nLoading:\",\n                modality,\n            )\n            print(\n                \"Path:\",\n                modality_path,\n            )\n\n        volume = load_dicom_volume(\n            modality_path\n        )\n\n        if volume is None:\n            print(\n                f\"Patient {patient_id}: \"\n                f\"could not load {modality}\"\n            )\n            return None\n\n        if volume.ndim != 3:\n            print(\n                f\"Patient {patient_id}: invalid volume shape \"\n                f\"for {modality}: {volume.shape}\"\n            )\n            return None\n\n        volume = normalize_mri(\n            volume\n        )\n\n        if not np.isfinite(volume).all():\n            print(\n                f\"Patient {patient_id}: non-finite values \"\n                f\"found in {modality}\"\n            )\n            return None\n\n        volumes[modality] = volume\n\n        if verbose:\n            print(\n                modality,\n                \"normalized shape:\",\n                volume.shape,\n            )\n\n    # --------------------------------------------------------\n    # Retrieve modalities\n    # --------------------------------------------------------\n\n    flair = volumes[\"FLAIR\"]\n    t1ce = volumes[\"T1wCE\"]\n    t2w = volumes[\"T2w\"]\n\n    flair_depth = int(\n        flair.shape[0]\n    )\n\n    t1ce_depth = int(\n        t1ce.shape[0]\n    )\n\n    t2w_depth = int(\n        t2w.shape[0]\n    )\n\n    # --------------------------------------------------------\n    # Select informative FLAIR slices\n    # --------------------------------------------------------\n\n    try:\n        (\n            selected_flair_indices,\n            scores,\n        ) = select_slice_indices(\n            [\n                flair,\n                t1ce,\n                t2w,\n            ],\n            num_slices=NUM_SLICES,\n        )\n\n    except Exception as exc:\n        print(\n            f\"Patient {patient_id}: \"\n            f\"slice selection failed:\",\n            repr(exc),\n        )\n        return None\n\n    if len(selected_flair_indices) != NUM_SLICES:\n        print(\n            f\"Patient {patient_id}: expected \"\n            f\"{NUM_SLICES} selected slices, got \"\n            f\"{len(selected_flair_indices)}\"\n        )\n        return None\n\n    # --------------------------------------------------------\n    # Preallocate final tensor\n    # --------------------------------------------------------\n\n    tensor = np.empty(\n        expected_shape,\n        dtype=np.uint8,\n    )\n\n    # --------------------------------------------------------\n    # Build multimodal slices\n    # --------------------------------------------------------\n\n    for output_index, flair_index in enumerate(\n        selected_flair_indices\n    ):\n        flair_index = int(\n            flair_index\n        )\n\n        t1ce_index = map_slice_index(\n            reference_index=flair_index,\n            reference_depth=flair_depth,\n            target_depth=t1ce_depth,\n        )\n\n        t2w_index = map_slice_index(\n            reference_index=flair_index,\n            reference_depth=flair_depth,\n            target_depth=t2w_depth,\n        )\n\n        flair_slice = cv2.resize(\n            flair[flair_index],\n            (\n                IMG_SIZE,\n                IMG_SIZE,\n            ),\n            interpolation=cv2.INTER_AREA,\n        )\n\n        t1ce_slice = cv2.resize(\n            t1ce[t1ce_index],\n            (\n                IMG_SIZE,\n                IMG_SIZE,\n            ),\n            interpolation=cv2.INTER_AREA,\n        )\n\n        t2w_slice = cv2.resize(\n            t2w[t2w_index],\n            (\n                IMG_SIZE,\n                IMG_SIZE,\n            ),\n            interpolation=cv2.INTER_AREA,\n        )\n\n        flair_slice = slice_to_255(\n            flair_slice\n        )\n\n        t1ce_slice = slice_to_255(\n            t1ce_slice\n        )\n\n        t2w_slice = slice_to_255(\n            t2w_slice\n        )\n\n        rgb = np.stack(\n            [\n                flair_slice,\n                t1ce_slice,\n                t2w_slice,\n            ],\n            axis=-1,\n        )\n\n        rgb = np.nan_to_num(\n            rgb,\n            nan=0.0,\n            posinf=255.0,\n            neginf=0.0,\n        )\n\n        np.clip(\n            rgb,\n            0.0,\n            255.0,\n            out=rgb,\n        )\n\n        tensor[output_index] = rgb.astype(\n            np.uint8\n        )\n\n        del (\n            flair_slice,\n            t1ce_slice,\n            t2w_slice,\n            rgb,\n        )\n\n    # --------------------------------------------------------\n    # Final validation\n    # --------------------------------------------------------\n\n    if tensor.shape != expected_shape:\n        print(\n            f\"Patient {patient_id}: obtained shape \"\n            f\"{tensor.shape}, expected {expected_shape}\"\n        )\n        return None\n\n    if not np.isfinite(\n        tensor.astype(np.float32)\n    ).all():\n        print(\n            f\"Patient {patient_id}: final tensor contains \"\n            \"NaN or infinite values\"\n        )\n        return None\n\n    # --------------------------------------------------------\n    # Save uint8 cache\n    # --------------------------------------------------------\n\n    if save_cache:\n        np.save(\n            cache_path,\n            tensor,\n        )\n\n    if verbose:\n        print(\n            \"\\nSelected FLAIR indices:\"\n        )\n        print(\n            selected_flair_indices\n        )\n\n        print(\n            \"\\nFinal tensor:\"\n        )\n        print(\n            \"Shape:\",\n            tensor.shape,\n        )\n        print(\n            \"Dtype:\",\n            tensor.dtype,\n        )\n        print(\n            \"Minimum:\",\n            tensor.min(),\n        )\n        print(\n            \"Maximum:\",\n            tensor.max(),\n        )\n        print(\n            \"Mean:\",\n            tensor.mean(),\n        )\n        print(\n            \"Standard deviation:\",\n            tensor.std(),\n        )\n\n        if save_cache:\n            print(\n                \"Saved cache:\",\n                cache_path,\n            )\n\n            print(\n                \"Cache size MB:\",\n                tensor.nbytes / 1024**2,\n            )\n\n    # --------------------------------------------------------\n    # Explicitly release full MRI volumes\n    # --------------------------------------------------------\n\n    del volumes\n    del flair\n    del t1ce\n    del t2w\n    del scores\n\n    # Return float32 only for immediate model compatibility.\n    return tensor.astype(\n        np.float32\n    )\n\n\n# ============================================================\n# Test one patient\n# ============================================================\n\nsample_id = int(\n    labels_df.iloc[0][\"BraTS21ID\"]\n)\n\nsample_tensor = build_patient_tensor(\n    patient_id=sample_id,\n    save_cache=False,\n    verbose=True,\n)\n\nprint(\n    \"\\nSample patient:\",\n    sample_id,\n)\n\nprint(\n    \"Tensor shape:\",\n    None\n    if sample_tensor is None\n    else sample_tensor.shape,\n)\n\nprint(\n    \"Returned dtype:\",\n    None\n    if sample_tensor is None\n    else sample_tensor.dtype,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:04.859699Z","iopub.execute_input":"2026-07-11T17:06:04.859948Z","iopub.status.idle":"2026-07-11T17:06:19.734758Z","shell.execute_reply.started":"2026-07-11T17:06:04.859926Z","shell.execute_reply":"2026-07-11T17:06:19.734062Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 7. Visual quality control\n# ============================================================\n\ndef show_patient(patient_id):\n    tensor = build_patient_tensor(\n        patient_id,\n        save_cache=False\n    )\n\n    if tensor is None:\n        print(\"Could not load patient\", patient_id)\n        return\n\n    print(\"Tensor shape:\", tensor.shape)\n    print(\"Min:\", tensor.min())\n    print(\"Max:\", tensor.max())\n\n    fig, axes = plt.subplots(\n        NUM_SLICES,\n        4,\n        figsize=(14, NUM_SLICES * 3)\n    )\n\n    for i in range(NUM_SLICES):\n\n        flair = tensor[i, :, :, 0]\n        t1ce = tensor[i, :, :, 1]\n        t2w = tensor[i, :, :, 2]\n\n        rgb = tensor[i].astype(np.uint8)\n\n        axes[i,0].imshow(flair, cmap=\"gray\")\n        axes[i,0].set_title(f\"Slice {i+1}\\nFLAIR\")\n        axes[i,0].axis(\"off\")\n\n        axes[i,1].imshow(t1ce, cmap=\"gray\")\n        axes[i,1].set_title(\"T1wCE\")\n        axes[i,1].axis(\"off\")\n\n        axes[i,2].imshow(t2w, cmap=\"gray\")\n        axes[i,2].set_title(\"T2w\")\n        axes[i,2].axis(\"off\")\n\n        axes[i,3].imshow(rgb)\n        axes[i,3].set_title(\"RGB\")\n        axes[i,3].axis(\"off\")\n\n    plt.tight_layout()\n    plt.show()\n\nshow_patient(sample_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:19.735509Z","iopub.execute_input":"2026-07-11T17:06:19.735774Z","iopub.status.idle":"2026-07-11T17:06:32.715438Z","shell.execute_reply.started":"2026-07-11T17:06:19.735748Z","shell.execute_reply":"2026-07-11T17:06:32.714640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 8. Precompute patient tensors\n# ============================================================\n\nimport gc\nimport time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\n\n\n# ============================================================\n# Configuration\n# ============================================================\n\n# For a quick test:\n# MAX_PATIENTS = 20\n\n# For the complete dataset:\nMAX_PATIENTS = None\n\nPROGRESS_EVERY = 25\n\nvalid_output_path = Path(\n    \"/kaggle/working/valid_patients.csv\"\n)\n\nfailed_output_path = Path(\n    \"/kaggle/working/failed_patients.csv\"\n)\n\n\n# ============================================================\n# Prepare rows to process\n# ============================================================\n\nrows_to_process = labels_df.copy()\n\nif MAX_PATIENTS is not None:\n    rows_to_process = rows_to_process.iloc[\n        :MAX_PATIENTS\n    ].copy()\n\nrows_to_process = rows_to_process.reset_index(\n    drop=True\n)\n\nexpected_shape = (\n    NUM_SLICES,\n    IMG_SIZE,\n    IMG_SIZE,\n    3,\n)\n\nprint(\n    \"Patients to process:\",\n    len(rows_to_process),\n)\n\nprint(\n    \"Cache directory:\",\n    CACHE_DIR,\n)\n\nprint(\n    \"Expected tensor shape:\",\n    expected_shape,\n)\n\n\n# ============================================================\n# Load previous progress if available\n# ============================================================\n\nvalid_rows = []\nfailed_patients = []\n\ncompleted_valid_ids = set()\ncompleted_failed_ids = set()\n\n\nif valid_output_path.exists():\n    previous_valid_df = pd.read_csv(\n        valid_output_path\n    )\n\n    if not previous_valid_df.empty:\n        valid_rows = previous_valid_df[\n            [\n                \"BraTS21ID\",\n                \"MGMT_value\",\n            ]\n        ].to_dict(\n            \"records\"\n        )\n\n        completed_valid_ids = set(\n            previous_valid_df[\n                \"BraTS21ID\"\n            ]\n            .astype(int)\n            .tolist()\n        )\n\n        print(\n            \"Loaded previous valid patients:\",\n            len(completed_valid_ids),\n        )\n\n\nif failed_output_path.exists():\n    previous_failed_df = pd.read_csv(\n        failed_output_path\n    )\n\n    if not previous_failed_df.empty:\n        failed_patients = (\n            previous_failed_df\n            .to_dict(\"records\")\n        )\n\n        completed_failed_ids = set(\n            previous_failed_df[\n                \"BraTS21ID\"\n            ]\n            .astype(int)\n            .tolist()\n        )\n\n        print(\n            \"Loaded previous failed patients:\",\n            len(completed_failed_ids),\n        )\n\n\ncompleted_ids = (\n    completed_valid_ids\n    | completed_failed_ids\n)\n\n\n# ============================================================\n# Precompute patient tensors\n# ============================================================\n\nstart_time = time.time()\n\nnewly_processed = 0\nnew_valid = 0\nnew_failed = 0\n\n\nfor processed_count, row in enumerate(\n    rows_to_process.itertuples(\n        index=False\n    ),\n    start=1,\n):\n    patient_id = int(\n        row.BraTS21ID\n    )\n\n    label = int(\n        row.MGMT_value\n    )\n\n    cache_path = (\n        CACHE_DIR\n        / f\"{patient_id:05d}.npy\"\n    )\n\n    # --------------------------------------------------------\n    # Skip patients already recorded and cached correctly\n    # --------------------------------------------------------\n\n    if patient_id in completed_valid_ids:\n        if cache_path.exists():\n            try:\n                cached = np.load(\n                    cache_path,\n                    mmap_mode=\"r\",\n                )\n\n                if (\n                    cached.shape == expected_shape\n                    and cached.dtype == np.uint8\n                ):\n                    del cached\n                    continue\n\n                del cached\n\n            except Exception:\n                pass\n\n        # Existing record is no longer valid because the cache\n        # is missing or incompatible.\n        valid_rows = [\n            item\n            for item in valid_rows\n            if int(item[\"BraTS21ID\"])\n            != patient_id\n        ]\n\n        completed_valid_ids.discard(\n            patient_id\n        )\n\n    if patient_id in completed_failed_ids:\n        continue\n\n    tensor = None\n\n    try:\n        tensor = build_patient_tensor(\n            patient_id=patient_id,\n            save_cache=True,\n            verbose=False,\n        )\n\n        if tensor is None:\n            raise ValueError(\n                \"build_patient_tensor returned None.\"\n            )\n\n        if tensor.shape != expected_shape:\n            raise ValueError(\n                f\"Unexpected returned shape: \"\n                f\"{tensor.shape}\"\n            )\n\n        # ----------------------------------------------------\n        # Verify the cached file instead of keeping the\n        # float32 tensor in memory.\n        # ----------------------------------------------------\n\n        if not cache_path.exists():\n            raise FileNotFoundError(\n                f\"Cache file was not created: \"\n                f\"{cache_path}\"\n            )\n\n        cached_tensor = np.load(\n            cache_path,\n            mmap_mode=\"r\",\n        )\n\n        if cached_tensor.shape != expected_shape:\n            raise ValueError(\n                f\"Cached shape is \"\n                f\"{cached_tensor.shape}, \"\n                f\"expected {expected_shape}\"\n            )\n\n        if cached_tensor.dtype != np.uint8:\n            raise ValueError(\n                f\"Cached dtype is \"\n                f\"{cached_tensor.dtype}, \"\n                \"expected uint8.\"\n            )\n\n        del cached_tensor\n\n        valid_rows.append({\n            \"BraTS21ID\": patient_id,\n            \"MGMT_value\": label,\n        })\n\n        completed_valid_ids.add(\n            patient_id\n        )\n\n        new_valid += 1\n\n    except Exception as exc:\n        failed_patients.append({\n            \"BraTS21ID\": patient_id,\n            \"reason\": repr(exc),\n        })\n\n        completed_failed_ids.add(\n            patient_id\n        )\n\n        new_failed += 1\n\n        print(\n            f\"\\nFailed patient {patient_id}: \"\n            f\"{repr(exc)}\"\n        )\n\n    finally:\n        if tensor is not None:\n            del tensor\n\n        newly_processed += 1\n\n        # Periodic cleanup\n        if newly_processed % 5 == 0:\n            gc.collect()\n\n    # --------------------------------------------------------\n    # Save progress regularly\n    # --------------------------------------------------------\n\n    if (\n        processed_count % PROGRESS_EVERY == 0\n        or processed_count\n        == len(rows_to_process)\n    ):\n        valid_df = pd.DataFrame(\n            valid_rows\n        ).drop_duplicates(\n            subset=\"BraTS21ID\",\n            keep=\"last\",\n        )\n\n        failed_df = pd.DataFrame(\n            failed_patients\n        ).drop_duplicates(\n            subset=\"BraTS21ID\",\n            keep=\"last\",\n        )\n\n        valid_df.to_csv(\n            valid_output_path,\n            index=False,\n        )\n\n        failed_df.to_csv(\n            failed_output_path,\n            index=False,\n        )\n\n        elapsed = (\n            time.time()\n            - start_time\n        )\n\n        print(\n            f\"Checked {processed_count}/\"\n            f\"{len(rows_to_process)} | \"\n            f\"Valid total: {len(valid_df)} | \"\n            f\"Failed total: {len(failed_df)} | \"\n            f\"New valid: {new_valid} | \"\n            f\"New failed: {new_failed} | \"\n            f\"Elapsed: {elapsed / 60:.1f} min\"\n        )\n\n        gc.collect()\n\n\n# ============================================================\n# Final summary tables\n# ============================================================\n\nvalid_df = pd.DataFrame(\n    valid_rows\n).drop_duplicates(\n    subset=\"BraTS21ID\",\n    keep=\"last\",\n).sort_values(\n    \"BraTS21ID\"\n).reset_index(\n    drop=True\n)\n\nfailed_df = pd.DataFrame(\n    failed_patients\n)\n\nif not failed_df.empty:\n    failed_df = failed_df.drop_duplicates(\n        subset=\"BraTS21ID\",\n        keep=\"last\",\n    ).sort_values(\n        \"BraTS21ID\"\n    ).reset_index(\n        drop=True\n    )\n\n\n# ============================================================\n# Save final preprocessing summaries\n# ============================================================\n\nvalid_df.to_csv(\n    valid_output_path,\n    index=False,\n)\n\nfailed_df.to_csv(\n    failed_output_path,\n    index=False,\n)\n\n\n# ============================================================\n# Final report\n# ============================================================\n\nelapsed = (\n    time.time()\n    - start_time\n)\n\nprint(\n    \"\\n\"\n    + \"=\" * 65\n)\n\nprint(\n    \"PRECOMPUTATION FINISHED\"\n)\n\nprint(\n    \"=\" * 65\n)\n\nprint(\n    \"Patients in input:\",\n    len(rows_to_process),\n)\n\nprint(\n    \"Valid patients:\",\n    len(valid_df),\n)\n\nprint(\n    \"Failed patients:\",\n    len(failed_df),\n)\n\nprint(\n    f\"Execution time: \"\n    f\"{elapsed / 60:.2f} minutes\"\n)\n\nif not valid_df.empty:\n    print(\n        \"\\nClass distribution:\"\n    )\n\n    print(\n        valid_df[\n            \"MGMT_value\"\n        ]\n        .value_counts()\n        .sort_index()\n    )\n\nif not failed_df.empty:\n    print(\n        \"\\nFirst failed patients:\"\n    )\n\n    print(\n        failed_df.head(10)\n    )\n\n\n# ============================================================\n# Cache verification\n# ============================================================\n\ncache_files = list(\n    CACHE_DIR.glob(\"*.npy\")\n)\n\nprint(\n    \"\\nNumber of cache files:\",\n    len(cache_files),\n)\n\nif cache_files:\n    sample_cache_path = (\n        cache_files[0]\n    )\n\n    sample_cache = np.load(\n        sample_cache_path,\n        mmap_mode=\"r\",\n    )\n\n    print(\n        \"Example cache file:\",\n        sample_cache_path,\n    )\n\n    print(\n        \"Cache dtype:\",\n        sample_cache.dtype,\n    )\n\n    print(\n        \"Cache shape:\",\n        sample_cache.shape,\n    )\n\n    print(\n        \"Cache size MB:\",\n        sample_cache.nbytes\n        / 1024**2,\n    )\n\n    del sample_cache\n\n\nprint(\n    \"\\nSaved:\"\n)\n\nprint(\n    valid_output_path\n)\n\nprint(\n    failed_output_path\n)\n\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T17:06:32.716321Z","iopub.execute_input":"2026-07-11T17:06:32.716583Z","iopub.status.idle":"2026-07-11T18:08:06.213013Z","shell.execute_reply.started":"2026-07-11T17:06:32.716566Z","shell.execute_reply":"2026-07-11T18:08:06.212446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 9. TensorFlow dataset\n#    Memory-efficient version\n# ============================================================\n\ndef load_cached_numpy(patient_id, label):\n    \"\"\"\n    Load one cached patient tensor from disk.\n\n    Cache format:\n        uint8, shape (NUM_SLICES, IMG_SIZE, IMG_SIZE, 3)\n\n    Returned format:\n        float32 in [0, 255]\n    \"\"\"\n\n    patient_id = int(patient_id)\n\n    path = CACHE_DIR / f\"{patient_id:05d}.npy\"\n\n    if not path.exists():\n        raise FileNotFoundError(\n            f\"Missing cached tensor: {path}\"\n        )\n\n    expected_shape = (\n        NUM_SLICES,\n        IMG_SIZE,\n        IMG_SIZE,\n        3,\n    )\n\n    # Memory-map the cached file instead of loading\n    # the complete file immediately.\n    cached = np.load(\n        path,\n        mmap_mode=\"r\",\n    )\n\n    if cached.shape != expected_shape:\n        raise ValueError(\n            f\"Patient {patient_id}: cached shape \"\n            f\"{cached.shape}, expected {expected_shape}\"\n        )\n\n    if cached.dtype != np.uint8:\n        raise ValueError(\n            f\"Patient {patient_id}: cached dtype \"\n            f\"{cached.dtype}, expected uint8. \"\n            \"Rebuild the cache.\"\n        )\n\n    # Convert only the current patient to float32.\n    x = cached.astype(\n        np.float32,\n        copy=True,\n    )\n\n    del cached\n\n    y = np.float32(label)\n\n    return x, y\n\n\ndef tf_load_patient(patient_id, label):\n    \"\"\"\n    Wrap the NumPy loader for use inside tf.data.\n    \"\"\"\n\n    x, y = tf.numpy_function(\n        func=load_cached_numpy,\n        inp=[\n            patient_id,\n            label,\n        ],\n        Tout=[\n            tf.float32,\n            tf.float32,\n        ],\n    )\n\n    x.set_shape(\n        (\n            NUM_SLICES,\n            IMG_SIZE,\n            IMG_SIZE,\n            3,\n        )\n    )\n\n    y.set_shape(())\n\n    return x, y\n\n\ndef augment_patient(x, y):\n    \"\"\"\n    Apply lightweight augmentation to one patient tensor.\n\n    The same transformation is applied across all selected\n    slices, preserving patient-level consistency.\n\n    Input and output range:\n        [0, 255]\n    \"\"\"\n\n    # --------------------------------------------------------\n    # Horizontal flip\n    # --------------------------------------------------------\n    do_flip = (\n        tf.random.uniform(\n            shape=(),\n            minval=0.0,\n            maxval=1.0,\n            dtype=tf.float32,\n        )\n        < 0.5\n    )\n\n    x = tf.cond(\n        do_flip,\n        lambda: tf.reverse(\n            x,\n            axis=[2],\n        ),\n        lambda: x,\n    )\n\n    # --------------------------------------------------------\n    # Mild brightness variation\n    # --------------------------------------------------------\n    brightness_delta = tf.random.uniform(\n        shape=(),\n        minval=-6.0,\n        maxval=6.0,\n        dtype=tf.float32,\n    )\n\n    x = tf.image.adjust_brightness(\n        x,\n        delta=brightness_delta,\n    )\n\n    # --------------------------------------------------------\n    # Mild contrast variation\n    # --------------------------------------------------------\n    contrast_factor = tf.random.uniform(\n        shape=(),\n        minval=0.95,\n        maxval=1.05,\n        dtype=tf.float32,\n    )\n\n    x = tf.image.adjust_contrast(\n        x,\n        contrast_factor=contrast_factor,\n    )\n\n    # --------------------------------------------------------\n    # Modality dropout\n    # --------------------------------------------------------\n    drop_probability = tf.random.uniform(\n        shape=(),\n        minval=0.0,\n        maxval=1.0,\n        dtype=tf.float32,\n    )\n\n    channel_to_drop = tf.random.uniform(\n        shape=(),\n        minval=0,\n        maxval=3,\n        dtype=tf.int32,\n    )\n\n    def drop_one_channel():\n        channel_mask = tf.ones(\n            shape=(3,),\n            dtype=tf.float32,\n        )\n\n        channel_mask = tf.tensor_scatter_nd_update(\n            tensor=channel_mask,\n            indices=tf.reshape(\n                channel_to_drop,\n                shape=(1, 1),\n            ),\n            updates=tf.constant(\n                [0.0],\n                dtype=tf.float32,\n            ),\n        )\n\n        return x * channel_mask\n\n    x = tf.cond(\n        drop_probability < 0.10,\n        drop_one_channel,\n        lambda: x,\n    )\n\n    # Keep EfficientNet-compatible range.\n    x = tf.clip_by_value(\n        x,\n        clip_value_min=0.0,\n        clip_value_max=255.0,\n    )\n\n    return x, y\n\n\ndef make_dataset(\n    frame,\n    training=False,\n    batch_size=BATCH_SIZE,\n):\n    \"\"\"\n    Create a patient-level tf.data pipeline with limited\n    buffering to reduce CPU RAM usage.\n    \"\"\"\n\n    if frame.empty:\n        raise ValueError(\n            \"The dataframe used to build the dataset is empty.\"\n        )\n\n    patient_ids = (\n        frame[\"BraTS21ID\"]\n        .values\n        .astype(np.int32)\n    )\n\n    labels = (\n        frame[\"MGMT_value\"]\n        .values\n        .astype(np.float32)\n    )\n\n    ds = tf.data.Dataset.from_tensor_slices(\n        (\n            patient_ids,\n            labels,\n        )\n    )\n\n    if training:\n        ds = ds.shuffle(\n            buffer_size=min(\n                len(frame),\n                32,\n            ),\n            seed=SEED,\n            reshuffle_each_iteration=True,\n        )\n\n    ds = ds.map(\n        tf_load_patient,\n        num_parallel_calls=1,\n        deterministic=True,\n    )\n\n    if training:\n        ds = ds.map(\n            augment_patient,\n            num_parallel_calls=1,\n            deterministic=True,\n        )\n\n    ds = ds.batch(\n        batch_size,\n        drop_remainder=False,\n    )\n\n    ds = ds.prefetch(1)\n\n    return ds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T18:08:06.216973Z","iopub.execute_input":"2026-07-11T18:08:06.217616Z","iopub.status.idle":"2026-07-11T18:08:06.234028Z","shell.execute_reply.started":"2026-07-11T18:08:06.217597Z","shell.execute_reply":"2026-07-11T18:08:06.233140Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 10. Model definition\n# ============================================================\n\nimport tensorflow as tf\nfrom tensorflow.keras import layers, Model\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras.optimizers import Adam\n\n\nclass GatedAttentionPooling(layers.Layer):\n\n\n    def __init__(\n        self,\n        hidden_dim=128,\n        dropout_rate=0.10,\n        **kwargs,\n    ):\n        super().__init__(**kwargs)\n\n        self.hidden_dim = hidden_dim\n        self.dropout_rate = dropout_rate\n\n        self.attention_tanh = layers.Dense(\n            hidden_dim,\n            activation=\"tanh\",\n            name=\"attention_tanh\",\n        )\n\n        self.attention_sigmoid = layers.Dense(\n            hidden_dim,\n            activation=\"sigmoid\",\n            name=\"attention_sigmoid\",\n        )\n\n        self.attention_dropout = layers.Dropout(\n            dropout_rate\n        )\n\n        self.attention_score = layers.Dense(\n            1,\n            use_bias=False,\n            name=\"attention_score\",\n        )\n\n        self.last_attention_weights = None\n\n    def call(self, inputs, training=None):\n        tanh_branch = self.attention_tanh(inputs)\n        sigmoid_branch = self.attention_sigmoid(inputs)\n\n        gated = tanh_branch * sigmoid_branch\n\n        gated = self.attention_dropout(\n            gated,\n            training=training,\n        )\n\n        logits = self.attention_score(gated)\n\n        weights = tf.nn.softmax(\n            logits,\n            axis=1,\n        )\n\n        self.last_attention_weights = weights\n\n        pooled = tf.reduce_sum(\n            inputs * weights,\n            axis=1,\n        )\n\n        return pooled\n\n    def get_config(self):\n        config = super().get_config()\n\n        config.update({\n            \"hidden_dim\": self.hidden_dim,\n            \"dropout_rate\": self.dropout_rate,\n        })\n\n        return config\n\n\ndef build_model(\n    backbone_trainable=False,\n):\n    # --------------------------------------------------------\n    # Input\n    # --------------------------------------------------------\n    input_tensor = layers.Input(\n        shape=(\n            NUM_SLICES,\n            IMG_SIZE,\n            IMG_SIZE,\n            3,\n        ),\n        name=\"patient_slices\",\n        dtype=tf.float32,\n    )\n\n    # --------------------------------------------------------\n    # Shared 2D CNN encoder\n    # --------------------------------------------------------\n    backbone = EfficientNetB0(\n        include_top=False,\n        weights=\"imagenet\",\n        input_shape=(\n            IMG_SIZE,\n            IMG_SIZE,\n            3,\n        ),\n        pooling=\"avg\",\n    )\n\n    backbone.trainable = backbone_trainable\n\n    slice_features = layers.TimeDistributed(\n        backbone,\n        name=\"slice_encoder\",\n    )(input_tensor)\n\n    # EfficientNetB0 pooled feature size is normally 1280.\n    # Output shape:\n    # (batch, NUM_SLICES, 1280)\n\n    # --------------------------------------------------------\n    # Slice-level feature projection\n    # --------------------------------------------------------\n    x = layers.TimeDistributed(\n        layers.LayerNormalization(),\n        name=\"slice_feature_normalization\",\n    )(slice_features)\n\n    x = layers.TimeDistributed(\n        layers.Dense(\n            256,\n            activation=tf.nn.swish,\n            kernel_regularizer=tf.keras.regularizers.l2(\n                1e-5\n            ),\n        ),\n        name=\"slice_projection\",\n    )(x)\n\n    x = layers.TimeDistributed(\n        layers.Dropout(0.30),\n        name=\"slice_dropout\",\n    )(x)\n\n    # --------------------------------------------------------\n    # Patient-level aggregation\n    # --------------------------------------------------------\n    x = GatedAttentionPooling(\n        hidden_dim=128,\n        dropout_rate=0.10,\n        name=\"attention_pooling\",\n    )(x)\n\n    # LayerNormalization is safer than BatchNormalization\n    # when BATCH_SIZE is small.\n    x = layers.LayerNormalization(\n        name=\"patient_feature_normalization\",\n    )(x)\n\n    # --------------------------------------------------------\n    # Classification head\n    # --------------------------------------------------------\n    x = layers.Dense(\n        128,\n        activation=tf.nn.swish,\n        kernel_regularizer=tf.keras.regularizers.l2(\n            1e-4\n        ),\n        name=\"patient_dense\",\n    )(x)\n\n    x = layers.Dropout(\n        0.45,\n        name=\"patient_dropout\",\n    )(x)\n\n    output = layers.Dense(\n        1,\n        activation=\"sigmoid\",\n        name=\"mgmt_probability\",\n        dtype=tf.float32,\n    )(x)\n\n    model = Model(\n        inputs=input_tensor,\n        outputs=output,\n        name=\"MGMT_EfficientNet_Attention\",\n    )\n\n    return model\n\n\ndef compile_model(\n    model,\n    learning_rate,\n):\n    optimizer = Adam(\n        learning_rate=learning_rate,\n        clipnorm=1.0,\n    )\n\n    model.compile(\n        optimizer=optimizer,\n        loss=tf.keras.losses.BinaryCrossentropy(\n            label_smoothing=0.02,\n        ),\n        metrics=[\n            tf.keras.metrics.BinaryAccuracy(\n                name=\"accuracy\",\n                threshold=0.5,\n            ),\n            tf.keras.metrics.AUC(\n                name=\"auc\",\n                curve=\"ROC\",\n                num_thresholds=200,\n            ),\n            tf.keras.metrics.AUC(\n                name=\"pr_auc\",\n                curve=\"PR\",\n                num_thresholds=200,\n            ),\n            tf.keras.metrics.Precision(\n                name=\"precision\",\n            ),\n            tf.keras.metrics.Recall(\n                name=\"recall\",\n            ),\n        ],\n    )\n\n\nmodel = build_model(\n    backbone_trainable=False\n)\n\ncompile_model(\n    model,\n    LEARNING_RATE_HEAD\n)\n\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T18:08:06.236754Z","iopub.execute_input":"2026-07-11T18:08:06.237274Z","iopub.status.idle":"2026-07-11T18:08:11.895951Z","shell.execute_reply.started":"2026-07-11T18:08:06.237257Z","shell.execute_reply":"2026-07-11T18:08:11.895323Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 11. Helper functions for evaluation\n# ============================================================\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    roc_curve,\n    precision_recall_curve,\n    average_precision_score,\n    confusion_matrix,\n    classification_report,\n    accuracy_score,\n    balanced_accuracy_score,\n    f1_score,\n    matthews_corrcoef,\n)\n\n\ndef find_best_threshold(y_true, probabilities):\n    \"\"\"\n    Compute the optimal threshold using Youden's J statistic.\n    \"\"\"\n\n    fpr, tpr, thresholds = roc_curve(\n        y_true,\n        probabilities,\n    )\n\n    j = tpr - fpr\n    best_index = np.argmax(j)\n\n    return float(thresholds[best_index])\n\n\ndef evaluate_predictions(\n    y_true,\n    probabilities,\n    threshold=None,\n    show_plots=True,\n):\n    \"\"\"\n    Evaluate patient-level predictions.\n    \"\"\"\n\n    y_true = np.asarray(y_true)\n    probabilities = np.asarray(probabilities)\n\n    if threshold is None:\n        threshold = find_best_threshold(\n            y_true,\n            probabilities,\n        )\n\n    predictions = (\n        probabilities >= threshold\n    ).astype(np.int32)\n\n    cm = confusion_matrix(\n        y_true,\n        predictions,\n    )\n\n    tn, fp, fn, tp = cm.ravel()\n\n    sensitivity = (\n        tp / (tp + fn + 1e-8)\n    )\n\n    specificity = (\n        tn / (tn + fp + 1e-8)\n    )\n\n    precision = (\n        tp / (tp + fp + 1e-8)\n    )\n\n    npv = (\n        tn / (tn + fn + 1e-8)\n    )\n\n    metrics = {\n        \"threshold\": threshold,\n        \"auc\": roc_auc_score(\n            y_true,\n            probabilities,\n        ),\n        \"pr_auc\": average_precision_score(\n            y_true,\n            probabilities,\n        ),\n        \"accuracy\": accuracy_score(\n            y_true,\n            predictions,\n        ),\n        \"balanced_accuracy\": balanced_accuracy_score(\n            y_true,\n            predictions,\n        ),\n        \"f1\": f1_score(\n            y_true,\n            predictions,\n        ),\n        \"mcc\": matthews_corrcoef(\n            y_true,\n            predictions,\n        ),\n        \"sensitivity\": sensitivity,\n        \"specificity\": specificity,\n        \"precision\": precision,\n        \"npv\": npv,\n    }\n\n    print(\"=\" * 60)\n\n    for key, value in metrics.items():\n        print(f\"{key:20s}: {value:.4f}\")\n\n    print(\"=\" * 60)\n\n    print(\"\\nClassification report\\n\")\n    print(\n        classification_report(\n            y_true,\n            predictions,\n            digits=4,\n        )\n    )\n\n    if show_plots:\n\n        # --------------------------------------------------\n        # Confusion matrix\n        # --------------------------------------------------\n\n        plt.figure(figsize=(5,5))\n\n        plt.imshow(\n            cm,\n            cmap=\"Blues\",\n        )\n\n        plt.title(\"Confusion Matrix\")\n        plt.xlabel(\"Predicted\")\n        plt.ylabel(\"True\")\n\n        for i in range(2):\n            for j in range(2):\n                plt.text(\n                    j,\n                    i,\n                    str(cm[i,j]),\n                    ha=\"center\",\n                    va=\"center\",\n                    fontsize=14,\n                )\n\n        plt.colorbar()\n        plt.tight_layout()\n        plt.show()\n\n        # --------------------------------------------------\n        # ROC\n        # --------------------------------------------------\n\n        fpr, tpr, _ = roc_curve(\n            y_true,\n            probabilities,\n        )\n\n        plt.figure(figsize=(6,5))\n\n        plt.plot(\n            fpr,\n            tpr,\n            lw=2,\n            label=f\"AUC = {metrics['auc']:.4f}\",\n        )\n\n        plt.plot(\n            [0,1],\n            [0,1],\n            \"--\",\n            color=\"gray\",\n        )\n\n        plt.xlabel(\"False Positive Rate\")\n        plt.ylabel(\"True Positive Rate\")\n        plt.title(\"ROC Curve\")\n        plt.grid(True)\n        plt.legend()\n\n        plt.show()\n\n        # --------------------------------------------------\n        # Precision-Recall\n        # --------------------------------------------------\n\n        precision_curve, recall_curve, _ = (\n            precision_recall_curve(\n                y_true,\n                probabilities,\n            )\n        )\n\n        plt.figure(figsize=(6,5))\n\n        plt.plot(\n            recall_curve,\n            precision_curve,\n            lw=2,\n            label=f\"PR-AUC = {metrics['pr_auc']:.4f}\",\n        )\n\n        plt.xlabel(\"Recall\")\n        plt.ylabel(\"Precision\")\n        plt.title(\"Precision-Recall Curve\")\n        plt.grid(True)\n        plt.legend()\n\n        plt.show()\n\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T18:08:11.896801Z","iopub.execute_input":"2026-07-11T18:08:11.897129Z","iopub.status.idle":"2026-07-11T18:08:11.909173Z","shell.execute_reply.started":"2026-07-11T18:08:11.897112Z","shell.execute_reply":"2026-07-11T18:08:11.908422Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 12. Train one cross-validation fold at a time\n# ============================================================\n\nimport gc\nimport os\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    accuracy_score,\n    balanced_accuracy_score,\n    f1_score,\n    matthews_corrcoef,\n    confusion_matrix,\n)\n\nfrom tensorflow.keras.callbacks import (\n    ModelCheckpoint,\n    EarlyStopping,\n    ReduceLROnPlateau,\n)\n\n\n# ============================================================\n# Choose the fold to train\n# ============================================================\n\n# Fold 1, 2, 3, 4 and 5 run the script for each one because \n# the cpu ram in kagglle not allowed to run the five together:\nCURRENT_FOLD = 2\n\n\nif CURRENT_FOLD < 1 or CURRENT_FOLD > N_FOLDS:\n    raise ValueError(\n        f\"CURRENT_FOLD must be between 1 and {N_FOLDS}.\"\n    )\n\n\n# ============================================================\n# Output directory\n# ============================================================\n\nCV_OUTPUT_DIR = Path(\"/kaggle/working/cv_results\")\nCV_OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n\nprediction_path = (\n    CV_OUTPUT_DIR\n    / f\"fold_{CURRENT_FOLD}_predictions.csv\"\n)\n\nmetrics_path = (\n    CV_OUTPUT_DIR\n    / f\"fold_{CURRENT_FOLD}_metrics.csv\"\n)\n\nhead_path = (\n    CV_OUTPUT_DIR\n    / f\"fold_{CURRENT_FOLD}_head.weights.h5\"\n)\n\nfine_path = (\n    CV_OUTPUT_DIR\n    / f\"fold_{CURRENT_FOLD}_best.weights.h5\"\n)\n\n\n# ============================================================\n# Avoid accidentally retraining a completed fold\n# ============================================================\n\nif prediction_path.exists() and metrics_path.exists():\n    raise FileExistsError(\n        f\"Fold {CURRENT_FOLD} already appears to be completed.\\n\"\n        f\"Predictions: {prediction_path}\\n\"\n        f\"Metrics: {metrics_path}\\n\\n\"\n        \"Change CURRENT_FOLD to the next fold or delete these \"\n        \"files if you intentionally want to rerun it.\"\n    )\n\n\n# ============================================================\n# TensorFlow memory configuration\n# ============================================================\n\ngpus = tf.config.list_physical_devices(\"GPU\")\n\nfor gpu in gpus:\n    try:\n        tf.config.experimental.set_memory_growth(\n            gpu,\n            True,\n        )\n    except RuntimeError:\n        pass\n\ntf.keras.backend.clear_session()\ngc.collect()\n\n\n# ============================================================\n# Create the exact same cross-validation splits\n# ============================================================\n\nskf = StratifiedKFold(\n    n_splits=N_FOLDS,\n    shuffle=True,\n    random_state=SEED,\n)\n\nsplits = list(\n    skf.split(\n        valid_df,\n        valid_df[\"MGMT_value\"],\n    )\n)\n\ntrain_idx, val_idx = splits[CURRENT_FOLD - 1]\n\n\n# ============================================================\n# Prepare train and validation dataframes\n# ============================================================\n\ntrain_fold = (\n    valid_df\n    .iloc[train_idx]\n    .reset_index(drop=True)\n)\n\nval_fold = (\n    valid_df\n    .iloc[val_idx]\n    .reset_index(drop=True)\n)\n\nprint(\"\\n\" + \"=\" * 75)\nprint(f\"TRAINING FOLD {CURRENT_FOLD}/{N_FOLDS}\")\nprint(\"=\" * 75)\n\nprint(\"\\nTrain patients:\", len(train_fold))\nprint(\"Validation patients:\", len(val_fold))\n\nprint(\"\\nTrain distribution:\")\nprint(\n    train_fold[\"MGMT_value\"]\n    .value_counts()\n    .sort_index()\n)\n\nprint(\"\\nValidation distribution:\")\nprint(\n    val_fold[\"MGMT_value\"]\n    .value_counts()\n    .sort_index()\n)\n\n\n# ============================================================\n# TensorFlow datasets\n# ============================================================\n\ntrain_ds = make_dataset(\n    train_fold,\n    training=True,\n    batch_size=BATCH_SIZE,\n)\n\nval_ds = make_dataset(\n    val_fold,\n    training=False,\n    batch_size=BATCH_SIZE,\n)\n\n\n# ============================================================\n# Stage 1 — train only the classification head\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"STAGE 1 — FROZEN EFFICIENTNET\")\nprint(\"=\" * 75)\n\nseed_everything(SEED + CURRENT_FOLD)\n\nmodel = build_model(\n    backbone_trainable=False,\n)\n\ncompile_model(\n    model,\n    learning_rate=LEARNING_RATE_HEAD,\n)\n\nhead_callbacks = [\n    ModelCheckpoint(\n        filepath=str(head_path),\n        monitor=\"val_auc\",\n        mode=\"max\",\n        save_best_only=True,\n        save_weights_only=True,\n        verbose=1,\n    ),\n\n    EarlyStopping(\n        monitor=\"val_auc\",\n        mode=\"max\",\n        patience=4,\n        min_delta=1e-4,\n        restore_best_weights=True,\n        verbose=1,\n    ),\n\n    ReduceLROnPlateau(\n        monitor=\"val_loss\",\n        mode=\"min\",\n        factor=0.5,\n        patience=2,\n        min_lr=1e-6,\n        verbose=1,\n    ),\n]\n\nmodel.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=EPOCHS_HEAD,\n    callbacks=head_callbacks,\n    verbose=1,\n)\n\nmodel.load_weights(str(head_path))\n\ndel head_callbacks\ngc.collect()\n\n\n# ============================================================\n# Stage 2 — fine-tune final EfficientNet layers\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 75)\nprint(\"STAGE 2 — LIMITED FINE-TUNING\")\nprint(\"=\" * 75)\n\nbackbone = model.get_layer(\n    \"slice_encoder\"\n).layer\n\nbackbone.trainable = True\n\nNUMBER_OF_FINE_TUNED_LAYERS = 15\n\nfor layer in backbone.layers[\n    :-NUMBER_OF_FINE_TUNED_LAYERS\n]:\n    layer.trainable = False\n\nfor layer in backbone.layers[\n    -NUMBER_OF_FINE_TUNED_LAYERS:\n]:\n    if isinstance(\n        layer,\n        tf.keras.layers.BatchNormalization,\n    ):\n        layer.trainable = False\n    else:\n        layer.trainable = True\n\ncompile_model(\n    model,\n    learning_rate=LEARNING_RATE_FINE,\n)\n\nfine_callbacks = [\n    ModelCheckpoint(\n        filepath=str(fine_path),\n        monitor=\"val_auc\",\n        mode=\"max\",\n        save_best_only=True,\n        save_weights_only=True,\n        verbose=1,\n    ),\n\n    EarlyStopping(\n        monitor=\"val_auc\",\n        mode=\"max\",\n        patience=5,\n        min_delta=1e-4,\n        restore_best_weights=True,\n        verbose=1,\n    ),\n\n    ReduceLROnPlateau(\n        monitor=\"val_loss\",\n        mode=\"min\",\n        factor=0.5,\n        patience=2,\n        min_lr=1e-7,\n        verbose=1,\n    ),\n]\n\nmodel.fit(\n    train_ds,\n    validation_data=val_ds,\n    epochs=EPOCHS_FINE,\n    callbacks=fine_callbacks,\n    verbose=1,\n)\n\nmodel.load_weights(str(fine_path))\n\ndel fine_callbacks\ngc.collect()\n\n\n# ============================================================\n# Validation predictions\n# ============================================================\n\nval_probabilities = model.predict(\n    val_ds,\n    verbose=1,\n).reshape(-1).astype(np.float32)\n\nval_labels = (\n    val_fold[\"MGMT_value\"]\n    .values\n    .astype(np.int32)\n)\n\nval_predictions = (\n    val_probabilities >= 0.5\n).astype(np.int32)\n\n\n# ============================================================\n# Metrics\n# ============================================================\n\nfold_auc = roc_auc_score(\n    val_labels,\n    val_probabilities,\n)\n\nfold_pr_auc = average_precision_score(\n    val_labels,\n    val_probabilities,\n)\n\nfold_accuracy = accuracy_score(\n    val_labels,\n    val_predictions,\n)\n\nfold_balanced_accuracy = balanced_accuracy_score(\n    val_labels,\n    val_predictions,\n)\n\nfold_f1 = f1_score(\n    val_labels,\n    val_predictions,\n    zero_division=0,\n)\n\nfold_mcc = matthews_corrcoef(\n    val_labels,\n    val_predictions,\n)\n\ncm = confusion_matrix(\n    val_labels,\n    val_predictions,\n    labels=[0, 1],\n)\n\ntn, fp, fn, tp = cm.ravel()\n\nsensitivity = tp / max(tp + fn, 1)\nspecificity = tn / max(tn + fp, 1)\n\ndiagnostic_threshold = find_best_threshold(\n    val_labels,\n    val_probabilities,\n)\n\n\n# ============================================================\n# Save patient-level predictions\n# ============================================================\n\nfold_predictions_df = pd.DataFrame({\n    \"fold\": CURRENT_FOLD,\n    \"original_index\": val_idx,\n    \"BraTS21ID\": (\n        valid_df.iloc[val_idx][\"BraTS21ID\"]\n        .values\n        .astype(int)\n    ),\n    \"MGMT_value\": val_labels,\n    \"probability\": val_probabilities,\n    \"prediction_0.5\": val_predictions,\n})\n\nfold_predictions_df.to_csv(\n    prediction_path,\n    index=False,\n)\n\n\n# ============================================================\n# Save fold metrics\n# ============================================================\n\nfold_metrics_df = pd.DataFrame([{\n    \"fold\": CURRENT_FOLD,\n    \"n_train\": len(train_fold),\n    \"n_validation\": len(val_fold),\n    \"roc_auc\": fold_auc,\n    \"pr_auc\": fold_pr_auc,\n    \"accuracy_0.5\": fold_accuracy,\n    \"balanced_accuracy_0.5\": fold_balanced_accuracy,\n    \"f1_0.5\": fold_f1,\n    \"mcc_0.5\": fold_mcc,\n    \"sensitivity_0.5\": sensitivity,\n    \"specificity_0.5\": specificity,\n    \"diagnostic_best_threshold\": diagnostic_threshold,\n    \"tn\": tn,\n    \"fp\": fp,\n    \"fn\": fn,\n    \"tp\": tp,\n}])\n\nfold_metrics_df.to_csv(\n    metrics_path,\n    index=False,\n)\n\n\n# ============================================================\n# Display fold result\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 75)\nprint(f\"FOLD {CURRENT_FOLD} COMPLETED\")\nprint(\"=\" * 75)\n\nprint(f\"ROC-AUC:             {fold_auc:.4f}\")\nprint(f\"PR-AUC:              {fold_pr_auc:.4f}\")\nprint(f\"Accuracy:            {fold_accuracy:.4f}\")\nprint(\n    f\"Balanced accuracy:   \"\n    f\"{fold_balanced_accuracy:.4f}\"\n)\nprint(f\"F1-score:            {fold_f1:.4f}\")\nprint(f\"MCC:                 {fold_mcc:.4f}\")\nprint(f\"Sensitivity:         {sensitivity:.4f}\")\nprint(f\"Specificity:         {specificity:.4f}\")\nprint(\n    f\"Diagnostic threshold: \"\n    f\"{diagnostic_threshold:.4f}\"\n)\n\nprint(\"\\nSaved predictions:\")\nprint(prediction_path)\n\nprint(\"\\nSaved metrics:\")\nprint(metrics_path)\n\nprint(\"\\nSaved model:\")\nprint(fine_path)\n\n\n# ============================================================\n# Strong memory cleanup\n# ============================================================\n\ndel val_probabilities\ndel val_predictions\ndel val_labels\ndel train_fold\ndel val_fold\ndel train_ds\ndel val_ds\ndel backbone\ndel model\n\ntf.keras.backend.clear_session()\n\ngc.collect()\ngc.collect()\n\nprint(\"\\nMemory cleanup completed.\")\nprint(\n    f\"Next: change CURRENT_FOLD to \"\n    f\"{CURRENT_FOLD + 1}.\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T18:27:26.906118Z","iopub.execute_input":"2026-07-11T18:27:26.906877Z","iopub.status.idle":"2026-07-11T18:45:54.282336Z","shell.execute_reply.started":"2026-07-11T18:27:26.906852Z","shell.execute_reply":"2026-07-11T18:45:54.281636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# 13. Flexible merge and OOF evaluation\n# ============================================================\n\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom sklearn.metrics import (\n    roc_auc_score,\n    average_precision_score,\n    accuracy_score,\n    balanced_accuracy_score,\n    f1_score,\n    matthews_corrcoef,\n    confusion_matrix,\n    classification_report,\n    ConfusionMatrixDisplay,\n    roc_curve,\n    precision_recall_curve,\n)\n\n\n# ============================================================\n# Choose the folds to evaluate\n# ============================================================\n\n# Quick test with two completed folds:\nFOLDS_TO_EVALUATE = [1, 2]\n\n# Final evaluation after completing all folds:\n# FOLDS_TO_EVALUATE = list(range(1, N_FOLDS + 1))\n\n# Automatic option: use every completed fold found on disk.\n# Set USE_ALL_AVAILABLE_FOLDS = True to activate it.\nUSE_ALL_AVAILABLE_FOLDS = False\n\n\n# ============================================================\n# Paths\n# ============================================================\n\nCV_OUTPUT_DIR = Path(\n    \"/kaggle/working/cv_results\"\n)\n\nFINAL_OUTPUT_DIR = Path(\n    \"/kaggle/working/final_results\"\n)\n\nFINAL_OUTPUT_DIR.mkdir(\n    parents=True,\n    exist_ok=True,\n)\n\n\n# ============================================================\n# Detect completed folds automatically, if requested\n# ============================================================\n\nif USE_ALL_AVAILABLE_FOLDS:\n    available_folds = []\n\n    for fold in range(1, N_FOLDS + 1):\n        prediction_path = (\n            CV_OUTPUT_DIR\n            / f\"fold_{fold}_predictions.csv\"\n        )\n\n        metrics_path = (\n            CV_OUTPUT_DIR\n            / f\"fold_{fold}_metrics.csv\"\n        )\n\n        if (\n            prediction_path.exists()\n            and metrics_path.exists()\n        ):\n            available_folds.append(fold)\n\n    FOLDS_TO_EVALUATE = available_folds\n\n\n# ============================================================\n# Validate fold selection\n# ============================================================\n\nif not FOLDS_TO_EVALUATE:\n    raise ValueError(\n        \"FOLDS_TO_EVALUATE is empty. \"\n        \"Select at least one completed fold.\"\n    )\n\nFOLDS_TO_EVALUATE = sorted(\n    set(\n        int(fold)\n        for fold in FOLDS_TO_EVALUATE\n    )\n)\n\ninvalid_folds = [\n    fold\n    for fold in FOLDS_TO_EVALUATE\n    if fold < 1 or fold > N_FOLDS\n]\n\nif invalid_folds:\n    raise ValueError(\n        f\"Invalid folds: {invalid_folds}. \"\n        f\"Valid fold numbers are between 1 and {N_FOLDS}.\"\n    )\n\nprint(\n    \"Folds selected for evaluation:\",\n    FOLDS_TO_EVALUATE,\n)\n\n\n# ============================================================\n# Build file lists\n# ============================================================\n\nprediction_files = [\n    CV_OUTPUT_DIR\n    / f\"fold_{fold}_predictions.csv\"\n    for fold in FOLDS_TO_EVALUATE\n]\n\nmetrics_files = [\n    CV_OUTPUT_DIR\n    / f\"fold_{fold}_metrics.csv\"\n    for fold in FOLDS_TO_EVALUATE\n]\n\n\n# ============================================================\n# Check selected fold files\n# ============================================================\n\nmissing_predictions = [\n    path\n    for path in prediction_files\n    if not path.exists()\n]\n\nmissing_metrics = [\n    path\n    for path in metrics_files\n    if not path.exists()\n]\n\nif missing_predictions or missing_metrics:\n\n    if missing_predictions:\n        print(\"\\nMissing prediction files:\")\n\n        for path in missing_predictions:\n            print(\"-\", path)\n\n    if missing_metrics:\n        print(\"\\nMissing metrics files:\")\n\n        for path in missing_metrics:\n            print(\"-\", path)\n\n    raise RuntimeError(\n        \"Some selected folds are not complete. \"\n        \"Train them first or remove them from \"\n        \"FOLDS_TO_EVALUATE.\"\n    )\n\n\n# ============================================================\n# Merge selected fold predictions\n# ============================================================\n\noof_df = pd.concat(\n    [\n        pd.read_csv(path)\n        for path in prediction_files\n    ],\n    ignore_index=True,\n)\n\noof_df = (\n    oof_df\n    .sort_values(\"original_index\")\n    .reset_index(drop=True)\n)\n\n\n# ============================================================\n# Validate merged predictions\n# ============================================================\n\nrequired_prediction_columns = [\n    \"fold\",\n    \"original_index\",\n    \"BraTS21ID\",\n    \"MGMT_value\",\n    \"probability\",\n]\n\nmissing_columns = [\n    column\n    for column in required_prediction_columns\n    if column not in oof_df.columns\n]\n\nif missing_columns:\n    raise ValueError(\n        \"Missing columns in prediction files: \"\n        f\"{missing_columns}\"\n    )\n\nif oof_df[\"original_index\"].duplicated().any():\n    duplicated_indices = (\n        oof_df.loc[\n            oof_df[\"original_index\"].duplicated(),\n            \"original_index\",\n        ]\n        .astype(int)\n        .tolist()\n    )\n\n    raise ValueError(\n        \"Some patients have duplicated OOF predictions. \"\n        f\"Duplicated indices: {duplicated_indices[:10]}\"\n    )\n\nif not set(\n    oof_df[\"fold\"].astype(int).unique()\n).issubset(\n    set(FOLDS_TO_EVALUATE)\n):\n    raise ValueError(\n        \"The merged predictions contain folds that were \"\n        \"not selected.\"\n    )\n\n\n# ============================================================\n# Partial versus complete OOF information\n# ============================================================\n\nevaluated_patients = len(oof_df)\ntotal_patients = len(valid_df)\n\nis_complete_oof = (\n    evaluated_patients == total_patients\n    and len(FOLDS_TO_EVALUATE) == N_FOLDS\n)\n\ncoverage = (\n    evaluated_patients\n    / max(total_patients, 1)\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"OOF COVERAGE\")\nprint(\"=\" * 80)\n\nprint(\n    \"Evaluated folds:\",\n    FOLDS_TO_EVALUATE,\n)\n\nprint(\n    \"Number of evaluated folds:\",\n    len(FOLDS_TO_EVALUATE),\n)\n\nprint(\n    f\"Evaluated patients: \"\n    f\"{evaluated_patients}/{total_patients}\"\n)\n\nprint(\n    f\"Patient coverage: \"\n    f\"{coverage * 100:.2f}%\"\n)\n\nif is_complete_oof:\n    print(\n        \"Status: complete OOF evaluation.\"\n    )\nelse:\n    print(\n        \"Status: partial OOF evaluation. \"\n        \"These results are preliminary and do not represent \"\n        \"the complete cross-validation experiment.\"\n    )\n\n\n# ============================================================\n# Merge selected fold metrics\n# ============================================================\n\nfold_results_df = pd.concat(\n    [\n        pd.read_csv(path)\n        for path in metrics_files\n    ],\n    ignore_index=True,\n)\n\nfold_results_df = (\n    fold_results_df\n    .sort_values(\"fold\")\n    .reset_index(drop=True)\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"RESULTS FOR SELECTED FOLDS\")\nprint(\"=\" * 80)\n\ndisplay_columns = [\n    \"fold\",\n    \"roc_auc\",\n    \"pr_auc\",\n    \"accuracy_0.5\",\n    \"balanced_accuracy_0.5\",\n    \"f1_0.5\",\n    \"mcc_0.5\",\n    \"sensitivity_0.5\",\n    \"specificity_0.5\",\n]\n\navailable_display_columns = [\n    column\n    for column in display_columns\n    if column in fold_results_df.columns\n]\n\ndisplay(\n    fold_results_df[\n        available_display_columns\n    ].round(4)\n)\n\n\n# ============================================================\n# Extract labels and probabilities\n# ============================================================\n\ny_true = (\n    oof_df[\"MGMT_value\"]\n    .values\n    .astype(np.int32)\n)\n\ny_prob = (\n    oof_df[\"probability\"]\n    .values\n    .astype(np.float32)\n)\n\nif not np.isfinite(y_prob).all():\n    raise ValueError(\n        \"OOF probabilities contain NaN or infinite values.\"\n    )\n\nif len(np.unique(y_true)) < 2:\n    raise ValueError(\n        \"The selected folds contain only one class. \"\n        \"ROC-AUC cannot be calculated.\"\n    )\n\n\n# ============================================================\n# Thresholds\n# ============================================================\n\noptimized_threshold = find_best_threshold(\n    y_true,\n    y_prob,\n)\n\ny_pred_optimized = (\n    y_prob >= optimized_threshold\n).astype(np.int32)\n\ny_pred_05 = (\n    y_prob >= 0.5\n).astype(np.int32)\n\n\n# ============================================================\n# Metrics with optimized threshold\n# ============================================================\n\noof_auc = roc_auc_score(\n    y_true,\n    y_prob,\n)\n\noof_pr_auc = average_precision_score(\n    y_true,\n    y_prob,\n)\n\noptimized_accuracy = accuracy_score(\n    y_true,\n    y_pred_optimized,\n)\n\noptimized_balanced_accuracy = (\n    balanced_accuracy_score(\n        y_true,\n        y_pred_optimized,\n    )\n)\n\noptimized_f1 = f1_score(\n    y_true,\n    y_pred_optimized,\n    zero_division=0,\n)\n\noptimized_mcc = matthews_corrcoef(\n    y_true,\n    y_pred_optimized,\n)\n\ncm_optimized = confusion_matrix(\n    y_true,\n    y_pred_optimized,\n    labels=[0, 1],\n)\n\ntn, fp, fn, tp = cm_optimized.ravel()\n\nsensitivity = (\n    tp / max(tp + fn, 1)\n)\n\nspecificity = (\n    tn / max(tn + fp, 1)\n)\n\nprecision = (\n    tp / max(tp + fp, 1)\n)\n\nnegative_predictive_value = (\n    tn / max(tn + fn, 1)\n)\n\n\n# ============================================================\n# Metrics with fixed threshold 0.5\n# ============================================================\n\naccuracy_05 = accuracy_score(\n    y_true,\n    y_pred_05,\n)\n\nbalanced_accuracy_05 = (\n    balanced_accuracy_score(\n        y_true,\n        y_pred_05,\n    )\n)\n\nf1_05 = f1_score(\n    y_true,\n    y_pred_05,\n    zero_division=0,\n)\n\nmcc_05 = matthews_corrcoef(\n    y_true,\n    y_pred_05,\n)\n\n\n# ============================================================\n# Print results\n# ============================================================\n\nevaluation_label = (\n    \"COMPLETE OOF RESULTS\"\n    if is_complete_oof\n    else \"PARTIAL OOF RESULTS\"\n)\n\nprint(\"\\n\" + \"=\" * 80)\nprint(evaluation_label)\nprint(\"=\" * 80)\n\nprint(\n    f\"Evaluated folds:              \"\n    f\"{FOLDS_TO_EVALUATE}\"\n)\n\nprint(\n    f\"Evaluated patients:           \"\n    f\"{evaluated_patients}\"\n)\n\nprint(\n    f\"ROC-AUC:                      \"\n    f\"{oof_auc:.4f}\"\n)\n\nprint(\n    f\"PR-AUC:                       \"\n    f\"{oof_pr_auc:.4f}\"\n)\n\nprint(\n    f\"Optimized threshold:          \"\n    f\"{optimized_threshold:.4f}\"\n)\n\nprint(\n    f\"Accuracy (optimized):         \"\n    f\"{optimized_accuracy:.4f}\"\n)\n\nprint(\n    f\"Balanced accuracy (optimized): \"\n    f\"{optimized_balanced_accuracy:.4f}\"\n)\n\nprint(\n    f\"F1-score (optimized):         \"\n    f\"{optimized_f1:.4f}\"\n)\n\nprint(\n    f\"MCC (optimized):              \"\n    f\"{optimized_mcc:.4f}\"\n)\n\nprint(\n    f\"Sensitivity:                  \"\n    f\"{sensitivity:.4f}\"\n)\n\nprint(\n    f\"Specificity:                  \"\n    f\"{specificity:.4f}\"\n)\n\nprint(\n    f\"Precision:                    \"\n    f\"{precision:.4f}\"\n)\n\nprint(\n    f\"Negative predictive value:    \"\n    f\"{negative_predictive_value:.4f}\"\n)\n\nprint(\"\\nMetrics at fixed threshold 0.5:\")\n\nprint(\n    f\"Accuracy at 0.5:              \"\n    f\"{accuracy_05:.4f}\"\n)\n\nprint(\n    f\"Balanced accuracy at 0.5:     \"\n    f\"{balanced_accuracy_05:.4f}\"\n)\n\nprint(\n    f\"F1-score at 0.5:              \"\n    f\"{f1_05:.4f}\"\n)\n\nprint(\n    f\"MCC at 0.5:                   \"\n    f\"{mcc_05:.4f}\"\n)\n\nprint(\"\\nConfusion matrix values:\")\nprint(f\"True negatives:  {tn}\")\nprint(f\"False positives: {fp}\")\nprint(f\"False negatives: {fn}\")\nprint(f\"True positives:  {tp}\")\n\nprint(\"\\nClassification report:\\n\")\n\nprint(\n    classification_report(\n        y_true,\n        y_pred_optimized,\n        target_names=[\n            \"MGMT Negative\",\n            \"MGMT Positive\",\n        ],\n        digits=4,\n        zero_division=0,\n    )\n)\n\n\n# ============================================================\n# Mean ± standard deviation across selected folds\n# ============================================================\n\nmetric_columns = [\n    \"roc_auc\",\n    \"pr_auc\",\n    \"accuracy_0.5\",\n    \"balanced_accuracy_0.5\",\n    \"f1_0.5\",\n    \"mcc_0.5\",\n    \"sensitivity_0.5\",\n    \"specificity_0.5\",\n]\n\nsummary_rows = []\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"MEAN ± STANDARD DEVIATION — SELECTED FOLDS\")\nprint(\"=\" * 80)\n\nfor column in metric_columns:\n    if column not in fold_results_df.columns:\n        continue\n\n    mean_value = (\n        fold_results_df[column].mean()\n    )\n\n    if len(fold_results_df) > 1:\n        std_value = (\n            fold_results_df[column].std()\n        )\n    else:\n        std_value = np.nan\n\n    summary_rows.append({\n        \"metric\": column,\n        \"mean\": mean_value,\n        \"std\": std_value,\n    })\n\n    if np.isnan(std_value):\n        print(\n            f\"{column:28s}: \"\n            f\"{mean_value:.4f} ± N/A\"\n        )\n    else:\n        print(\n            f\"{column:28s}: \"\n            f\"{mean_value:.4f} ± {std_value:.4f}\"\n        )\n\nfold_summary_df = pd.DataFrame(\n    summary_rows\n)\n\n\n# ============================================================\n# Output filename suffix\n# ============================================================\n\nfold_suffix = \"_\".join(\n    str(fold)\n    for fold in FOLDS_TO_EVALUATE\n)\n\nevaluation_prefix = (\n    \"complete\"\n    if is_complete_oof\n    else f\"partial_folds_{fold_suffix}\"\n)\n\n\n# ============================================================\n# Confusion matrix\n# ============================================================\n\nfig, ax = plt.subplots(\n    figsize=(7, 6)\n)\n\ndisplay_cm = ConfusionMatrixDisplay(\n    confusion_matrix=cm_optimized,\n    display_labels=[\n        \"MGMT Negative\",\n        \"MGMT Positive\",\n    ],\n)\n\ndisplay_cm.plot(\n    cmap=\"Blues\",\n    values_format=\"d\",\n    ax=ax,\n    colorbar=True,\n)\n\nax.set_title(\n    (\n        \"Complete OOF Confusion Matrix\"\n        if is_complete_oof\n        else (\n            \"Partial OOF Confusion Matrix\\n\"\n            f\"Folds {FOLDS_TO_EVALUATE}\"\n        )\n    ),\n    fontsize=14,\n)\n\nax.grid(False)\n\nplt.tight_layout()\n\nconfusion_matrix_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_confusion_matrix.png\"\n)\n\nplt.savefig(\n    confusion_matrix_path,\n    dpi=300,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\n\n# ============================================================\n# ROC curve\n# ============================================================\n\nfpr, tpr, roc_thresholds = roc_curve(\n    y_true,\n    y_prob,\n)\n\nplt.figure(\n    figsize=(7, 6)\n)\n\nplt.plot(\n    fpr,\n    tpr,\n    linewidth=2,\n    label=f\"ROC-AUC = {oof_auc:.4f}\",\n)\n\nplt.plot(\n    [0, 1],\n    [0, 1],\n    linestyle=\"--\",\n    linewidth=1.5,\n    label=\"Random classifier\",\n)\n\nplt.xlabel(\n    \"False Positive Rate\"\n)\n\nplt.ylabel(\n    \"True Positive Rate\"\n)\n\nplt.title(\n    (\n        \"Complete Out-of-Fold ROC Curve\"\n        if is_complete_oof\n        else (\n            \"Partial Out-of-Fold ROC Curve\\n\"\n            f\"Folds {FOLDS_TO_EVALUATE}\"\n        )\n    )\n)\n\nplt.legend(\n    loc=\"lower right\"\n)\n\nplt.grid(\n    alpha=0.3\n)\n\nplt.tight_layout()\n\nroc_curve_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_roc_curve.png\"\n)\n\nplt.savefig(\n    roc_curve_path,\n    dpi=300,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\n\n# ============================================================\n# Precision–Recall curve\n# ============================================================\n\nprecision_curve, recall_curve, pr_thresholds = (\n    precision_recall_curve(\n        y_true,\n        y_prob,\n    )\n)\n\npositive_prevalence = float(\n    y_true.mean()\n)\n\nplt.figure(\n    figsize=(7, 6)\n)\n\nplt.plot(\n    recall_curve,\n    precision_curve,\n    linewidth=2,\n    label=f\"PR-AUC = {oof_pr_auc:.4f}\",\n)\n\nplt.axhline(\n    y=positive_prevalence,\n    linestyle=\"--\",\n    linewidth=1.5,\n    label=(\n        f\"Positive prevalence = \"\n        f\"{positive_prevalence:.4f}\"\n    ),\n)\n\nplt.xlabel(\n    \"Recall\"\n)\n\nplt.ylabel(\n    \"Precision\"\n)\n\nplt.title(\n    (\n        \"Complete OOF Precision–Recall Curve\"\n        if is_complete_oof\n        else (\n            \"Partial OOF Precision–Recall Curve\\n\"\n            f\"Folds {FOLDS_TO_EVALUATE}\"\n        )\n    )\n)\n\nplt.legend(\n    loc=\"lower left\"\n)\n\nplt.grid(\n    alpha=0.3\n)\n\nplt.tight_layout()\n\npr_curve_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_precision_recall_curve.png\"\n)\n\nplt.savefig(\n    pr_curve_path,\n    dpi=300,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\n\n# ============================================================\n# Fold-by-fold ROC-AUC plot\n# ============================================================\n\nplt.figure(\n    figsize=(8, 5)\n)\n\nplt.bar(\n    fold_results_df[\"fold\"].astype(str),\n    fold_results_df[\"roc_auc\"],\n)\n\nplt.axhline(\n    y=fold_results_df[\"roc_auc\"].mean(),\n    linestyle=\"--\",\n    linewidth=1.5,\n    label=(\n        \"Mean ROC-AUC = \"\n        f\"{fold_results_df['roc_auc'].mean():.4f}\"\n    ),\n)\n\nplt.xlabel(\n    \"Fold\"\n)\n\nplt.ylabel(\n    \"ROC-AUC\"\n)\n\nplt.title(\n    \"ROC-AUC by Selected Cross-Validation Fold\"\n)\n\nplt.ylim(\n    0.0,\n    1.0,\n)\n\nplt.legend()\n\nplt.grid(\n    axis=\"y\",\n    alpha=0.3,\n)\n\nplt.tight_layout()\n\nfold_auc_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_fold_roc_auc.png\"\n)\n\nplt.savefig(\n    fold_auc_path,\n    dpi=300,\n    bbox_inches=\"tight\",\n)\n\nplt.show()\n\n\n# ============================================================\n# Save prediction and metric tables\n# ============================================================\n\noof_df[\"prediction_0.5\"] = (\n    y_pred_05\n)\n\noof_df[\"prediction_optimized\"] = (\n    y_pred_optimized\n)\n\noof_df[\"optimized_threshold\"] = (\n    optimized_threshold\n)\n\noof_df[\"evaluation_is_complete\"] = (\n    is_complete_oof\n)\n\noof_df[\"evaluated_folds\"] = (\n    \",\".join(\n        str(fold)\n        for fold in FOLDS_TO_EVALUATE\n    )\n)\n\nfinal_metrics_df = pd.DataFrame([{\n    \"evaluation_type\": (\n        \"complete\"\n        if is_complete_oof\n        else \"partial\"\n    ),\n    \"evaluated_folds\": \",\".join(\n        str(fold)\n        for fold in FOLDS_TO_EVALUATE\n    ),\n    \"n_evaluated_folds\": len(\n        FOLDS_TO_EVALUATE\n    ),\n    \"n_evaluated_patients\": (\n        evaluated_patients\n    ),\n    \"total_patients\": total_patients,\n    \"coverage\": coverage,\n    \"roc_auc\": oof_auc,\n    \"pr_auc\": oof_pr_auc,\n    \"accuracy_optimized\": (\n        optimized_accuracy\n    ),\n    \"balanced_accuracy_optimized\": (\n        optimized_balanced_accuracy\n    ),\n    \"f1_score_optimized\": optimized_f1,\n    \"mcc_optimized\": optimized_mcc,\n    \"accuracy_0.5\": accuracy_05,\n    \"balanced_accuracy_0.5\": (\n        balanced_accuracy_05\n    ),\n    \"f1_score_0.5\": f1_05,\n    \"mcc_0.5\": mcc_05,\n    \"sensitivity\": sensitivity,\n    \"specificity\": specificity,\n    \"precision\": precision,\n    \"negative_predictive_value\": (\n        negative_predictive_value\n    ),\n    \"optimized_threshold\": (\n        optimized_threshold\n    ),\n    \"true_negatives\": tn,\n    \"false_positives\": fp,\n    \"false_negatives\": fn,\n    \"true_positives\": tp,\n}])\n\n\noof_predictions_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_oof_predictions.csv\"\n)\n\nfold_results_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_fold_results.csv\"\n)\n\nfold_summary_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_fold_summary.csv\"\n)\n\nfinal_metrics_path = (\n    FINAL_OUTPUT_DIR\n    / f\"{evaluation_prefix}_metrics.csv\"\n)\n\n\noof_df.to_csv(\n    oof_predictions_path,\n    index=False,\n)\n\nfold_results_df.to_csv(\n    fold_results_path,\n    index=False,\n)\n\nfold_summary_df.to_csv(\n    fold_summary_path,\n    index=False,\n)\n\nfinal_metrics_df.to_csv(\n    final_metrics_path,\n    index=False,\n)\n\n\n# ============================================================\n# Final saved files report\n# ============================================================\n\nprint(\"\\n\" + \"=\" * 80)\nprint(\"FILES SAVED\")\nprint(\"=\" * 80)\n\nprint(oof_predictions_path)\nprint(fold_results_path)\nprint(fold_summary_path)\nprint(final_metrics_path)\nprint(confusion_matrix_path)\nprint(roc_curve_path)\nprint(pr_curve_path)\nprint(fold_auc_path)\n\nif not is_complete_oof:\n    print(\n        \"\\nImportant: these are preliminary results based only \"\n        f\"on folds {FOLDS_TO_EVALUATE}. Run all {N_FOLDS} folds \"\n        \"before reporting the final model performance.\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T18:58:01.346700Z","iopub.execute_input":"2026-07-11T18:58:01.347138Z","iopub.status.idle":"2026-07-11T18:58:03.587652Z","shell.execute_reply.started":"2026-07-11T18:58:01.347114Z","shell.execute_reply":"2026-07-11T18:58:03.586918Z"}},"outputs":[],"execution_count":null}]}