{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":29653,"databundleVersionId":2420395,"isSourceIdPinned":false}],"dockerImageVersionId":31400,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"![brain_baner](http://www.mf-data-science.fr/images/projects/brain_baner.jpg)","metadata":{}},{"cell_type":"markdown","source":"<h1 style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\">Context</h1>\n\nThe goal of this competition, initiated by the **Radiological Society of North America *(RSNA)*** in partnership with the **Medical Image Computing and Computer Assisted Intervention Society *(the MICCAI Society)*** is to predict the methylation of the **MGMT promoter**, which is an important gene biomarker for treatment of brain tumors.\n\nThese predictions will be based on a database of **MRI *(magnetic resonance imaging)*** scans of several hundred patients.\n\n<h1 style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\">Data</h1>\n\nEach independent case has a dedicated folder identified by a five-digit number. Within each of these “case” folders, there are four sub-folders, each of them corresponding to each of the structural multi-parametric MRI (mpMRI) scans, in DICOM format. The exact mpMRI scans included are:\n\n- Fluid Attenuated Inversion Recovery (FLAIR)\n- T1-weighted pre-contrast (T1w)\n- T1-weighted post-contrast (T1Gd)\n- T2-weighted (T2)\n\n| ![brain_baner](http://www.mf-data-science.fr/images/projects/brain_tumor_types.png) | \n|:--:| \n| *Examples of the four MR sequence types included in this work* |\n\n<h1 style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\">Acknowledgement</h1>\n\nThis Notebook is inspired from *Ammar Alhaj Ali* work :\n- [🧠Brain Tumor 3D [Training]](https://www.kaggle.com/ammarnassanalhajali/brain-tumor-3d-training)\n- [🧠Brain Tumor 3D [Inference]](https://www.kaggle.com/ammarnassanalhajali/brain-tumor-3d-inference)","metadata":{}},{"cell_type":"markdown","source":"<h1 style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\">Summary</h1>\n\n1. [Exploratory data analysis (EDA)](#section_1)      \n    1.1. [Submission sample & train.csv](#section_1_1)      \n    1.2. [MRI train data](#section_1_2)      \n    1.3. [Data cleaning](#section_1_3)      \n\n2. [Preprocessing](#section_2)      \n    2.1. [Crop and resize the images](#section_2_1)      \n    2.2. [Equalization CLAHE](#section_2_2)      \n    2.3. [Denoising filter](#section_2_3)      \n    2.4. [Global preprocessing function](#section_2_4)      \n    \n3. [Development of supervised models](#section_3)      \n    3.1. [Multimodal inputs CNN from scratch](#section_3_1)      \n    3.2. [Define loaders for images sequences 4 MRI types](#section_3_2)      \n    3.3. [Define folds](#section_3_3)      \n    3.4. [Keras custom data generator](#section_3_4)      \n    3.5. [Define CNN Multi-inputs model](#section_3_5)      \n    \n4. [Test of trained final model](#section_4)     \n5. [Try another approach: Transfer Learning](#section_5)","metadata":{}},{"cell_type":"markdown","source":"# <span style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\" id=\"section_1\">Exploratory data analysis (EDA)</span>","metadata":{}},{"cell_type":"markdown","source":"First, we have to load the usefull Python libraries :","metadata":{"execution":{"iopub.status.busy":"2021-08-04T14:37:44.731175Z","iopub.execute_input":"2021-08-04T14:37:44.731595Z","iopub.status.idle":"2021-08-04T14:37:44.739822Z","shell.execute_reply.started":"2021-08-04T14:37:44.731563Z","shell.execute_reply":"2021-08-04T14:37:44.737907Z"}}},{"cell_type":"code","source":"# Import Python libraries\nimport os\nimport glob\nimport re\nimport math\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom matplotlib import animation, rc\nimport seaborn as sns\nfrom tqdm.notebook import tqdm\nimport pydicom as dicom\nimport cv2\nfrom PIL import Image\nimport gc\n\n# TensorFlow and Keras imports\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras.layers import (\n    Input, Conv3D, MaxPool3D, BatchNormalization, \n    Dense, Dropout, GlobalAveragePooling3D, concatenate, Activation\n)\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.utils import plot_model, Sequence\nfrom tensorflow.keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau\nfrom tensorflow.keras import mixed_precision\nfrom sklearn.model_selection import StratifiedKFold, train_test_split\n\nprint(f\"TensorFlow version: {tf.__version__}\")\nprint(f\"GPU Available: {tf.config.list_physical_devices('GPU')}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:05.594956Z","iopub.execute_input":"2026-07-05T08:54:05.595236Z","iopub.status.idle":"2026-07-05T08:54:23.894360Z","shell.execute_reply.started":"2026-07-05T08:54:05.595211Z","shell.execute_reply":"2026-07-05T08:54:23.893602Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Configure GPU memory growth to avoid OOM errors\ngpus = tf.config.experimental.list_physical_devices('GPU')\nif gpus:\n    try:\n        for gpu in gpus:\n            tf.config.experimental.set_memory_growth(gpu, True)\n        print(\"GPU memory growth enabled.\")\n    except RuntimeError as e:\n        print(e)\nelse:\n    print(\"No GPU found — running on CPU.\")\n\n# Mixed precision disabled: float16 breaks MaxPool3D on CPU.\n# The main speedup comes from the .npy cache, not from mixed precision.\nfrom tensorflow.keras import mixed_precision\nmixed_precision.set_global_policy('float32')\nprint(f\"Compute policy: {mixed_precision.global_policy().name}\")\n\nnp.random.seed(42)\ntf.random.set_seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:23.895996Z","iopub.execute_input":"2026-07-05T08:54:23.896525Z","iopub.status.idle":"2026-07-05T08:54:23.904143Z","shell.execute_reply.started":"2026-07-05T08:54:23.896500Z","shell.execute_reply":"2026-07-05T08:54:23.903212Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_1_1\">Submission sample & train.csv</span>\n\nWe will already look at the exemple **submission file** to see exactly what we need to predict in the end.","metadata":{}},{"cell_type":"code","source":"# Define dataset paths\ninput_path = \"/kaggle/input/competitions/rsna-miccai-brain-tumor-radiogenomic-classification\"\ntrain_labels_file = \"train_labels.csv\"\n\n# Load training labels from the correct file\ntrain_labels = pd.read_csv(os.path.join(input_path, train_labels_file))\n\n# Display basic info\nprint(f\"Total training patients loaded: {len(train_labels)}\")\nprint(f\"Missing MGMT values: {train_labels['MGMT_value'].isnull().sum()}\")\nprint(\"\\nFirst 5 records:\")\ntrain_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:23.905404Z","iopub.execute_input":"2026-07-05T08:54:23.905767Z","iopub.status.idle":"2026-07-05T08:54:23.960545Z","shell.execute_reply.started":"2026-07-05T08:54:23.905728Z","shell.execute_reply":"2026-07-05T08:54:23.959812Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We can see that the data structure is the same as for the submission file, knowing that here **MGMT_value is indeed equal to 0 or 1 and no longer a probability**.\n\nLet's look at the distribution of the values of this variable in the train set :","metadata":{}},{"cell_type":"code","source":"# Visualize MGMT value distribution\nsns.set_style(\"whitegrid\")\nplt.figure(figsize=(8, 6))\nax = sns.countplot(data=train_labels, x=\"MGMT_value\", palette=\"viridis\")\n\n# Annotate bars with counts\nfor p in ax.patches:\n    ax.annotate(\n        format(p.get_height(), '.0f'), \n        (p.get_x() + p.get_width() / 2., p.get_height()),\n        ha='center', va='center', \n        xytext=(0, 10), \n        textcoords='offset points', fontsize=10\n    )\n\nplt.title(\"MGMT Promoter Methylation Status Distribution\", fontsize=14, fontweight='bold')\nplt.xlabel(\"MGMT Value (0 = Unmethylated, 1 = Methylated)\")\nplt.ylabel(\"Count\")\nsns.despine()\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:23.962139Z","iopub.execute_input":"2026-07-05T08:54:23.962459Z","iopub.status.idle":"2026-07-05T08:54:24.212901Z","shell.execute_reply.started":"2026-07-05T08:54:23.962436Z","shell.execute_reply":"2026-07-05T08:54:24.212230Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The distribution is almost equal between class 0 and class 1. The presence of MGMT is slightly higher. A total of **585 patients are tested in the train set**.","metadata":{}},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_1_2\">MRI train data</span>\n\nTrain data contains one record per patient. For each patient, four sub-files are available *(FLAIR, T1w, T1wCE and T2w)* in which the MRI image sequences are distributed.\n\n![data_structure](http://www.mf-data-science.fr/images/projects/data_structure.jpg)\n\nWe are going to take a look at what an MRI image looks like :","metadata":{}},{"cell_type":"code","source":"# Function to display sample MRI slices from a patient\ndef plot_examples(patient_row=0, mri_type='FLAIR', num_to_plot=5):\n    \"\"\"\n    Display dynamically sampled valid MRI slices for a given patient.\n    \n    Parameters:\n    - patient_row: int, row index in train_labels DataFrame\n    - mri_type: str, one of ['FLAIR', 'T1w', 'T1wCE', 'T2w']\n    - num_to_plot: int, number of slices to display\n    \"\"\"\n    patient_id = str(train_labels.loc[patient_row, 'BraTS21ID']).zfill(5)\n    folder_path = os.path.join(input_path, 'train', patient_id, mri_type)\n    \n    if not os.path.exists(folder_path):\n        print(f\"Warning: Folder not found for patient {patient_id}\")\n        return\n    \n    # Get and sort DICOM files numerically\n    all_files = sorted(\n        glob.glob(os.path.join(folder_path, \"*.dcm\")),\n        key=lambda var: [int(x) if x.isdigit() else x for x in re.split('([0-9]+)', var)]\n    )\n    \n    # Filter valid (non-black) images\n    valid_images = []\n    valid_names = []\n    for f in all_files:\n        try:\n            dcm = dicom.dcmread(f)\n            img = dcm.pixel_array\n            if np.max(img) > 0 and np.mean(img) > 1.0:  # Filter black/noisy slices\n                valid_images.append(img)\n                valid_names.append(os.path.basename(f))\n        except:\n            continue\n    \n    if len(valid_images) == 0:\n        print(f\"No valid images found for {patient_id}\")\n        return\n    \n    # Dynamic uniform sampling\n    total = len(valid_images)\n    if total >= num_to_plot:\n        step = total / num_to_plot\n        indices = [int(i * step) for i in range(num_to_plot)]\n        indices[-1] = min(indices[-1], total - 1)  # Safety check\n        display_imgs = [valid_images[i] for i in indices]\n        display_names = [valid_names[i] for i in indices]\n    else:\n        display_imgs = valid_images\n        display_names = valid_names\n    \n    # Plot\n    fig, axes = plt.subplots(1, len(display_imgs), figsize=(20, 6))\n    if len(display_imgs) == 1:\n        axes = [axes]\n    \n    for idx, (img, name) in enumerate(zip(display_imgs, display_names)):\n        axes[idx].imshow(img, cmap='gray')\n        axes[idx].set_title(f\"{mri_type}\\n{name}\", fontsize=9)\n        axes[idx].axis('off')\n    \n    mgmt_val = train_labels.loc[patient_row, 'MGMT_value']\n    plt.suptitle(f\"Patient: {patient_id} | MGMT: {mgmt_val} | {mri_type} Scans\", \n                 fontsize=14, fontweight='bold', y=1.02)\n    plt.tight_layout()\n    plt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:24.213830Z","iopub.execute_input":"2026-07-05T08:54:24.214203Z","iopub.status.idle":"2026-07-05T08:54:24.224769Z","shell.execute_reply.started":"2026-07-05T08:54:24.214179Z","shell.execute_reply":"2026-07-05T08:54:24.224102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example of FLAIR scans\nplot_examples(0, 'FLAIR')\n\n# Example of FLAIR scans\nplot_examples(3, 'FLAIR')\n\n# Example of T1wCE scans\nplot_examples(12, 'T1wCE')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:24.225932Z","iopub.execute_input":"2026-07-05T08:54:24.226353Z","iopub.status.idle":"2026-07-05T08:54:37.732081Z","shell.execute_reply.started":"2026-07-05T08:54:24.226301Z","shell.execute_reply":"2026-07-05T08:54:37.731102Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"To better understand the representation of these MRI images, we can also **create an animation to visualize the sequence of images** of a certain category for a given patient.","metadata":{}},{"cell_type":"code","source":"# Function to create MRI animation\ndef create_animation(row=0, cat='FLAIR'):\n    \"\"\"\n    Returns an animation of MRI images for a given patient and modality.\n    \"\"\"\n    folder = str(train_labels.loc[row, 'BraTS21ID']).zfill(5)\n    path_file = os.path.join(input_path, 'train', folder, cat)\n    \n    t_paths = sorted(\n        glob.glob(os.path.join(path_file, \"*\")), \n        key=lambda x: int(x[:-4].split(\"-\")[-1]),\n    )\n    images = []\n    for filename in t_paths:\n        try:\n            data_file = dicom.dcmread(filename)\n            data = data_file.pixel_array\n            data = data - np.min(data)\n            if np.max(data) != 0:\n                data = data / np.max(data)\n            data = (data * 255).astype(np.uint8)\n            if data.max() == 0:\n                continue\n            images.append(data)\n        except:\n            continue\n    \n    if not images:\n        print(\"No valid images for animation.\")\n        return None\n\n    fig = plt.figure(figsize=(8, 8))\n    plt.axis('off')\n    im = plt.imshow(images[0], cmap=\"gray\", animated=True)\n\n    def animate_func(i):\n        im.set_array(images[i])\n        return [im]\n\n    animated = animation.FuncAnimation(\n        fig, animate_func, frames=len(images), interval=1000//24)\n    return animated","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:37.733401Z","iopub.execute_input":"2026-07-05T08:54:37.733995Z","iopub.status.idle":"2026-07-05T08:54:37.744257Z","shell.execute_reply.started":"2026-07-05T08:54:37.733923Z","shell.execute_reply":"2026-07-05T08:54:37.743513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Generate animation for a sample patient\nani_1 = create_animation(row=3, cat='FLAIR')","metadata":{"_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:37.745171Z","iopub.execute_input":"2026-07-05T08:54:37.745563Z","iopub.status.idle":"2026-07-05T08:54:38.912199Z","shell.execute_reply.started":"2026-07-05T08:54:37.745525Z","shell.execute_reply":"2026-07-05T08:54:38.911445Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display the animation\nani_1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:38.913265Z","iopub.execute_input":"2026-07-05T08:54:38.913644Z","iopub.status.idle":"2026-07-05T08:54:38.919218Z","shell.execute_reply.started":"2026-07-05T08:54:38.913613Z","shell.execute_reply":"2026-07-05T08:54:38.918260Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We are now going to check the number of DCM files to check if their number is the same for each category and for each patient. For that, we will complete a copy of train.csv with the calculated informations:","metadata":{}},{"cell_type":"code","source":"# Add slice count features for training set\nscan_categories = [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]\ntrain_dataset = train_labels.copy()\n\nprint(\"Counting DICOM files per patient and modality...\")\nfor scan in scan_categories:\n    counts = []\n    for pid in tqdm(train_dataset.BraTS21ID, desc=f\"Counting {scan}\"):\n        folder = os.path.join(input_path, \"train\", str(pid).zfill(5), scan)\n        if os.path.exists(folder):\n            counts.append(len(os.listdir(folder)))\n        else:\n            counts.append(0)  # Handle missing folders gracefully\n    train_dataset[f\"{scan}_count\"] = counts\n\nprint(\"Slice count features added.\")\nprint(train_dataset[[\"BraTS21ID\"] + [f\"{s}_count\" for s in scan_categories]].head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:54:38.923752Z","iopub.execute_input":"2026-07-05T08:54:38.924032Z","iopub.status.idle":"2026-07-05T08:55:29.085821Z","shell.execute_reply.started":"2026-07-05T08:54:38.924011Z","shell.execute_reply":"2026-07-05T08:55:29.084671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize slice count distributions for training set\nfig = plt.figure(figsize=(20, 12))\nfor i, scan in enumerate(scan_categories, 1):\n    ax = plt.subplot(2, 2, i)\n    sns.countplot(data=train_dataset, x=f\"{scan}_count\", ax=ax, color='steelblue')\n    ax.set_title(f\"Distribution of {scan} Slice Counts (Train)\", fontsize=12, fontweight='bold')\n    ax.set_xlabel(\"Number of Slices\")\n    ax.set_ylabel(\"Frequency\")\n    plt.xticks(rotation=45)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:29.087034Z","iopub.execute_input":"2026-07-05T08:55:29.087779Z","iopub.status.idle":"2026-07-05T08:55:31.488284Z","shell.execute_reply.started":"2026-07-05T08:55:29.087750Z","shell.execute_reply":"2026-07-05T08:55:31.487345Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Note that some values for each scan category are over-represented. On the other hand, the span ranges of the counters are important. This may be due, for example, to the use of different X-ray machines ...\n\nHowever, if we only consider patients with maximum values, the amount of data available may be too low for a complex machine learning algorithm.","metadata":{}},{"cell_type":"code","source":"# [Optional] Filter patients with most common slice counts across all modalities\n# Note: This is a strict filter and may return very few samples. Use with caution.\n\nmode_filters = {\n    scan: int(train_dataset[f\"{scan}_count\"].mode()[0]) \n    for scan in scan_categories\n}\n\nprint(\"Mode slice counts per modality:\", mode_filters)\n\n# Apply filter (commented out by default - enable if needed)\n# filtered_df = train_dataset[\n#     (train_dataset[\"FLAIR_count\"] == mode_filters[\"FLAIR\"]) &\n#     (train_dataset[\"T1w_count\"] == mode_filters[\"T1w\"]) &\n#     (train_dataset[\"T1wCE_count\"] == mode_filters[\"T1wCE\"]) &\n#     (train_dataset[\"T2w_count\"] == mode_filters[\"T2w\"])\n# ]\n# print(f\"Patients matching all mode counts: {len(filtered_df)}\")\n# filtered_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:31.489419Z","iopub.execute_input":"2026-07-05T08:55:31.489797Z","iopub.status.idle":"2026-07-05T08:55:31.496416Z","shell.execute_reply.started":"2026-07-05T08:55:31.489761Z","shell.execute_reply":"2026-07-05T08:55:31.495644Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"There are only 151 patients whose scans contain the maximum of DCM images out of the 585 at the start. We therefore keep the entire dataset for the moment. Lets have a look to the **test dataset** :","metadata":{}},{"cell_type":"code","source":"# Build test dataset by scanning test directory dynamically\ntest_dataset = pd.DataFrame(columns=[\"BraTS21ID\"] + [f\"{s}_count\" for s in scan_categories])\n\ntest_patients = [d for d in os.listdir(os.path.join(input_path, \"test\")) \n                 if os.path.isdir(os.path.join(input_path, \"test\", d))]\ntest_dataset[\"BraTS21ID\"] = [int(p) for p in test_patients]\n\nprint(f\"Found {len(test_dataset)} patients in test set.\")\n\nfor scan in scan_categories:\n    counts = []\n    for pid in tqdm(test_dataset.BraTS21ID, desc=f\"Counting test {scan}\"):\n        folder = os.path.join(input_path, \"test\", str(pid).zfill(5), scan)\n        if os.path.exists(folder):\n            counts.append(len(os.listdir(folder)))\n        else:\n            counts.append(0)\n    test_dataset[f\"{scan}_count\"] = counts\n\nprint(\"Test dataset ready.\")\ntest_dataset.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:31.497448Z","iopub.execute_input":"2026-07-05T08:55:31.497764Z","iopub.status.idle":"2026-07-05T08:55:39.675162Z","shell.execute_reply.started":"2026-07-05T08:55:31.497731Z","shell.execute_reply":"2026-07-05T08:55:39.674190Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Preview test dataset\ntest_dataset.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:39.676335Z","iopub.execute_input":"2026-07-05T08:55:39.676898Z","iopub.status.idle":"2026-07-05T08:55:39.684948Z","shell.execute_reply.started":"2026-07-05T08:55:39.676870Z","shell.execute_reply":"2026-07-05T08:55:39.684133Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize slice count distributions for test set (Combined with Cell 29)\nfig = plt.figure(figsize=(20, 12))\nfor i, scan in enumerate(scan_categories, 1):\n    ax = plt.subplot(2, 2, i)\n    sns.countplot(data=test_dataset, x=f\"{scan}_count\", ax=ax, color='coral')\n    ax.set_title(f\"Distribution of {scan} Slice Counts (Test)\", fontsize=12, fontweight='bold')\n    ax.set_xlabel(\"Number of Slices\")\n    ax.set_ylabel(\"Frequency\")\n    plt.xticks(rotation=45)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:39.685868Z","iopub.execute_input":"2026-07-05T08:55:39.686284Z","iopub.status.idle":"2026-07-05T08:55:40.873801Z","shell.execute_reply.started":"2026-07-05T08:55:39.686242Z","shell.execute_reply":"2026-07-05T08:55:40.873157Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_1_3\">Data cleaning</span>\n\nIn the animation of the scans projected above, we notice that a certain number of images at the beginning or at the end of the sequence has a lot of black area.\n\nThese images will therefore be useless in our models and may even cause over-training.\n**When creating the image sequences, we will therefore start from the central image of each folder and we will then take the same number of images upstream and downstream**.\n\nLet's take an example from a single image :","metadata":{}},{"cell_type":"code","source":"# Analyze a sample image: percentage of non-zero pixels\nsample_pid = \"00005\"\nsample_img_path = os.path.join(input_path, \"train\", sample_pid, \"FLAIR\", \"Image-80.dcm\")\n\nif os.path.exists(sample_img_path):\n    sample_img = dicom.dcmread(sample_img_path).pixel_array\n    non_zero_pct = np.sum(sample_img != 0) / (sample_img.shape[0] * sample_img.shape[1]) * 100\n    \n    plt.figure(figsize=(6, 6))\n    plt.imshow(sample_img, cmap=\"gray\")\n    plt.title(f\"Sample Image\\nNon-zero pixels: {non_zero_pct:.2f}%\", fontsize=12, fontweight='bold')\n    plt.axis(\"off\")\n    plt.show()\n    print(f\"Image shape: {sample_img.shape}\")\n    print(f\"Non-zero pixel percentage: {non_zero_pct:.2f}%\")\nelse:\n    print(f\"Sample file not found: {sample_img_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:40.874769Z","iopub.execute_input":"2026-07-05T08:55:40.875080Z","iopub.status.idle":"2026-07-05T08:55:40.999230Z","shell.execute_reply.started":"2026-07-05T08:55:40.875055Z","shell.execute_reply":"2026-07-05T08:55:40.998535Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"In this example, only 7% of the image has colored pixels.      \nWe will now look at the distribution of the **colorization rates of the images in the full FLAIR folder** :","metadata":{}},{"cell_type":"code","source":"# Analyze non-zero pixel distribution across all FLAIR slices for one patient\npatient_id = \"00005\"\nfolder_path = os.path.join(input_path, \"train\", patient_id, \"FLAIR\")\n\nnon_zero_ratios = []\nif os.path.exists(folder_path):\n    for fname in tqdm(os.listdir(folder_path), desc=f\"Analyzing {patient_id}\"):\n        if fname.endswith(\".dcm\"):\n            try:\n                img = dicom.dcmread(os.path.join(folder_path, fname)).pixel_array\n                ratio = np.sum(img != 0) / (img.shape[0] * img.shape[1])\n                non_zero_ratios.append(round(ratio, 2))\n            except:\n                continue\n    \n    if non_zero_ratios:\n        plt.figure(figsize=(10, 5))\n        sns.histplot(non_zero_ratios, bins=20, kde=True, color='teal')\n        plt.xlabel(\"Ratio of Non-Zero Pixels\")\n        plt.ylabel(\"Frequency\")\n        plt.title(f\"Non-Zero Pixel Ratio Distribution - Patient {patient_id} (FLAIR)\", fontsize=12, fontweight='bold')\n        plt.tight_layout()\n        plt.show()\n        print(f\"Analyzed {len(non_zero_ratios)} slices.\")\n    else:\n        print(\"No valid slices found.\")\nelse:\n    print(f\"Folder not found: {folder_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:41.000247Z","iopub.execute_input":"2026-07-05T08:55:41.000619Z","iopub.status.idle":"2026-07-05T08:55:42.019543Z","shell.execute_reply.started":"2026-07-05T08:55:41.000549Z","shell.execute_reply":"2026-07-05T08:55:42.018855Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"It can be seen that the majority of the images are completely black.\n\n# <span style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\" id=\"section_2\">Preprocessing</span>\nWe will be using several preprocessing techniques on our images.\n- Crop images to reduce black areas.\n- Resize images.\n- Application of a denoising filter.\n\n**We will not apply image equalization** as the different types of scans already have voluntary contrast variations.\n\n## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_2_1\">Crop and resize the images</span>","metadata":{}},{"cell_type":"code","source":"# Function to crop and resize a 2D image\ndef crop_resize_img(img, scale=1.0, dim=(244, 244)):\n    \"\"\"\n    Crop central region and resize to target dimensions.\n    \n    Parameters:\n    - img: 2D numpy array\n    - scale: float, crop scale factor (1.0 = full image)\n    - dim: tuple, target (width, height)\n    \n    Returns:\n    - Cropped and resized 2D array\n    \"\"\"\n    h, w = img.shape[:2]\n    center_x, center_y = w // 2, h // 2\n    new_w, new_h = int(w * scale), int(h * scale)\n    \n    x1 = max(0, center_x - new_w // 2)\n    x2 = min(w, center_x + new_w // 2)\n    y1 = max(0, center_y - new_h // 2)\n    y2 = min(h, center_y + new_h // 2)\n    \n    cropped = img[y1:y2, x1:x2]\n    resized = cv2.resize(cropped, dim, interpolation=cv2.INTER_LINEAR)\n    return resized","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.020707Z","iopub.execute_input":"2026-07-05T08:55:42.021070Z","iopub.status.idle":"2026-07-05T08:55:42.026977Z","shell.execute_reply.started":"2026-07-05T08:55:42.021043Z","shell.execute_reply":"2026-07-05T08:55:42.026187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test crop_resize_img function on sample image\nif 'sample_img' in locals():\n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n    \n    axes[0].imshow(sample_img, cmap=\"gray\")\n    axes[0].set_title(f\"Original\\nShape: {sample_img.shape}\")\n    axes[0].axis(\"off\")\n    \n    cropped = crop_resize_img(sample_img, scale=0.7, dim=(244, 244))\n    axes[1].imshow(cropped, cmap=\"gray\")\n    axes[1].set_title(f\"Cropped (scale=0.7)\\nShape: {cropped.shape}\")\n    axes[1].axis(\"off\")\n    \n    plt.suptitle(\"Crop & Resize Function Test\", fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\nelse:\n    print(\"sample_img not found. Run Cell 14 first.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.028034Z","iopub.execute_input":"2026-07-05T08:55:42.028368Z","iopub.status.idle":"2026-07-05T08:55:42.316799Z","shell.execute_reply.started":"2026-07-05T08:55:42.028331Z","shell.execute_reply":"2026-07-05T08:55:42.315923Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_2_4\">Global preprocessing function</span>\n\nWe can now create a global preprocessing function that will be applied to our MRI images before entering the neural network models. This function will take over the various treatments seen previously.","metadata":{}},{"cell_type":"code","source":"# Wrapper function for standard MRI preprocessing pipeline\ndef mri_preprocessor(img, scale=0.8, dim=(240, 240)):\n    \"\"\"\n    Apply standard preprocessing: crop + resize.\n    \n    Parameters:\n    - img: 2D numpy array (MRI slice)\n    - scale: float, crop scale factor\n    - dim: tuple, target dimensions (width, height)\n    \n    Returns:\n    - Preprocessed 2D array\n    \"\"\"\n    return crop_resize_img(img, scale=scale, dim=dim)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.318009Z","iopub.execute_input":"2026-07-05T08:55:42.318338Z","iopub.status.idle":"2026-07-05T08:55:42.322848Z","shell.execute_reply.started":"2026-07-05T08:55:42.318297Z","shell.execute_reply":"2026-07-05T08:55:42.321971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test mri_preprocessor on sample image\nif 'sample_img' in locals():\n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n    \n    axes[0].imshow(sample_img, cmap=\"gray\")\n    axes[0].set_title(f\"Original\\nShape: {sample_img.shape}\")\n    axes[0].axis(\"off\")\n    \n    preprocessed = mri_preprocessor(sample_img, scale=0.8, dim=(240, 240))\n    axes[1].imshow(preprocessed, cmap=\"gray\")\n    axes[1].set_title(f\"Preprocessed\\nShape: {preprocessed.shape}\")\n    axes[1].axis(\"off\")\n    \n    plt.suptitle(\"MRI Preprocessor Function Test\", fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\nelse:\n    print(\"sample_img not found. Run Cell 14 first.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.323791Z","iopub.execute_input":"2026-07-05T08:55:42.324089Z","iopub.status.idle":"2026-07-05T08:55:42.589017Z","shell.execute_reply.started":"2026-07-05T08:55:42.324056Z","shell.execute_reply":"2026-07-05T08:55:42.588031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <span style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\" id=\"section_3\">Development of supervised models</span>\n\nWe will test several models of neural networks and test the main metrics to determine the best model.\n\n- **CNN** from scratch (Baseline).\n- **Learning transfer**.\n- **Fine Tuning**.\n\nThe metrics tested will be the **accuracy** in train and validation.","metadata":{}},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_3_1\">Multimodal inputs CNN from scratch</span>\n\nAs part of this modeling, it is necessary to develop the models on each of the 4 types of scans available to patients:\n\n- Fluid Attenuated Inversion Recovery (FLAIR).\n- T1-weighted pre-contrast (T1w).\n- T1-weighted post-contrast (T1Gd).\n- T2-weighted (T2).\n\nAs a reminder, these models are stored in the variable:\n```Python\nscan_categories = [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]\n```\n\nThe arrays of the patient images will be loaded into the variable X, the labels of MGMT_value into the variable y. To obtain a binary classification, we will use a global **sigmoid activation layer** to the results of our 4 models *(classifier layer)*.\n\nWe will test the models with **a sequence of images** *(as for a video)* by taking the **middle 24 images** *(not entirely black)* of each patient for each type of scan :","metadata":{}},{"cell_type":"code","source":"# [Optional] Quick verification of slice counts using pre-computed train_dataset\n# This avoids re-scanning the filesystem\n\n#print(\"Slice count statistics per modality (from train_dataset):\")\n#for scan in SCAN_CATEGORIES:\n    #col = f\"{scan}_count\"\n    #if col in train_dataset.columns:\n        #print(f\"\\n{scan}:\")\n        #print(f\"  Mean: {train_dataset[col].mean():.1f}\")\n        #print(f\"  Median: {train_dataset[col].median():.1f}\")\n        #print(f\"  Min: {train_dataset[col].min()}, Max: {train_dataset[col].max()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.590140Z","iopub.execute_input":"2026-07-05T08:55:42.590512Z","iopub.status.idle":"2026-07-05T08:55:42.594557Z","shell.execute_reply.started":"2026-07-05T08:55:42.590467Z","shell.execute_reply":"2026-07-05T08:55:42.593899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Model and Training Hyperparameters  (OPTIMIZED)\n# =============================================================================\n\nIMAGE_SIZE   = 64     # Slice width/height\nNUM_IMAGES   = 20     # REDUCED from 30 → -33% memory & compute per volume\nBATCH_SIZE   = 8      # INCREASED from 2 → better GPU utilization with float16\n\nLEARNING_RATE       = 8e-5\nN_SPLITS            = 4\nEPOCHS              = 25    # max per fold; early-stopping will cut sooner\nPATIENCE_EARLY_STOP = 6\nPATIENCE_REDUCE_LR  = 3\n\nRANDOM_STATE    = 42\nSCAN_CATEGORIES = [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]\n\n# OPTIMIZATION 2: Pre-cache dir\nCACHE_DIR = '/kaggle/working/mri_cache'\n\nprint(f\"Config: {IMAGE_SIZE}x{IMAGE_SIZE}x{NUM_IMAGES}, Batch={BATCH_SIZE}, Epochs={EPOCHS}\")\nprint(f\"Modalities: {SCAN_CATEGORIES}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.595881Z","iopub.execute_input":"2026-07-05T08:55:42.596253Z","iopub.status.idle":"2026-07-05T08:55:42.610246Z","shell.execute_reply.started":"2026-07-05T08:55:42.596215Z","shell.execute_reply":"2026-07-05T08:55:42.609491Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# [Optional] Descriptive statistics for slice counts\n# Requires slice_counts_df from Cell 21 (if kept)\n\nif 'slice_counts_df' in locals():\n    print(\"Slice count statistics by modality:\")\n    print(slice_counts_df.groupby(\"scan_type\")[\"num_slices\"].describe())\nelse:\n    print(\"slice_counts_df not found. Use train_dataset columns instead:\")\n    for scan in SCAN_CATEGORIES:\n        col = f\"{scan}_count\"\n        if col in train_dataset.columns:\n            print(f\"\\n{scan}_count statistics:\")\n            print(train_dataset[col].describe())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.611410Z","iopub.execute_input":"2026-07-05T08:55:42.611726Z","iopub.status.idle":"2026-07-05T08:55:42.638369Z","shell.execute_reply.started":"2026-07-05T08:55:42.611695Z","shell.execute_reply":"2026-07-05T08:55:42.637724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# [Optional] Top patients by slice count (for debugging/EDA)\nif 'slice_counts_df' in locals():\n    print(\"Top 20 entries by slice count:\")\n    print(slice_counts_df.sort_values(\"num_slices\", ascending=False).head(20))\nelse:\n    print(\"Skipping: slice_counts_df not available.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.639313Z","iopub.execute_input":"2026-07-05T08:55:42.639731Z","iopub.status.idle":"2026-07-05T08:55:42.644759Z","shell.execute_reply.started":"2026-07-05T08:55:42.639703Z","shell.execute_reply":"2026-07-05T08:55:42.644058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare training DataFrame with formatted IDs and fold column\ntrain_df = train_labels.copy()\ntrain_df['BraTS21ID5'] = train_df['BraTS21ID'].apply(lambda x: str(x).zfill(5))\ntrain_df[\"Fold\"] = -1  # Initialize fold column\n\nprint(f\"train_df shape: {train_df.shape}\")\nprint(train_df[[\"BraTS21ID\", \"BraTS21ID5\", \"MGMT_value\", \"Fold\"]].head(3))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.645941Z","iopub.execute_input":"2026-07-05T08:55:42.646567Z","iopub.status.idle":"2026-07-05T08:55:42.661329Z","shell.execute_reply.started":"2026-07-05T08:55:42.646525Z","shell.execute_reply":"2026-07-05T08:55:42.660610Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_3_2\">Define loaders for images sequences 4 MRI types</span>","metadata":{}},{"cell_type":"code","source":"# Core function: Load and preprocess 3D MRI volume for a single modality\ndef load_dicom_images_3d(scan_id, mri_type=\"FLAIR\", num_images=30, img_size=64, split=\"train\", random_sampling=True):\n    \"\"\"\n    Load, filter, sample, and preprocess a 3D MRI volume.\n    \n    Parameters:\n    - scan_id: str, patient BraTS21ID (zero-padded)\n    - mri_type: str, modality name (FLAIR/T1w/T1wCE/T2w)\n    - num_images: int, target number of slices\n    - img_size: int, resize dimension for each slice\n    - split: str, 'train' or 'test'\n    - random_sampling: bool, use random sampling if True, uniform if False\n    \n    Returns:\n    - 3D numpy array of shape (img_size, img_size, num_images)\n    \"\"\"\n    folder_path = os.path.join(input_path, split, scan_id, mri_type)\n    \n    # Get sorted DICOM files (numerical order)\n    files = sorted(\n        glob.glob(os.path.join(folder_path, \"*.dcm\")),\n        key=lambda var: [int(x) if x.isdigit() else x for x in re.split('([0-9]+)', var)]\n    )\n    \n    # Read and filter valid (non-black) images\n    valid_images = []\n    for f in files:\n        try:\n            dcm = dicom.dcmread(f)\n            img = dcm.pixel_array\n            # Filter: keep images with meaningful content\n            if np.max(img) > 0 and np.mean(img) > 1.0:\n                valid_images.append(img)\n        except Exception as e:\n            continue  # Skip corrupted files\n    \n    # Handle empty case: return zero volume\n    if len(valid_images) == 0:\n        return np.zeros((img_size, img_size, num_images), dtype=np.float32)\n    \n    # Sampling strategy\n    total = len(valid_images)\n    if total > num_images:\n        if random_sampling:\n            # Random sampling for training (augmentation effect)\n            indices = np.sort(np.random.choice(total, num_images, replace=False))\n        else:\n            # Uniform sampling for validation/test (consistency)\n            step = total / num_images\n            indices = [int(i * step) for i in range(num_images)]\n            indices[-1] = min(indices[-1], total - 1)  # Safety bound\n        selected = [valid_images[i] for i in indices]\n    else:\n        # Use all available slices if fewer than target\n        selected = valid_images\n    \n    # Preprocess each slice: resize + normalize\n    processed = []\n    for img in selected:\n        img = cv2.resize(img, (img_size, img_size))\n        if np.max(img) > 0:\n            img = img / np.max(img)  # Normalize to [0, 1]\n        processed.append(img)\n    \n    # Pad with zeros if needed to reach target depth\n    while len(processed) < num_images:\n        processed.append(np.zeros((img_size, img_size), dtype=np.float32))\n    \n    # Stack and transpose to (H, W, D)\n    volume = np.stack(processed, axis=-1).astype(np.float32)\n    return volume","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.662703Z","iopub.execute_input":"2026-07-05T08:55:42.663038Z","iopub.status.idle":"2026-07-05T08:55:42.680067Z","shell.execute_reply.started":"2026-07-05T08:55:42.663013Z","shell.execute_reply":"2026-07-05T08:55:42.679195Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Pre-cache all DICOM volumes as .npy\n\n**This is the #1 fix for Kaggle timeout.**  \nReading DICOM files during every training step causes hundreds of disk reads per batch.  \nThe cell below converts every patient's 4 modalities to `.npy` arrays **once** (~15-30 min).  \nAfter this, each batch loads in milliseconds, training speed increases by ~10-20x.\n","metadata":{}},{"cell_type":"code","source":"# =============================================================================\n# OPTIMIZATION 3: Pre-cache volumes as .npy — run ONCE before training\n# =============================================================================\nimport os\n\nos.makedirs(os.path.join(CACHE_DIR, 'train'), exist_ok=True)\nos.makedirs(os.path.join(CACHE_DIR, 'test'),  exist_ok=True)\n\n\ndef _cache_path(pid_str, mri_type, split='train'):\n    return os.path.join(CACHE_DIR, split, f'{pid_str}_{mri_type}.npy')\n\n\ndef build_cache(df, split='train'):\n    \"\"\"Pre-process each patient's volumes and save as .npy for fast loading.\"\"\"\n    new_files = 0\n    for pid_raw in tqdm(df['BraTS21ID'].values, desc=f'Caching {split}'):\n        pid_str = str(pid_raw).zfill(5)\n        for mri_type in SCAN_CATEGORIES:\n            path = _cache_path(pid_str, mri_type, split)\n            if not os.path.exists(path):\n                vol = load_dicom_images_3d(\n                    scan_id=pid_str,\n                    mri_type=mri_type,\n                    num_images=NUM_IMAGES,\n                    img_size=IMAGE_SIZE,\n                    split=split,\n                    random_sampling=False\n                )\n                np.save(path, vol)\n                new_files += 1\n    print(f'Done. {new_files} new volumes cached to {CACHE_DIR}/{split}/')\n\n\ndef load_volume(pid_str, mri_type, split='train'):\n    \"\"\"Load pre-cached volume; fall back to DICOM if cache miss.\"\"\"\n    path = _cache_path(pid_str, mri_type, split)\n    if os.path.exists(path):\n        return np.load(path)\n    return load_dicom_images_3d(\n        pid_str, mri_type, NUM_IMAGES, IMAGE_SIZE, split, random_sampling=False\n    )\n\n\n# Build train cache\nbuild_cache(train_labels, split='train')\n\n# Build test cache\n_test_dir = os.path.join(input_path, 'test')\n_test_ids  = [int(d) for d in os.listdir(_test_dir)\n              if os.path.isdir(os.path.join(_test_dir, d))]\nbuild_cache(pd.DataFrame({'BraTS21ID': _test_ids}), split='test')\n\nprint(\"All caches ready — training batches will now load in milliseconds!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T08:55:42.685443Z","iopub.execute_input":"2026-07-05T08:55:42.685827Z","iopub.status.idle":"2026-07-05T10:10:02.526986Z","shell.execute_reply.started":"2026-07-05T08:55:42.685802Z","shell.execute_reply":"2026-07-05T10:10:02.525991Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# [Optional] Load and preprocess a single 2D DICOM image\n# Note: Main pipeline uses load_dicom_images_3d, so this is for debugging only.\n\ndef load_dicom_image(path, img_size=IMAGE_SIZE, preproc=True):\n    \"\"\"\n    Load and optionally preprocess a single 2D MRI slice.\n    \n    Parameters:\n    - path: str, path to .dcm file\n    - img_size: int, target dimension\n    - preproc: bool, apply crop+resize if True\n    \n    Returns:\n    - 2D numpy array\n    \"\"\"\n    data = dicom.dcmread(path).pixel_array   # \n    \n    if preproc:\n        return mri_preprocessor(data, scale=0.8, dim=(img_size, img_size))\n    else:\n        return cv2.resize(data, (img_size, img_size))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:02.528144Z","iopub.execute_input":"2026-07-05T10:10:02.528407Z","iopub.status.idle":"2026-07-05T10:10:02.534399Z","shell.execute_reply.started":"2026-07-05T10:10:02.528385Z","shell.execute_reply":"2026-07-05T10:10:02.533853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test single image loading and preprocessing\nsample_path = os.path.join(input_path, \"train\", \"00046\", \"FLAIR\", \"Image-90.dcm\")\nif os.path.exists(sample_path):\n    original = dicom.dcmread(sample_path).pixel_array   # \n    processed = load_dicom_image(sample_path, preproc=True)\n    \n    fig, axes = plt.subplots(1, 2, figsize=(12, 6))\n    axes[0].imshow(original, cmap=\"gray\")\n    axes[0].set_title(f\"Original\\nShape: {original.shape}\")\n    axes[0].axis(\"off\")\n    \n    axes[1].imshow(processed, cmap=\"gray\")\n    axes[1].set_title(f\"Preprocessed\\nShape: {processed.shape}\")\n    axes[1].axis(\"off\")\n    \n    plt.suptitle(\"Single Image Loader Test\", fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\nelse:\n    print(f\"File not found: {sample_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:02.535254Z","iopub.execute_input":"2026-07-05T10:10:02.535551Z","iopub.status.idle":"2026-07-05T10:10:02.778016Z","shell.execute_reply.started":"2026-07-05T10:10:02.535527Z","shell.execute_reply":"2026-07-05T10:10:02.777400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test 3D volume loading for a single patient\ntest_vol = load_dicom_images_3d(\"00046\", mri_type=\"FLAIR\", num_images=NUM_IMAGES, img_size=IMAGE_SIZE)\nprint(f\"3D Volume shape: {test_vol.shape}\")\nprint(f\"Expected shape: ({IMAGE_SIZE}, {IMAGE_SIZE}, {NUM_IMAGES})\")\n\n# Plot middle slice\nmid_idx = test_vol.shape[2] // 2\nplt.figure(figsize=(6, 6))\nplt.imshow(test_vol[:, :, mid_idx], cmap=\"gray\")\nplt.title(f\"FLAIR Volume - Slice {mid_idx}\", fontsize=12, fontweight='bold')\nplt.axis(\"off\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:02.778962Z","iopub.execute_input":"2026-07-05T10:10:02.779353Z","iopub.status.idle":"2026-07-05T10:10:04.271131Z","shell.execute_reply.started":"2026-07-05T10:10:02.779328Z","shell.execute_reply":"2026-07-05T10:10:04.270150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify shape consistency across all 4 MRI modalities\nfor scan in SCAN_CATEGORIES:\n    vol = load_dicom_images_3d(\"00046\", mri_type=scan, num_images=NUM_IMAGES, img_size=IMAGE_SIZE)\n    print(f\"{scan}: shape = {vol.shape}, dtype = {vol.dtype}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:04.272755Z","iopub.execute_input":"2026-07-05T10:10:04.273008Z","iopub.status.idle":"2026-07-05T10:10:08.648722Z","shell.execute_reply.started":"2026-07-05T10:10:04.272972Z","shell.execute_reply":"2026-07-05T10:10:08.647969Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Thanks to this function, we obtain **4 ordered sequences *(one for each Scan type)* of 24 images of dimension 120 x 120** and the preprocessing has been applied.\n\nTo couple our four types of scans, we will use a **multi-modal approach** to create our model. We will integrate 4 different inputs for a single final classifier.","metadata":{}},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_3_3\">Define folds</span>","metadata":{}},{"cell_type":"code","source":"# Assign Stratified K-Fold labels to training DataFrame\nfrom sklearn.model_selection import StratifiedKFold\n\nskf = StratifiedKFold(n_splits=N_SPLITS, shuffle=True, random_state=RANDOM_STATE)\ntrain_df[\"Fold\"] = -1  # Reset fold column\n\nfor fold_idx, (_, valid_idx) in enumerate(skf.split(train_df[\"BraTS21ID\"], train_df[\"MGMT_value\"]), start=1):\n    train_df.loc[valid_idx, \"Fold\"] = fold_idx\n\nprint(\"Fold assignment completed.\")\nprint(train_df[\"Fold\"].value_counts().sort_index())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:08.649501Z","iopub.execute_input":"2026-07-05T10:10:08.649862Z","iopub.status.idle":"2026-07-05T10:10:08.665699Z","shell.execute_reply.started":"2026-07-05T10:10:08.649827Z","shell.execute_reply":"2026-07-05T10:10:08.664957Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_3_4\">Keras custom data generator</span>\nThanks to the Sequence module of the Keras library, we are going to create a personalized image generator. This will prevent us from creating Numpy arrays or Tensors containing all the sequences which would quickly overload the memory.","metadata":{}},{"cell_type":"code","source":"# =============================================================================\n# Custom Keras Sequence — loads from .npy cache  (OPTIMIZED for Keras 3)\n# =============================================================================\nclass Dataset(Sequence):\n    def __init__(self, df, is_train=True, batch_size=BATCH_SIZE,\n                 shuffle=False, augment=False, **kwargs):\n        super().__init__(**kwargs)          # \n        self.patient_ids = df['BraTS21ID'].values\n        self.is_train    = is_train\n        self.batch_size  = batch_size\n        self.shuffle     = shuffle\n        self.augment     = augment\n        self.split_type  = 'train' if is_train else 'test'\n\n        if 'MGMT_value' in df.columns and df['MGMT_value'].notnull().any():\n            self.labels = df['MGMT_value'].values\n        else:\n            self.labels = None\n\n        self.indices = np.arange(len(self.patient_ids))\n        self.on_epoch_end()\n\n    def __len__(self):\n        return math.ceil(len(self.patient_ids) / self.batch_size)\n\n    def _augment_vol(self, vol):\n        if np.random.rand() > 0.5:\n            vol = vol[:, ::-1, :]\n        if np.random.rand() > 0.5:\n            vol = vol[::-1, :, :]\n        return vol\n\n    def __getitem__(self, idx):\n        batch_idx = self.indices[idx * self.batch_size : (idx + 1) * self.batch_size]\n        batch_ids = self.patient_ids[batch_idx]\n\n        batch_modalities = []\n        for modality in SCAN_CATEGORIES:\n            mod_imgs = []\n            for pid in batch_ids:\n                pid_str = str(pid).zfill(5)\n                vol = load_volume(pid_str, modality, self.split_type)\n                if self.augment:\n                    vol = self._augment_vol(vol)\n                mod_imgs.append(vol)\n            batch_modalities.append(\n                np.expand_dims(np.stack(mod_imgs, axis=0), axis=-1)\n            )\n\n        central_slices = []\n        for mod_batch in batch_modalities:\n            center_idx = mod_batch.shape[3] // 2\n            central_slices.append(mod_batch[:, :, :, center_idx, :])\n        batch_slices_2d = np.concatenate(central_slices, axis=-1)\n\n        # tuple مطلوب في Keras 3 — list تسبب TypeError في from_generator\n        model_inputs = tuple(batch_modalities + [batch_slices_2d])\n\n        if self.labels is not None:\n            return model_inputs, self.labels[batch_idx].astype(np.float32)\n        return model_inputs\n\n    def on_epoch_end(self):\n        if self.shuffle:\n            np.random.shuffle(self.indices)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:08.666762Z","iopub.execute_input":"2026-07-05T10:10:08.667058Z","iopub.status.idle":"2026-07-05T10:10:08.679025Z","shell.execute_reply.started":"2026-07-05T10:10:08.667036Z","shell.execute_reply":"2026-07-05T10:10:08.678359Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Once the generators are created, we can project an image to check:","metadata":{}},{"cell_type":"markdown","source":"## <span style=\"color:#3c99dc; font-size:18px; text-transform: uppercase; font-weight:bold\" id=\"section_3_5\">Define CNN Multi-inputs model</span>","metadata":{}},{"cell_type":"code","source":"# Quick diagnostic: Test Dataset class on a small subset\nsample_ids = train_df.head(2)\ntest_loader = Dataset(sample_ids, is_train=True, batch_size=2, shuffle=False)\n\nbatch_X, batch_y = test_loader[0]\n\nprint(\"Batch X type:\", type(batch_X))\nprint(\"Number of model inputs:\", len(batch_X))\n\nfor i, input_array in enumerate(batch_X):\n    print(f\"Input {i} shape:\", input_array.shape)\n\nprint(\"Labels shape:\", batch_y.shape)\nprint(\"Labels:\", batch_y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:08.679964Z","iopub.execute_input":"2026-07-05T10:10:08.680345Z","iopub.status.idle":"2026-07-05T10:10:08.711937Z","shell.execute_reply.started":"2026-07-05T10:10:08.680308Z","shell.execute_reply":"2026-07-05T10:10:08.711253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# 3D ResNet + EfficientNetB0\n# =============================================================================\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom tensorflow.keras.layers import Add, Activation\n\n\ndef create_cnn_model(learning_rate=1e-4):\n    \"\"\"\n    Hybrid model:\n    - 3D ResNet branch for the four MRI modalities\n    - 2D EfficientNetB0 branch for the central multimodal slice\n    - Feature fusion for MGMT prediction\n    \"\"\"\n\n    # ============================================================\n    # INPUT SHAPES\n    # ============================================================\n    input_shape_3d = (IMAGE_SIZE, IMAGE_SIZE, NUM_IMAGES, 1)\n    input_shape_2d = (IMAGE_SIZE, IMAGE_SIZE, 4)\n\n    # 4 MRI 3D inputs\n    inputs_flair = Input(shape=input_shape_3d, name=\"FLAIR_input\")\n    inputs_t1w   = Input(shape=input_shape_3d, name=\"T1w_input\")\n    inputs_t1wce = Input(shape=input_shape_3d, name=\"T1wCE_input\")\n    inputs_t2w   = Input(shape=input_shape_3d, name=\"T2w_input\")\n\n    # 2D multimodal input for EfficientNetB0\n    input_2d = Input(shape=input_shape_2d, name=\"EfficientNet_2D_input\")\n\n    # ============================================================\n    # 3D RESNET BLOCK\n    # ============================================================\n    def resnet3d_block(x, filters, stride=1):\n        shortcut = x\n\n        x = Conv3D(filters, kernel_size=3, strides=stride, padding=\"same\", use_bias=False)(x)\n        x = BatchNormalization()(x)\n        x = Activation(\"relu\")(x)\n\n        x = Conv3D(filters, kernel_size=3, strides=1, padding=\"same\", use_bias=False)(x)\n        x = BatchNormalization()(x)\n\n        if shortcut.shape[-1] != filters or stride != 1:\n            shortcut = Conv3D(filters, kernel_size=1, strides=stride, padding=\"same\", use_bias=False)(shortcut)\n            shortcut = BatchNormalization()(shortcut)\n\n        x = Add()([x, shortcut])\n        x = Activation(\"relu\")(x)\n        return x\n\n    # ============================================================\n    # 3D RESNET FEATURE EXTRACTOR\n    # ============================================================\n    def resnet3d_feature_extractor(x):\n        x = Conv3D(32, kernel_size=7, strides=2, padding=\"same\", use_bias=False)(x)\n        x = BatchNormalization()(x)\n        x = Activation(\"relu\")(x)\n        x = MaxPool3D(pool_size=2)(x)\n\n        x = resnet3d_block(x, 32, stride=1)\n        x = resnet3d_block(x, 32, stride=1)\n\n        x = resnet3d_block(x, 64, stride=2)\n        x = resnet3d_block(x, 64, stride=1)\n\n        x = resnet3d_block(x, 128, stride=2)\n        x = resnet3d_block(x, 128, stride=1)\n\n        x = GlobalAveragePooling3D()(x)\n        return x\n\n    # ============================================================\n    # 3D BRANCH: ONE RESNET3D PER MRI MODALITY\n    # ============================================================\n    f_flair  = resnet3d_feature_extractor(inputs_flair)\n    f_t1w    = resnet3d_feature_extractor(inputs_t1w)\n    f_t1wce  = resnet3d_feature_extractor(inputs_t1wce)\n    f_t2w    = resnet3d_feature_extractor(inputs_t2w)\n\n    features_3d = concatenate(\n        [f_flair, f_t1w, f_t1wce, f_t2w],\n        name=\"fusion_3d_modalities\"\n    )\n\n    # ============================================================\n    # 2D BRANCH: EFFICIENTNETB0\n    # ============================================================\n    efficientnet = EfficientNetB0(\n        include_top=False,\n        weights=None,\n        input_tensor=input_2d,\n        pooling=\"avg\"\n    )\n    features_2d = efficientnet.output\n\n    # ============================================================\n    # FUSION\n    # ============================================================\n    combined = concatenate(\n        [features_3d, features_2d],\n        name=\"fusion_3d_resnet_2d_efficientnet\"\n    )\n\n    # ============================================================\n    # CLASSIFICATION HEAD\n    # ============================================================\n    x = Dense(512, activation=\"relu\")(combined)\n    x = BatchNormalization()(x)\n    x = Dropout(0.5)(x)\n\n    x = Dense(128, activation=\"relu\")(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.4)(x)\n\n    x = Dense(64, activation=\"relu\")(x)\n    x = BatchNormalization()(x)\n    x = Dropout(0.3)(x)\n\n    output = Dense(1, activation=\"sigmoid\", dtype=\"float32\", name=\"MGMT_output\")(x)\n\n    # ============================================================\n    # BUILD & COMPILE\n    # ============================================================\n    model = Model(\n        inputs=[inputs_flair, inputs_t1w, inputs_t1wce, inputs_t2w, input_2d],\n        outputs=output,\n        name=\"Hybrid_3DResNet_EfficientNetB0\"\n    )\n\n    model.compile(\n        optimizer=keras.optimizers.Adam(\n            learning_rate=learning_rate,\n            clipnorm=1.0\n        ),\n        loss=\"binary_crossentropy\",\n        metrics=[\"accuracy\", keras.metrics.AUC(name=\"auc\")]\n    )\n\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:08.712959Z","iopub.execute_input":"2026-07-05T10:10:08.713274Z","iopub.status.idle":"2026-07-05T10:10:08.731348Z","shell.execute_reply.started":"2026-07-05T10:10:08.713237Z","shell.execute_reply":"2026-07-05T10:10:08.730549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Instantiate model for architecture inspection\n# Note: A fresh model will be rebuilt inside each CV fold during training\nmodel_inspect = create_cnn_model(learning_rate=1e-4)\nprint(\"Model built and compiled successfully.\")\nprint(f\"Total parameters: {model_inspect.count_params():,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:08.732425Z","iopub.execute_input":"2026-07-05T10:10:08.732923Z","iopub.status.idle":"2026-07-05T10:10:12.476433Z","shell.execute_reply.started":"2026-07-05T10:10:08.732892Z","shell.execute_reply":"2026-07-05T10:10:12.475787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test forward pass del modello ibrido su un batch reale\nsample_ids = train_df.head(2)\ntest_loader = Dataset(sample_ids, is_train=True, batch_size=2, shuffle=False)\n\nbatch_X, batch_y = test_loader[0]\n\ntest_model = create_cnn_model(learning_rate=1e-4)\n\npreds = test_model.predict(batch_X)\n\nprint(\"Numero input dati al modello:\", len(batch_X))\n\nfor i, input_array in enumerate(batch_X):\n    print(f\"Input {i} shape:\", input_array.shape)\n\nprint(\"Predictions shape:\", preds.shape)\nprint(\"Predictions:\", preds)\nprint(\"Labels shape:\", batch_y.shape)\nprint(\"Labels:\", batch_y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:12.477286Z","iopub.execute_input":"2026-07-05T10:10:12.477971Z","iopub.status.idle":"2026-07-05T10:10:26.474808Z","shell.execute_reply.started":"2026-07-05T10:10:12.477948Z","shell.execute_reply":"2026-07-05T10:10:26.474051Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize and save model architecture\nplot_model(model_inspect, show_shapes=True, show_layer_names=True, \n           to_file=\"model_architecture.png\", dpi=100)\nprint(\"Architecture diagram saved as 'model_architecture.png'\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:26.475900Z","iopub.execute_input":"2026-07-05T10:10:26.476247Z","iopub.status.idle":"2026-07-05T10:10:36.549843Z","shell.execute_reply.started":"2026-07-05T10:10:26.476223Z","shell.execute_reply":"2026-07-05T10:10:36.549100Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Then we **train the multi-input model** with the network defined above:","metadata":{}},{"cell_type":"code","source":"# (Debug subset removed — using full dataset with cache)\n# train_df = train_df.sample(n=20, ...) -- no longer needed\nprint(f\"train_df has {len(train_df)} patients (full dataset).\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T10:10:36.551612Z","iopub.execute_input":"2026-07-05T10:10:36.551925Z","iopub.status.idle":"2026-07-05T10:10:36.557130Z","shell.execute_reply.started":"2026-07-05T10:10:36.551902Z","shell.execute_reply":"2026-07-05T10:10:36.556000Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# Training Loop\n# =============================================================================\nfrom sklearn.model_selection import StratifiedKFold\nimport gc\n\nprint(f\"Training on full dataset: {len(train_df)} patients.\")\n\nN_SPLITS      = 4\nEPOCHS        = 25\nLEARNING_RATE = 8e-5\nPATIENCE_ES   = 6\nPATIENCE_RLR  = 3\n\n# Custom callback: prints all metrics clearly after each epoch\nclass MetricsLogger(tf.keras.callbacks.Callback):\n    def on_epoch_end(self, epoch, logs=None):\n        logs = logs or {}\n        parts = [f\"Epoch {epoch+1}\"]\n        for k, v in logs.items():\n            parts.append(f\"{k}: {v:.4f}\")\n        print(\" — \".join(parts))\n\nskf = StratifiedKFold(n_splits=N_SPLITS, shuffle=True, random_state=42)\ncv_results        = []\nlast_fold_history = None\n\nX = train_df['BraTS21ID'].values\ny = train_df['MGMT_value'].values\n\nfor fold, (train_idx, valid_idx) in enumerate(skf.split(X, y), start=1):\n    print(f\"\\n{'='*40} FOLD {fold}/{N_SPLITS} {'='*40}\")\n\n    df_train_fold = train_df.iloc[train_idx].reset_index(drop=True)\n    df_valid_fold = train_df.iloc[valid_idx].reset_index(drop=True)\n\n    train_gen = Dataset(df_train_fold, is_train=True,\n                        batch_size=BATCH_SIZE, shuffle=True, augment=True)\n    valid_gen = Dataset(df_valid_fold, is_train=True,\n                        batch_size=BATCH_SIZE, shuffle=False, augment=False)\n\n    cnn_model = create_cnn_model(learning_rate=LEARNING_RATE)\n\n    ckpt_path = f'/kaggle/working/best_model_fold_{fold}.keras'\n    callbacks = [\n        MetricsLogger(),\n        ModelCheckpoint(ckpt_path, monitor='val_auc', mode='max',\n                        save_best_only=True, verbose=1),\n        EarlyStopping(monitor='val_auc', mode='max',\n                      patience=PATIENCE_ES, restore_best_weights=True, verbose=1),\n        ReduceLROnPlateau(monitor='val_auc', mode='max', factor=0.8,  # ← FIX 3\n                          patience=PATIENCE_RLR, min_lr=1e-7, verbose=1),\n    ]\n\n    #modifica per XAI\n    history = cnn_model.fit(\n        train_gen,\n        validation_data=valid_gen,\n        epochs=EPOCHS,\n        callbacks=callbacks,\n        verbose=1,\n    )\n\n    manual_save_path = f\"/kaggle/working/final_model_fold_{fold}.keras\"\n    cnn_model.save(manual_save_path)\n    print(f\"Modello salvato in: {manual_save_path}\")\n\n    best_val_auc = max(history.history['val_auc'])\n    cv_results.append({'fold': fold, 'best_val_auc': best_val_auc})\n    last_fold_history = history\n\n    del cnn_model, train_gen, valid_gen\n    gc.collect()\n    tf.keras.backend.clear_session()\n\n    print(f\"Fold {fold} done. Best Val AUC: {best_val_auc:.4f}\")\n\nprint(\"\\nCross-Validation complete!\")","metadata":{"_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.busy":"2026-07-05T11:18:01.600216Z","iopub.execute_input":"2026-07-05T11:18:01.601072Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display Cross-Validation Results Summary\ncv_results_df = pd.DataFrame(cv_results)\nprint(cv_results_df)\nprint(f\"\\nMean Val AUC: {cv_results_df['best_val_auc'].mean():.4f} (+/- {cv_results_df['best_val_auc'].std():.4f})\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best checkpoint and select a correctly classified sample for XAI\n# XAI maps should be generated only for cases correctly classified by the model.\n\nfrom tensorflow.keras.models import load_model\nimport numpy as np\n\n# 1. Load the best checkpoint\nbest_model_path = \"/kaggle/working/best_model_fold_3.keras\"\n\nxai_model = load_model(best_model_path)\n\nprint(\"Model loaded successfully.\")\nprint(\"Number of inputs required by the model:\", len(xai_model.inputs))\nprint(\"Output shape:\", xai_model.output_shape)\n\n# 2. Search for a correctly classified patient\ncorrect_sample_found = False\n\nfor idx in range(len(train_df)):\n\n    # Select one patient\n    sample_ids = train_df.iloc[[idx]]\n\n    # Create a one-patient dataset\n    xai_loader = Dataset(\n        sample_ids,\n        is_train=True,\n        batch_size=1,\n        shuffle=False\n    )\n\n    # Load patient inputs and true label\n    xai_inputs, xai_label = xai_loader[0]\n\n    # Run model prediction\n    xai_prediction = xai_model.predict(\n        xai_inputs,\n        verbose=0\n    )\n\n    predicted_probability = float(xai_prediction[0][0])\n    predicted_class = int(predicted_probability >= 0.5)\n    true_label = int(xai_label[0])\n\n    # The XAI methods must explain the same class predicted by the model\n    target_class = predicted_class\n\n    # Keep only correctly classified cases\n    if predicted_class == true_label:\n\n        correct_sample_found = True\n\n        print(\"Correctly classified patient found.\")\n        print(\"Patient index:\", idx)\n        print(\"True label:\", true_label)\n        print(\"Predicted probability:\", predicted_probability)\n        print(\"Predicted class:\", predicted_class)\n        print(\"Target class used for XAI:\", target_class)\n\n        print(\"Number of inputs provided by the patient:\", len(xai_inputs))\n\n        for i, arr in enumerate(xai_inputs):\n            print(f\"Input {i} shape:\", arr.shape)\n\n        break\n\nif not correct_sample_found:\n    raise ValueError(\"No correctly classified patient was found in the selected dataframe.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute Integrated Gradients on FLAIR input\n\nimport tensorflow as tf\nimport numpy as np\n\n# Convert all inputs to TensorFlow tensors\nxai_tensors = [\n    tf.convert_to_tensor(arr, dtype=tf.float32)\n    for arr in xai_inputs\n]\n\n# Baseline: completely zero volume for the FLAIR input only\nbaseline_flair = tf.zeros_like(xai_tensors[0])\n\n\ndef integrated_gradients_flair(\n    model,\n    inputs,\n    baseline,\n    steps=32\n):\n    # Generate interpolation coefficients between the baseline and the real input\n    alphas = tf.linspace(0.0, 1.0, steps + 1)\n\n    # Initialize the gradient accumulator\n    accumulated_gradients = tf.zeros_like(inputs[0])\n\n    for alpha in alphas:\n        # Interpolate only the FLAIR input\n        interpolated_flair = baseline + alpha * (inputs[0] - baseline)\n\n        # Keep the other inputs unchanged\n        current_inputs = [\n            interpolated_flair,\n            inputs[1],\n            inputs[2],\n            inputs[3],\n            inputs[4]\n        ]\n\n        # Compute the gradient of the model output with respect to the interpolated FLAIR input\n        with tf.GradientTape() as tape:\n            tape.watch(interpolated_flair)\n            prediction = model(current_inputs, training=False)\n            target = prediction[:, 0]\n\n        gradients = tape.gradient(target, interpolated_flair)\n\n        # Accumulate gradients along the interpolation path\n        accumulated_gradients += gradients\n\n    # Average gradients over all interpolation steps\n    average_gradients = accumulated_gradients / tf.cast(\n        len(alphas),\n        tf.float32\n    )\n\n    # Compute Integrated Gradients\n    integrated_gradients = (\n        inputs[0] - baseline\n    ) * average_gradients\n\n    return integrated_gradients\n\n\n# Apply Integrated Gradients to the FLAIR input\nig_flair = integrated_gradients_flair(\n    xai_model,\n    xai_tensors,\n    baseline_flair,\n    steps=32\n)\n\nprint(\"Integrated Gradients computed successfully.\")\nprint(\"Attribution map shape:\", ig_flair.shape)\nprint(\"Minimum value:\", tf.reduce_min(ig_flair).numpy())\nprint(\"Maximum value:\", tf.reduce_max(ig_flair).numpy())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize Integrated Gradients on FLAIR input \n\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# Extract the original FLAIR volume and the Integrated Gradients attribution volume\nflair_volume = xai_inputs[0][0, :, :, :, 0]\nig_volume = ig_flair.numpy()[0, :, :, :, 0]\n\n# Select the slice with the highest absolute attribution\nslice_scores = np.sum(np.abs(ig_volume), axis=(0, 1))\nbest_slice = int(np.argmax(slice_scores))\n\nflair_slice = flair_volume[:, :, best_slice]\nig_slice = ig_volume[:, :, best_slice]\n\n\n# 1. Absolute attribution visualization\n# Shows how important each region is, regardless of sign\n\nig_abs = np.abs(ig_slice)\n\nif ig_abs.max() > 0:\n    ig_norm = ig_abs / ig_abs.max()\nelse:\n    ig_norm = ig_abs\n\nplt.figure(figsize=(6, 6))\nplt.imshow(flair_slice, cmap=\"gray\")\nplt.imshow(ig_norm, cmap=\"jet\", alpha=0.45)\nplt.title(f\"Integrated Gradients, Absolute Map - FLAIR - Slice {best_slice}\")\nplt.axis(\"off\")\nplt.show()\n\n\n# 2. Signed attribution visualization\n# Shows positive and negative contributions separately\n\nmax_abs = np.max(np.abs(ig_slice))\n\nif max_abs > 0:\n    ig_signed = ig_slice / max_abs\nelse:\n    ig_signed = ig_slice\n\nfig, axes = plt.subplots(1, 3, figsize=(15, 5))\n\n# Original FLAIR slice\naxes[0].imshow(flair_slice, cmap=\"gray\")\naxes[0].set_title(\"Original FLAIR\")\naxes[0].axis(\"off\")\n\n# Signed Integrated Gradients heatmap\nim = axes[1].imshow(ig_signed, cmap=\"seismic\", vmin=-1, vmax=1)\naxes[1].set_title(\"Integrated Gradients, Signed Map\")\naxes[1].axis(\"off\")\nplt.colorbar(im, ax=axes[1], fraction=0.046, pad=0.04)\n\n# Overlay signed attribution map\naxes[2].imshow(flair_slice, cmap=\"gray\")\naxes[2].imshow(ig_signed, cmap=\"seismic\", alpha=0.5, vmin=-1, vmax=1)\naxes[2].set_title(f\"Signed Overlay - Slice {best_slice}\")\naxes[2].axis(\"off\")\n\nplt.suptitle(\n    f\"Integrated Gradients on FLAIR input - Slice {best_slice}\",\n    fontsize=14\n)\n\nplt.tight_layout()\nplt.show()\n\n\n# Summary values\nprint(\"Selected slice:\", best_slice)\nprint(\"Minimum attribution:\", ig_slice.min())\nprint(\"Maximum attribution:\", ig_slice.max())\nprint(\"Sum of absolute attributions:\", slice_scores[best_slice])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# SmoothGrad Integrated Gradients Function \n# To reduce noise by averaging attribution maps obtained \n# from multiple versions of the same input.\n\nimport numpy as np\nimport tensorflow as tf\nfrom scipy.ndimage import gaussian_filter\n\n\ndef smoothgrad_integrated_gradients(\n    model,\n    inputs,\n    target_class=target_class,\n    ig_steps=32,\n    n_samples=10,\n    noise_sigma=0.03,\n    gaussian_sigma=1.0\n):\n    \"\"\"\n    Compute SmoothGrad Integrated Gradients for a TensorFlow/Keras model.\n    \"\"\"\n\n    # Check whether the model has multiple inputs\n    multi_input = isinstance(inputs, (list, tuple))\n\n    if not multi_input:\n        inputs = [inputs]\n\n    # Convert all inputs to TensorFlow tensors\n    inputs = [\n        tf.cast(x, tf.float32)\n        for x in inputs\n    ]\n\n    # Create a black baseline for each input\n    baselines = [\n        tf.zeros_like(x)\n        for x in inputs\n    ]\n\n    # Initialize variables to accumulate attribution maps\n    accumulated_attributions = [\n        tf.zeros_like(x)\n        for x in inputs\n    ]\n\n    # Repeat Integrated Gradients on multiple noisy versions of the inputs\n    for sample_index in range(n_samples):\n\n        noisy_inputs = []\n\n        # Add Gaussian noise to each input\n        for x in inputs:\n\n            input_range = tf.reduce_max(x) - tf.reduce_min(x)\n\n            noise_std = (\n                noise_sigma\n                * tf.maximum(input_range, 1e-8)\n            )\n\n            noise = tf.random.normal(\n                shape=tf.shape(x),\n                mean=0.0,\n                stddev=noise_std,\n                dtype=tf.float32\n            )\n\n            noisy_inputs.append(x + noise)\n\n        # Interpolation path from the baseline to the noisy input\n        alphas = tf.linspace(\n            0.0,\n            1.0,\n            ig_steps + 1\n        )\n\n        gradients_sum = [\n            tf.zeros_like(x)\n            for x in inputs\n        ]\n\n        previous_gradients = None\n\n        for alpha in alphas:\n\n            interpolated_inputs = [\n                baseline + alpha * (x - baseline)\n                for x, baseline in zip(\n                    noisy_inputs,\n                    baselines\n                )\n            ]\n\n            with tf.GradientTape() as tape:\n\n                # Watch all interpolated inputs to compute gradients\n                for interpolated in interpolated_inputs:\n                    tape.watch(interpolated)\n\n                # Forward pass through the model\n                if multi_input:\n                    predictions = model(\n                        interpolated_inputs,\n                        training=False\n                    )\n                else:\n                    predictions = model(\n                        interpolated_inputs[0],\n                        training=False\n                    )\n\n                # Binary sigmoid output case\n                if predictions.shape[-1] == 1:\n\n                    score = predictions[:, 0]\n\n                    # If target_class is 0, explain the negative class\n                    if target_class == 0:\n                        score = 1.0 - score\n\n                # Multi-class softmax output case\n                else:\n\n                    if target_class is None:\n                        selected_class = tf.argmax(\n                            predictions[0]\n                        )\n                    else:\n                        selected_class = target_class\n\n                    score = predictions[\n                        :,\n                        selected_class\n                    ]\n\n            # Compute gradients of the selected score with respect to the interpolated inputs\n            gradients = tape.gradient(\n                score,\n                interpolated_inputs\n            )\n\n            # Trapezoidal approximation of the integral\n            if previous_gradients is not None:\n\n                for i in range(len(inputs)):\n                    gradients_sum[i] += (\n                        previous_gradients[i]\n                        + gradients[i]\n                    ) / 2.0\n\n            previous_gradients = gradients\n\n        # Average gradients along the integration path\n        average_gradients = [\n            gradient_sum / float(ig_steps)\n            for gradient_sum in gradients_sum\n        ]\n\n        # Compute Integrated Gradients attribution maps\n        for i in range(len(inputs)):\n\n            attribution = (\n                noisy_inputs[i] - baselines[i]\n            ) * average_gradients[i]\n\n            accumulated_attributions[i] += attribution\n\n    # Average attribution maps over all noisy samples\n    final_attributions = [\n        attribution / float(n_samples)\n        for attribution in accumulated_attributions\n    ]\n\n    # Post-processing of the attribution maps\n    processed_maps = []\n\n    for attribution in final_attributions:\n\n        attribution = attribution.numpy()\n\n        # Remove batch dimension\n        attribution = attribution[0]\n\n        # If a final channel dimension exists, remove it by averaging\n        if attribution.ndim == 4:\n            attribution = np.mean(\n                attribution,\n                axis=-1\n            )\n\n        # Use absolute attribution magnitude\n        attribution = np.abs(attribution)\n\n        # Apply Gaussian smoothing to reduce noise\n        attribution = gaussian_filter(\n            attribution,\n            sigma=gaussian_sigma\n        )\n\n        # Normalize attribution map between 0 and 1\n        attribution = (\n            attribution - attribution.min()\n        ) / (\n            attribution.max()\n            - attribution.min()\n            + 1e-8\n        )\n\n        processed_maps.append(attribution)\n\n    return processed_maps\n\n\n# Apply SmoothGrad Integrated Gradients to the trained model and selected patient\nsgig_maps = smoothgrad_integrated_gradients(\n    model=xai_model,\n    inputs=xai_inputs,\n    target_class=target_class,\n    ig_steps=32,\n    n_samples=10,\n    noise_sigma=0.03,\n    gaussian_sigma=1.0\n)\n\n# In this notebook, the first input corresponds to the FLAIR sequence\nsgig_flair = sgig_maps[0]\n\nprint(\"SmoothGrad Integrated Gradients completed successfully.\")\nprint(\"Number of attribution maps obtained:\", len(sgig_maps))\nprint(\"FLAIR attribution map shape:\", sgig_flair.shape)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compare standard Integrated Gradients with SmoothGrad Integrated Gradients\n\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# Original FLAIR volume\nflair_volume = xai_inputs[0][0, :, :, :, 0]\n\n# Standard Integrated Gradients volume\nig_original_volume = ig_flair.numpy()[0, :, :, :, 0]\n\n# SmoothGrad Integrated Gradients volume\nsgig_volume = sgig_flair\n\n# Use the same slice selected in the previous XAI visualization\nflair_slice = flair_volume[:, :, best_slice]\n\n# Standard Integrated Gradients slice\nig_original_slice = np.abs(\n    ig_original_volume[:, :, best_slice]\n)\n\n# SmoothGrad Integrated Gradients slice\nsgig_slice = sgig_volume[:, :, best_slice]\n\n\n# Normalize the standard Integrated Gradients slice between 0 and 1\nig_original_slice = (\n    ig_original_slice - ig_original_slice.min()\n) / (\n    ig_original_slice.max()\n    - ig_original_slice.min()\n    + 1e-8\n)\n\n\n# Normalize the SmoothGrad Integrated Gradients slice between 0 and 1\nsgig_slice = (\n    sgig_slice - sgig_slice.min()\n) / (\n    sgig_slice.max()\n    - sgig_slice.min()\n    + 1e-8\n)\n\n\n# Create a four-panel figure:\n# 1. Original FLAIR image\n# 2. Standard Integrated Gradients map\n# 3. SmoothGrad Integrated Gradients map\n# 4. SmoothGrad Integrated Gradients overlay on the FLAIR image\nfig, axes = plt.subplots(1, 4, figsize=(20, 5))\n\n\n# 1. Original MRI\naxes[0].imshow(flair_slice, cmap=\"gray\")\naxes[0].set_title(f\"Original FLAIR\\nSlice {best_slice}\")\naxes[0].axis(\"off\")\n\n\n# 2. Standard Integrated Gradients\naxes[1].imshow(ig_original_slice, cmap=\"jet\")\naxes[1].set_title(\"Standard Integrated\\nGradients\")\naxes[1].axis(\"off\")\n\n\n# 3. SmoothGrad Integrated Gradients\naxes[2].imshow(sgig_slice, cmap=\"jet\")\naxes[2].set_title(\"SmoothGrad Integrated\\nGradients\")\naxes[2].axis(\"off\")\n\n\n# 4. SmoothGrad Integrated Gradients overlay\naxes[3].imshow(flair_slice, cmap=\"gray\")\naxes[3].imshow(\n    sgig_slice,\n    cmap=\"jet\",\n    alpha=0.45\n)\naxes[3].set_title(\"SmoothGrad-IG Overlay\")\naxes[3].axis(\"off\")\n\n\nplt.tight_layout()\nplt.show()\n\n\n# Print numerical information about the selected slice\nprint(\"Compared slice:\", best_slice)\nprint(\"SmoothGrad-IG minimum value:\", sgig_slice.min())\nprint(\"SmoothGrad-IG maximum value:\", sgig_slice.max())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize the most relevant FLAIR slices according to SmoothGrad-Integrated Gradients\n# This cell selects the slices with the highest brain-specific attribution scores.\n# The score is computed only inside an approximate brain mask, in order to avoid\n# the influence of the black MRI background.\n\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nbrain_slice_scores = []\n\n# Compute a brain-specific SmoothGrad-IG score for each slice\nfor slice_idx in range(sgig_volume.shape[2]):\n\n    current_flair = flair_volume[:, :, slice_idx]\n    current_sgig = sgig_volume[:, :, slice_idx]\n\n    # Select non-background pixels from the FLAIR image\n    positive_pixels = current_flair[current_flair > 0]\n\n    # If the slice contains no valid brain pixels, assign zero score\n    if positive_pixels.size == 0:\n        brain_slice_scores.append(0.0)\n        continue\n\n    # Create an approximate brain mask to exclude the black background\n    current_brain_mask = current_flair > np.percentile(\n        positive_pixels,\n        5\n    )\n\n    # Extract SmoothGrad-IG values only inside the brain region\n    brain_values = current_sgig[current_brain_mask]\n\n    # If no brain attribution values are available, assign zero score\n    if brain_values.size == 0:\n        brain_slice_scores.append(0.0)\n        continue\n\n    # Keep only the top 10% strongest attribution values inside the brain\n    current_threshold = np.percentile(\n        brain_values,\n        90\n    )\n\n    strong_values = brain_values[\n        brain_values >= current_threshold\n    ]\n\n    # Compute the brain-specific attribution score for the current slice\n    brain_slice_scores.append(\n        float(np.sum(strong_values))\n    )\n\nbrain_slice_scores = np.array(brain_slice_scores)\n\n# Select the 6 slices with the highest brain-specific attribution score\ntop_brain_slices = np.argsort(\n    brain_slice_scores\n)[-6:][::-1]\n\n# Create a 2x3 grid to visualize the selected slices\nfig, axes = plt.subplots(2, 3, figsize=(15, 10))\n\nfor ax, slice_idx in zip(\n    axes.ravel(),\n    top_brain_slices\n):\n\n    current_flair = flair_volume[:, :, slice_idx]\n    current_sgig = sgig_volume[:, :, slice_idx]\n\n    # Select non-background pixels\n    positive_pixels = current_flair[current_flair > 0]\n\n    # Create an approximate brain mask\n    current_brain_mask = current_flair > np.percentile(\n        positive_pixels,\n        5\n    )\n\n    # Extract SmoothGrad-IG values inside the brain\n    brain_values = current_sgig[current_brain_mask]\n\n    # Compute the threshold corresponding to the top 10% brain attribution values\n    current_threshold = np.percentile(\n        brain_values,\n        90\n    )\n\n    # Create a thresholded attribution map.\n    # Values outside the brain or below the threshold are set to NaN,\n    # so they become transparent in the overlay.\n    current_thresholded = np.where(\n        current_brain_mask\n        & (current_sgig >= current_threshold),\n        current_sgig,\n        np.nan\n    )\n\n    # Create a colormap where NaN values are transparent\n    current_cmap = plt.cm.jet.copy()\n    current_cmap.set_bad(alpha=0)\n\n    # Show the original FLAIR slice\n    ax.imshow(\n        current_flair,\n        cmap=\"gray\"\n    )\n\n    # Overlay the strongest SmoothGrad-IG attributions\n    ax.imshow(\n        current_thresholded,\n        cmap=current_cmap,\n        alpha=0.65,\n        vmin=current_threshold,\n        vmax=1\n    )\n\n    ax.set_title(\n        f\"Slice {slice_idx}\\n\"\n        f\"Brain score: {brain_slice_scores[slice_idx]:.2f}\"\n    )\n\n    ax.axis(\"off\")\n\nplt.suptitle(\n    \"Slices with highest brain-specific SmoothGrad-IG attribution\",\n    fontsize=16\n)\n\nplt.tight_layout()\nplt.show()\n\nprint(\n    \"Selected slices with brain-specific score:\",\n    top_brain_slices\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define Gradient x Input\n\nimport numpy as np\nimport tensorflow as tf\nfrom scipy.ndimage import gaussian_filter\n\n\ndef gradient_x_input(\n    model,\n    inputs,\n    target_class=None,\n    smooth_sigma=1.0,\n    use_absolute=True\n):\n    \"\"\"\n    Compute Gradient x Input attribution maps for a Keras model.\n\n    Parameters\n    ----------\n    model : tf.keras.Model\n        Trained Keras model used for prediction.\n\n    inputs : list or tuple\n        Model inputs for one patient.\n        In this notebook, the first input corresponds to the FLAIR sequence.\n\n    target_class : int or None\n        Target class for the explanation.\n        If None, the predicted class is explained automatically.\n\n    smooth_sigma : float\n        Standard deviation of the Gaussian filter used to smooth the attribution maps.\n        Set to 0 to disable smoothing.\n\n    use_absolute : bool\n        If True, the absolute attribution values are used.\n        This highlights the regions with the strongest contribution regardless of sign.\n\n    Returns\n    -------\n    processed_maps : list of numpy.ndarray\n        One normalized attribution map for each model input.\n    \"\"\"\n\n    # Ensure that the inputs are provided as a list\n    if not isinstance(inputs, (list, tuple)):\n        inputs = [inputs]\n\n    # Convert all inputs to TensorFlow tensors\n    inputs = [\n        tf.convert_to_tensor(x, dtype=tf.float32)\n        for x in inputs\n    ]\n\n    with tf.GradientTape() as tape:\n\n        # Watch all input tensors\n        for x in inputs:\n            tape.watch(x)\n\n        # Forward pass\n        predictions = model(\n            inputs,\n            training=False\n        )\n\n        # Binary sigmoid output case\n        if predictions.shape[-1] == 1:\n\n            probability = predictions[:, 0]\n\n            # If no target class is provided, explain the predicted class\n            if target_class is None:\n                target_class = int(probability[0] >= 0.5)\n\n            # For class 1, use p\n            # For class 0, use 1 minus p\n            if target_class == 1:\n                score = probability\n            else:\n                score = 1.0 - probability\n\n        # Softmax output case\n        else:\n\n            if target_class is None:\n                target_class = tf.argmax(predictions[0])\n\n            score = predictions[:, target_class]\n\n    # Compute gradients of the selected score with respect to the model inputs\n    gradients = tape.gradient(\n        score,\n        inputs\n    )\n\n    processed_maps = []\n\n    for x, grad in zip(inputs, gradients):\n\n        # Gradient x Input\n        attribution = x * grad\n\n        # Convert to NumPy and remove batch dimension\n        attribution = attribution.numpy()[0]\n\n        # If there is a final channel dimension, remove it by averaging\n        if attribution.ndim == 4:\n            attribution = np.mean(\n                attribution,\n                axis=-1\n            )\n\n        # Use absolute values if required\n        if use_absolute:\n            attribution = np.abs(attribution)\n\n        # Optional spatial smoothing\n        if smooth_sigma is not None and smooth_sigma > 0:\n            attribution = gaussian_filter(\n                attribution,\n                sigma=smooth_sigma\n            )\n\n        # Normalize between 0 and 1\n        attribution = (\n            attribution - attribution.min()\n        ) / (\n            attribution.max() - attribution.min() + 1e-8\n        )\n\n        processed_maps.append(attribution)\n\n    return processed_maps\n\n\nprint(\"Gradient x Input function defined successfully.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Compute Gradient x Input on the selected correctly classified patient.\n\ngxi_maps = gradient_x_input(\n    model=xai_model,\n    inputs=xai_inputs,\n    target_class=target_class,\n    smooth_sigma=1.0,\n    use_absolute=True\n)\n\n# In this notebook, the first input corresponds to the FLAIR sequence.\ngxi_flair = gxi_maps[0]\n\nprint(\"Gradient x Input computed successfully.\")\nprint(\"Number of attribution maps obtained:\", len(gxi_maps))\nprint(\"FLAIR Gradient x Input shape:\", gxi_flair.shape)\nprint(\"Minimum value:\", gxi_flair.min())\nprint(\"Maximum value:\", gxi_flair.max())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Resize and visualize the Gradient x Input map on the original FLAIR MRI volume.\n\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# Original FLAIR volume\nflair_volume = xai_inputs[0][0, :, :, :, 0]\n\n# Gradient x Input map for the FLAIR input\ngxi_volume = gxi_flair\n\n# Select the same slice already used for XAI visualization\nflair_slice = flair_volume[:, :, best_slice]\ngxi_slice = gxi_volume[:, :, best_slice]\n\n# Create an approximate brain mask to exclude the black MRI background\npositive_pixels = flair_slice[flair_slice > 0]\n\nbrain_mask = flair_slice > np.percentile(\n    positive_pixels,\n    5\n)\n\n# Extract Gradient x Input values only inside the brain region\ngxi_brain_values = gxi_slice[brain_mask]\n\n# Keep only the top 10% Gradient x Input values inside the brain\ngxi_threshold = np.percentile(\n    gxi_brain_values,\n    90\n)\n\ngxi_thresholded = np.where(\n    brain_mask & (gxi_slice >= gxi_threshold),\n    gxi_slice,\n    np.nan\n)\n\n# Create a colormap where NaN values are transparent\ntransparent_cmap = plt.cm.jet.copy()\ntransparent_cmap.set_bad(alpha=0)\n\n# Visualization\nfig, axes = plt.subplots(\n    1,\n    3,\n    figsize=(15, 5)\n)\n\n# Original FLAIR image\naxes[0].imshow(flair_slice, cmap=\"gray\")\naxes[0].set_title(f\"Original FLAIR\\nSlice {best_slice}\")\naxes[0].axis(\"off\")\n\n# Gradient x Input heatmap\naxes[1].imshow(\n    gxi_slice,\n    cmap=\"jet\",\n    vmin=0,\n    vmax=1\n)\naxes[1].set_title(\"Gradient x Input\\nFLAIR\")\naxes[1].axis(\"off\")\n\n# Gradient x Input overlay on FLAIR\naxes[2].imshow(flair_slice, cmap=\"gray\")\naxes[2].imshow(\n    gxi_thresholded,\n    cmap=transparent_cmap,\n    alpha=0.65,\n    vmin=gxi_threshold,\n    vmax=1\n)\naxes[2].set_title(\"Gradient x Input overlay\\nwith brain mask\")\naxes[2].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n\nprint(\"Displayed slice:\", best_slice)\nprint(\"Gradient x Input threshold:\", gxi_threshold)\nprint(\"Gradient x Input minimum value:\", gxi_slice.min())\nprint(\"Gradient x Input maximum value:\", gxi_slice.max())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visual comparison between SmoothGrad Integrated Gradients and Gradient x Input\n# on the same FLAIR slice.\n\nimport matplotlib.pyplot as plt\nimport numpy as np\n\n# Slice to compare\nflair_slice = flair_volume[:, :, best_slice]\nsgig_slice = sgig_volume[:, :, best_slice]\ngxi_slice = gxi_flair[:, :, best_slice]\n\n# Brain mask\npositive_pixels = flair_slice[flair_slice > 0]\n\nbrain_mask = flair_slice > np.percentile(\n    positive_pixels,\n    5\n)\n\n# SmoothGrad Integrated Gradients: keep the top 10% of attribution values inside the brain\nsgig_brain_values = sgig_slice[brain_mask]\n\nsgig_threshold = np.percentile(\n    sgig_brain_values,\n    90\n)\n\nsgig_thresholded = np.where(\n    brain_mask & (sgig_slice >= sgig_threshold),\n    sgig_slice,\n    np.nan\n)\n\n# Gradient x Input: keep the top 10% of attribution values inside the brain\ngxi_brain_values = gxi_slice[brain_mask]\n\ngxi_threshold = np.percentile(\n    gxi_brain_values,\n    90\n)\n\ngxi_thresholded = np.where(\n    brain_mask & (gxi_slice >= gxi_threshold),\n    gxi_slice,\n    np.nan\n)\n\n# Transparent colormap for thresholded maps\ntransparent_cmap = plt.cm.jet.copy()\ntransparent_cmap.set_bad(alpha=0)\n\nfig, axes = plt.subplots(\n    1,\n    4,\n    figsize=(20, 5)\n)\n\n# Original FLAIR\naxes[0].imshow(\n    flair_slice,\n    cmap=\"gray\"\n)\naxes[0].set_title(\n    f\"Original FLAIR\\nSlice {best_slice}\"\n)\naxes[0].axis(\"off\")\n\n# SmoothGrad Integrated Gradients overlay\naxes[1].imshow(\n    flair_slice,\n    cmap=\"gray\"\n)\naxes[1].imshow(\n    sgig_thresholded,\n    cmap=transparent_cmap,\n    alpha=0.65,\n    vmin=sgig_threshold,\n    vmax=1\n)\naxes[1].set_title(\n    \"SmoothGrad-IG\\nTop 10%\"\n)\naxes[1].axis(\"off\")\n\n# Gradient x Input overlay\naxes[2].imshow(\n    flair_slice,\n    cmap=\"gray\"\n)\naxes[2].imshow(\n    gxi_thresholded,\n    cmap=transparent_cmap,\n    alpha=0.65,\n    vmin=gxi_threshold,\n    vmax=1\n)\naxes[2].set_title(\n    \"Gradient x Input\\nTop 10%\"\n)\naxes[2].axis(\"off\")\n\n# Continuous Gradient x Input map\naxes[3].imshow(\n    flair_slice,\n    cmap=\"gray\"\n)\naxes[3].imshow(\n    gxi_slice,\n    cmap=\"jet\",\n    alpha=0.45,\n    vmin=0,\n    vmax=1\n)\naxes[3].set_title(\n    \"Continuous\\nGradient x Input\"\n)\naxes[3].axis(\"off\")\n\nplt.tight_layout()\nplt.show()\n\nprint(\"Compared slice:\", best_slice)\nprint(\"SmoothGrad-IG threshold:\", sgig_threshold)\nprint(\"Gradient x Input threshold:\", gxi_threshold)\nprint(\"Gradient x Input minimum value:\", gxi_slice.min())\nprint(\"Gradient x Input maximum value:\", gxi_slice.max())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot Training History (Last Fold)\nif last_fold_history:\n    fig, axes = plt.subplots(1, 2, figsize=(15, 5))\n    metrics = [\"loss\", \"auc\"]\n    titles = [\"Binary Crossentropy Loss\", \"ROC AUC Score\"]\n    \n    for i, (metric, title) in enumerate(zip(metrics, titles)):\n        axes[i].plot(last_fold_history.history[metric], label=\"Train\")\n        axes[i].plot(last_fold_history.history[f\"val_{metric}\"], label=\"Validation\")\n        axes[i].set_title(title, fontsize=12, fontweight='bold')\n        axes[i].set_xlabel(\"Epochs\")\n        axes[i].set_ylabel(metric.upper())\n        axes[i].legend()\n        axes[i].grid(True, alpha=0.3)\n        \n    plt.tight_layout()\n    plt.show()\nelse:\n    print(\"No training history found. Run the training cell first.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <span style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\" id=\"section_4\">Test of trained final model</span>\n\nThe best model was saved during training. We will therefore load it and test the predictions on the file and the submission images.","metadata":{}},{"cell_type":"code","source":"# Dynamically build test DataFrame by scanning the test directory\ntest_dir = os.path.join(input_path, \"test\")\ntest_patients = [d for d in os.listdir(test_dir) if os.path.isdir(os.path.join(test_dir, d))]\n\ntest_df = pd.DataFrame({\"BraTS21ID\": [int(p) for p in test_patients]})\ntest_df = test_df.sort_values(\"BraTS21ID\").reset_index(drop=True)\n\nprint(f\"Successfully created test_df with {len(test_df)} patients.\")\ntest_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize test dataset generator (is_train=False disables label loading)\ntest_generator = Dataset(test_df, is_train=False, batch_size=1, shuffle=False, augment=False)\nprint(f\"Test generator initialized. Total batches: {len(test_generator)}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify data shapes match model expectations (Optional Debug)\nsample_batch = test_generator[0]\nprint(f\"Batch type: {type(sample_batch)}\")\nprint(f\"Modalities count: {len(sample_batch)}\")\nprint(f\"Single modality shape: {sample_batch[0].shape}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best saved model and run inference\nfrom tensorflow.keras.models import load_model\nmodel_path = \"best_fold_1.h5\"\nif os.path.exists(model_path):\n    print(f\"Loading best model from: {model_path}\")\n    inference_model = load_model(model_path, compile=False)\n    inference_model.compile(optimizer=\"adam\", loss=\"binary_crossentropy\",\n                            metrics=[keras.metrics.AUC(name=\"auc\")])\nelse:\n    print(\"Warning: Saved model not found. Using last trained model in memory.\")\n    inference_model = cnn_model if 'cnn_model' in locals() else create_cnn_model()\n# Direct inference — bypasses Dataset/Keras3 compatibility issue entirely\nprint(\"Generating predictions...\")\nall_preds = []\nfor pid in tqdm(test_df['BraTS21ID'].values, desc='Predicting'):\n    pid_str = str(pid).zfill(5)\n    batch_modalities = []\n    for modality in SCAN_CATEGORIES:\n        vol = load_volume(pid_str, modality, 'test')\n        batch_modalities.append(np.expand_dims(vol[np.newaxis], axis=-1))\n    central_slices = [m[:, :, :, m.shape[3] // 2, :] for m in batch_modalities]\n    batch_slices_2d = np.concatenate(central_slices, axis=-1)\n    inputs = batch_modalities + [batch_slices_2d]\n    pred = inference_model.predict(inputs, verbose=0)\n    all_preds.append(float(pred[0, 0]))\npreds_flat = np.array(all_preds, dtype=np.float32)\nprint(f\"Predictions shape: {preds_flat.shape}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Inspect prediction range and values\nprint(f\"Prediction range: [{preds_flat.min():.4f}, {preds_flat.max():.4f}]\")\nprint(\"First 5 predictions:\", preds_flat[:5])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Construct final submission DataFrame\nsubmission = pd.DataFrame({\n    \"BraTS21ID\": test_df[\"BraTS21ID\"],\n    \"MGMT_value\": preds_flat\n})\nprint(\"Submission DataFrame preview:\")\nsubmission.head(10)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save submission file for Kaggle upload\nsubmission.to_csv(\"submission.csv\", index=False)\nprint(\"✅ submission.csv saved successfully!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The submission.csv file will be used for the competition *(evaluation with AUC under ROC curve)*. We can also look at the **distribution of the predicted probabilities** :","metadata":{}},{"cell_type":"code","source":"# Visualize prediction distribution\nplt.figure(figsize=(8, 6))\nplt.hist(submission[\"MGMT_value\"], bins=30, edgecolor=\"black\", color=\"steelblue\")\nplt.axvline(x=0.5, color=\"red\", linestyle=\"--\", label=\"Decision Threshold (0.5)\")\nplt.title(\"Distribution of Predicted MGMT Probabilities\", fontsize=14, fontweight=\"bold\")\nplt.xlabel(\"Predicted Probability\")\nplt.ylabel(\"Frequency\")\nplt.legend()\nplt.grid(axis=\"y\", alpha=0.3)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# <span style=\"color:#0b0a2d; font-size:24px; text-transform: uppercase; font-weight:bold\" id=\"section_5\">Try another approach: Transfer Learning</span>\n\nIn order to complete these models, we will try another approach using **Transfer Learning methods**. We will use a pre-trained deep model to detect the features *(like EfficientNet ...)* and an LSTM layer for the final classification on the matrices obtained. This approach is available in the Notebook :\n\n<span style=\"font-size:18px\">[🧠Brain Tumor - Transfert Learning MRI - All MRI](https://www.kaggle.com/michaelfumery/brain-tumor-transfert-learning-mri-all-mri/)</span>\n\n<span style=\"color:red; font-size:18px\">Don't forget to **upvote** if this Notebook helped you!</span>","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}