{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":2425289,"sourceType":"datasetVersion","datasetId":1467572},{"sourceId":11030525,"sourceType":"datasetVersion","datasetId":6869765}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"!pip install efficientnet_pytorch\n!pip install -U albumentations\n\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:00:41.861446Z","iopub.execute_input":"2025-03-26T14:00:41.861749Z","iopub.status.idle":"2025-03-26T14:00:53.139744Z","shell.execute_reply.started":"2025-03-26T14:00:41.861719Z","shell.execute_reply":"2025-03-26T14:00:53.138661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport re\nimport time\nimport random\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.nn import functional as F\nfrom tqdm import tqdm\nfrom sklearn.metrics import confusion_matrix\n\nfrom efficientnet_pytorch import EfficientNet\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import StratifiedKFold\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:00:53.140769Z","iopub.execute_input":"2025-03-26T14:00:53.141055Z","iopub.status.idle":"2025-03-26T14:00:53.165037Z","shell.execute_reply.started":"2025-03-26T14:00:53.141029Z","shell.execute_reply":"2025-03-26T14:00:53.164232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Analysis","metadata":{}},{"cell_type":"code","source":"# Load train labels\nlabels_df = pd.read_csv(\"/kaggle/input/brats21-t-labels/train_labels.csv\")\nlabels_df[\"BraTS21ID\"] = labels_df[\"BraTS21ID\"].astype(str).str.zfill(5)  # Ensure ID format matches folders\n\n# Define dataset path\ntrain_path = \"/kaggle/input/rsna-miccai-png/train/\"\n\n# Get only existing patient IDs from the dataset\nexisting_patient_ids = sorted(os.listdir(train_path))  # Get actual patient folders\nbad_ids = ['00109', '00123', '00709']  # IDs to remove\npatient_ids = [pid for pid in existing_patient_ids if pid not in bad_ids]\n\n# Initialize lists for image counts\nFlair_files, T1w_files, T1wCE_files, T2w_files = [], [], [], []\nmissing_patients = []\n\n# Iterate over the actual patient folders\nfor patient_id in tqdm(patient_ids):\n    patient_path = os.path.join(train_path, patient_id)\n    \n    # Count images for each type, set to 0 if the folder is missing\n    Flair_files.append(len(os.listdir(os.path.join(patient_path, \"FLAIR\"))) if os.path.exists(os.path.join(patient_path, \"FLAIR\")) else 0)\n    T1w_files.append(len(os.listdir(os.path.join(patient_path, \"T1w\"))) if os.path.exists(os.path.join(patient_path, \"T1w\")) else 0)\n    T1wCE_files.append(len(os.listdir(os.path.join(patient_path, \"T1wCE\"))) if os.path.exists(os.path.join(patient_path, \"T1wCE\")) else 0)\n    T2w_files.append(len(os.listdir(os.path.join(patient_path, \"T2w\"))) if os.path.exists(os.path.join(patient_path, \"T2w\")) else 0)\n\n# Create DataFrame\nmri_counts_df = pd.DataFrame({\n    \"Patient ID\": patient_ids,\n    \"Flair\": Flair_files,\n    \"T1w\": T1w_files,\n    \"T1wCE\": T1wCE_files,\n    \"T2w\": T2w_files\n})\n\n# Calculate total number of images per MRI type\ntotal_flair = sum(Flair_files)\ntotal_t1w = sum(T1w_files)\ntotal_t1wce = sum(T1wCE_files)\ntotal_t2w = sum(T2w_files)\n\n# Print dataset overview\nprint(f\"Total Patients in Dataset: {len(patient_ids)}\")\nprint(f\"Total MRI Scans Available:\")\nprint(f\"- FLAIR: {total_flair} scans\")\nprint(f\"- T1w: {total_t1w} scans\")\nprint(f\"- T1wCE: {total_t1wce} scans\")\nprint(f\"- T2w: {total_t2w} scans\")\n\n# Print missing patients (those in CSV but not in dataset)\ncsv_patients = set(labels_df[\"BraTS21ID\"])\ndataset_patients = set(patient_ids)\nmissing_from_dataset = csv_patients - dataset_patients\nif missing_from_dataset:\n    print(f\"Patients in CSV but missing from dataset: {missing_from_dataset}\")\n\n# Display dataset structure\nprint(mri_counts_df.head(10))\n\n# Visualization: Distribution of images per MRI type\nplt.figure(figsize=(12, 6))\nplt.hist([Flair_files, T1w_files, T1wCE_files, T2w_files], bins=20, label=[\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"], alpha=0.7)\nplt.xlabel(\"Number of Images per Patient\")\nplt.ylabel(\"Frequency\")\nplt.title(\"Distribution of MRI Image Counts Across Patients\")\nplt.legend()\nplt.show()\n\n# Bar plot: Total MRI scans per type\nplt.figure(figsize=(8, 5))\nplt.bar([\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"], [total_flair, total_t1w, total_t1wce, total_t2w], color=['blue', 'orange', 'green', 'red'])\nplt.xlabel(\"MRI Type\")\nplt.ylabel(\"Total Number of Scans\")\nplt.title(\"Total MRI Scans Per Type\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:00:53.165784Z","iopub.execute_input":"2025-03-26T14:00:53.166022Z","iopub.status.idle":"2025-03-26T14:01:17.495776Z","shell.execute_reply.started":"2025-03-26T14:00:53.166003Z","shell.execute_reply":"2025-03-26T14:01:17.495109Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## modalities disribution","metadata":{}},{"cell_type":"code","source":"# Set figure size\nplt.figure(figsize=(12, 8))\n\n# Skip 'Patient ID' column\nmri_types = [\"Flair\", \"T1w\", \"T1wCE\", \"T2w\"]\n\n# Plot each MRI type\nfor i, mri_type in enumerate(mri_types):\n    plt.subplot(2, 2, i + 1)\n    colors = [\"blue\", \"orange\", \"green\", \"red\"]\n    plt.bar(mri_counts_df[\"Patient ID\"], mri_counts_df[mri_type], color=colors[i], alpha=0.7)\n    plt.xlabel(\"Patient ID\")\n    plt.ylabel(f\"Number of {mri_type} scans\")\n    plt.title(f\"Scans per Patient ({mri_type})\")\n    plt.xticks([])  # Hide x-axis labels for clarity\n    plt.grid(axis=\"y\", linestyle=\"--\", alpha=0.5)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:17.497282Z","iopub.execute_input":"2025-03-26T14:01:17.497509Z","iopub.status.idle":"2025-03-26T14:01:20.950476Z","shell.execute_reply.started":"2025-03-26T14:01:17.497490Z","shell.execute_reply":"2025-03-26T14:01:20.949602Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## stats Min Max","metadata":{}},{"cell_type":"code","source":"# Compute statistics for each MRI type based on mri_counts_df\nstats = {\n    \"MRI Type\": [\"Flair\", \"T1w\", \"T1wCE\", \"T2w\"],\n    \"Min\": [mri_counts_df['Flair'].min(), mri_counts_df['T1w'].min(), mri_counts_df['T1wCE'].min(), mri_counts_df['T2w'].min()],\n    \"Max\": [mri_counts_df['Flair'].max(), mri_counts_df['T1w'].max(), mri_counts_df['T1wCE'].max(), mri_counts_df['T2w'].max()],\n    \"Avg\": [mri_counts_df['Flair'].mean(), mri_counts_df['T1w'].mean(), mri_counts_df['T1wCE'].mean(), mri_counts_df['T2w'].mean()]\n}\n\n# Convert to DataFrame for easy visualization\nstats_df = pd.DataFrame(stats)\n\n# Display stats\nprint(stats_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:20.951843Z","iopub.execute_input":"2025-03-26T14:01:20.952116Z","iopub.status.idle":"2025-03-26T14:01:20.960243Z","shell.execute_reply.started":"2025-03-26T14:01:20.952094Z","shell.execute_reply":"2025-03-26T14:01:20.959270Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# MRI Type Paths Collection","metadata":{}},{"cell_type":"markdown","source":"## Path Collection and sorting","metadata":{}},{"cell_type":"code","source":"# Define dataset path\ntrain_path = \"/kaggle/input/rsna-miccai-png/train/\"\n\n# ✅ Function to extract sequence number for sorting\ndef extract_number(filename):\n    match = re.search(r\"(\\d+)\", filename)\n    return int(match.group()) if match else float(\"inf\")  # Default to infinity if no number found\n\ndef natural_sort(image_paths):\n    \"\"\"Sort image paths numerically based on extracted numbers.\"\"\"\n    return sorted(image_paths, key=lambda x: extract_number(os.path.basename(x)))\n\n# ✅ Dictionary to store MRI image paths for each patient\npatient_image_paths = {}\n\n# ✅ Collect all image paths and sort them\nfor patient_id in tqdm(patient_ids, desc=\"Collecting image paths\", leave=True):\n    patient_image_paths[patient_id] = {}\n\n    for mri_type in [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]:\n        mri_path = os.path.join(train_path, patient_id, mri_type)\n\n        if os.path.exists(mri_path):\n            images = [os.path.join(mri_path, img) for img in os.listdir(mri_path)]\n            sorted_images = natural_sort(images)  # Apply corrected sorting\n            patient_image_paths[patient_id][mri_type] = sorted_images\n        else:\n            patient_image_paths[patient_id][mri_type] = []  # No images found\n\nprint(f\"\\n✅ Collected MRI image paths for {len(patient_ids)} patients.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:20.961129Z","iopub.execute_input":"2025-03-26T14:01:20.961444Z","iopub.status.idle":"2025-03-26T14:01:23.112441Z","shell.execute_reply.started":"2025-03-26T14:01:20.961412Z","shell.execute_reply":"2025-03-26T14:01:23.111599Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Debuggin the Path Collection","metadata":{}},{"cell_type":"code","source":"# ✅ Debugging: Check sorted file paths for a few patients\nfor patient_id in list(patient_image_paths.keys())[:1]:  # Show first 3 patients\n    for mri_type in [\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\"]:\n        images = patient_image_paths[patient_id][mri_type]\n        \n        print(f\"\\n🔍 **Patient {patient_id}, MRI Type {mri_type}**\")\n        print(f\"📁 Found {len(images)} images. Showing first 10:\")\n\n        for img in images[:10]:  # Print first 10 images to verify sorting\n            print(img)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:23.113287Z","iopub.execute_input":"2025-03-26T14:01:23.113575Z","iopub.status.idle":"2025-03-26T14:01:23.128201Z","shell.execute_reply.started":"2025-03-26T14:01:23.113542Z","shell.execute_reply":"2025-03-26T14:01:23.127302Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"## CFG essential parts","metadata":{}},{"cell_type":"code","source":"# Set device to GPU if available, otherwise use CPU\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Set a fixed seed for reproducibility\nseed = 123\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\n# Apply the seed setting function\nseed_everything(seed)\n\n# Configuration class for model parameters\nclass CFG:\n    img_size = 256       # Image resizing dimensions\n    n_frames = 14        # Number of frames per patient used in training\n    \n    cnn_features = 256   # Features extracted from CNN backbone\n    lstm_hidden = 32     # Hidden layer size for LSTM (if used)\n    \n    n_fold = 5           # Number of cross-validation folds\n    n_epochs = 20        # Total training epochs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:23.129146Z","iopub.execute_input":"2025-03-26T14:01:23.129432Z","iopub.status.idle":"2025-03-26T14:01:23.202399Z","shell.execute_reply.started":"2025-03-26T14:01:23.129406Z","shell.execute_reply":"2025-03-26T14:01:23.201630Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## model arch","metadata":{}},{"cell_type":"code","source":"class CNN(nn.Module):\n    def __init__(self, dropout_rate=0.3):\n        super().__init__()\n        self.map = nn.Conv2d(in_channels=1, out_channels=3, kernel_size=1)  # Convert grayscale to RGB\n        \n        # Load the pretrained EfficientNet\n        self.net = EfficientNet.from_pretrained(\"efficientnet-b0\")\n        \n        # Replace the fully connected layer with a smaller one + dropout\n        n_features = self.net._fc.in_features\n        self.net._fc = nn.Sequential(\n            nn.Dropout(dropout_rate),\n            nn.Linear(in_features=n_features, out_features=CFG.cnn_features, bias=True)\n        )\n        \n    def forward(self, x):\n        x = F.relu(self.map(x))\n        return self.net(x)\n\nclass Model(nn.Module):\n    def __init__(self, dropout_rate=0.3):\n        super(Model, self).__init__()\n        self.cnn = CNN(dropout_rate=dropout_rate)\n        \n        # Add dropout to LSTM for regularization\n        self.rnn = nn.LSTM(\n            input_size=CFG.cnn_features, \n            hidden_size=CFG.lstm_hidden, \n            num_layers=2, \n            batch_first=True,\n            dropout=dropout_rate if 2 > 1 else 0  # Only apply dropout between LSTM layers\n        )\n        \n        # Add dropout before final classification\n        self.dropout = nn.Dropout(dropout_rate)\n        self.fc = nn.Linear(CFG.lstm_hidden, 1, bias=True)\n        \n    def forward(self, x):\n        batch_size, timesteps, C, H, W = x.size()\n        c_in = x.view(batch_size * timesteps, C, H, W)\n        c_out = self.cnn(c_in)\n        r_in = c_out.view(batch_size, timesteps, -1)\n        output, (hn, cn) = self.rnn(r_in)\n        \n        # Apply dropout before classification\n        x = self.dropout(hn[-1])\n        return self.fc(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:23.203150Z","iopub.execute_input":"2025-03-26T14:01:23.203379Z","iopub.status.idle":"2025-03-26T14:01:23.212354Z","shell.execute_reply.started":"2025-03-26T14:01:23.203349Z","shell.execute_reply":"2025-03-26T14:01:23.211690Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model testing","metadata":{}},{"cell_type":"code","source":"model = Model()\nx = torch.zeros((5, 14, 1, 256, 256))  # (Batch, Timesteps, Channels, Height, Width)\nout = model(x)\nprint(out,x.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:23.213025Z","iopub.execute_input":"2025-03-26T14:01:23.213265Z","iopub.status.idle":"2025-03-26T14:01:34.293353Z","shell.execute_reply.started":"2025-03-26T14:01:23.213246Z","shell.execute_reply":"2025-03-26T14:01:34.292514Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Prepocessing","metadata":{}},{"cell_type":"markdown","source":"## data LOading function","metadata":{}},{"cell_type":"code","source":"# ✅ Load and preprocess MRI slices\ndef load_image(path):\n    image = cv2.imread(path, 0)  # Read in grayscale\n    if image is None:\n        return np.zeros((CFG.img_size, CFG.img_size))  # Return blank image if loading fails\n    \n    image = cv2.resize(image, (CFG.img_size, CFG.img_size)) / 255.0  # Resize and normalize\n    return image.astype('float32')\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:34.294128Z","iopub.execute_input":"2025-03-26T14:01:34.294351Z","iopub.status.idle":"2025-03-26T14:01:34.298488Z","shell.execute_reply.started":"2025-03-26T14:01:34.294332Z","shell.execute_reply":"2025-03-26T14:01:34.297730Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## The Best Selection Approach: (*select_best_central_images*)","metadata":{}},{"cell_type":"code","source":"def select_best_central_images(image_paths, n_frames=14, spread_threshold=100):\n    \"\"\"\n    Select the best n_frames (default 14) central images from a list of image paths.\n    \n    Args:\n        image_paths: List of paths to images\n        n_frames: Number of frames to select (default 14)\n        spread_threshold: Threshold for determining selection strategy\n        \n    Returns:\n        List of exactly n_frames selected image paths\n    \"\"\"\n    num_images = len(image_paths)\n    \n    # If we have fewer images than requested frames, duplicate some images to reach n_frames\n    if num_images < n_frames:\n        # Duplicate images to reach n_frames\n        selected_images = image_paths.copy()\n        # Keep duplicating the last image until we reach n_frames\n        while len(selected_images) < n_frames:\n            selected_images.append(image_paths[-1])\n        return selected_images\n    \n    # If we have exactly n_frames, return all\n    if num_images == n_frames:\n        return image_paths\n        \n    # Compute nonzero pixel count for each image\n    brain_pixel_counts = np.array([np.sum(cv2.imread(img, cv2.IMREAD_GRAYSCALE) > 0) for img in image_paths])\n    \n    # Find the slice with the largest brain area\n    best_idx = np.argmax(brain_pixel_counts)\n    \n    # Strategy 1: For smaller sequences, get evenly distributed slices around the best slice\n    if num_images <= spread_threshold:\n        # Calculate how many slices to take before and after the best slice\n        half_n = n_frames // 2\n        \n        # Calculate start and end indices, ensuring we don't go out of bounds\n        start_idx = max(0, best_idx - half_n)\n        end_idx = min(num_images, best_idx + half_n + (n_frames % 2))  # Add 1 more if n_frames is odd\n        \n        # Get the selected images based on calculated indices\n        selected_indices = list(range(start_idx, end_idx))\n        \n        # If we didn't get enough images, add more from one end or the other\n        if len(selected_indices) < n_frames:\n            deficit = n_frames - len(selected_indices)\n            if start_idx == 0:  # We hit the beginning, so add more from the end\n                additional_indices = list(range(end_idx, min(end_idx + deficit, num_images)))\n                selected_indices.extend(additional_indices)\n            else:  # We hit the end, so add more from the beginning\n                additional_indices = list(range(max(0, start_idx - deficit), start_idx))\n                selected_indices = additional_indices + selected_indices\n        \n        # If we still don't have enough (unlikely), take equally spaced images from the entire set\n        if len(selected_indices) < n_frames:\n            selected_indices = list(range(0, num_images, max(1, num_images // n_frames)))[:n_frames]\n            \n            # If still not enough (very unlikely), duplicate the last index\n            while len(selected_indices) < n_frames:\n                selected_indices.append(selected_indices[-1])\n    \n    # Strategy 2: For larger sequences, select evenly spaced slices with focus on the center region\n    else:\n        # Calculate indices for slices before the best slice\n        step_before = max(1, best_idx // (n_frames // 2))\n        before_indices = [best_idx - i * step_before for i in range(1, (n_frames // 2) + 1)]\n        before_indices = [i for i in before_indices if i >= 0]\n        \n        # Calculate indices for slices after the best slice\n        step_after = max(1, (num_images - best_idx - 1) // (n_frames // 2))\n        after_indices = [best_idx + i * step_after for i in range(1, (n_frames // 2) + 1)]\n        after_indices = [i for i in after_indices if i < num_images]\n        \n        # Combine all indices with the best index in the middle\n        selected_indices = before_indices + [best_idx] + after_indices\n        \n        # Sort indices and ensure we have exactly n_frames\n        selected_indices = sorted(selected_indices)\n        \n        # If we have too many indices, remove those furthest from the best_idx\n        if len(selected_indices) > n_frames:\n            # Sort by distance from best_idx\n            sorted_by_distance = sorted(selected_indices, key=lambda x: abs(x - best_idx))\n            selected_indices = sorted(sorted_by_distance[:n_frames])\n        \n        # If we don't have enough indices, add more from remaining slices\n        if len(selected_indices) < n_frames:\n            remaining_indices = list(set(range(num_images)) - set(selected_indices))\n            # Sort remaining indices by distance from best_idx\n            additional_indices = sorted(remaining_indices, key=lambda x: abs(x - best_idx))\n            # Add as many as needed\n            additional_indices = additional_indices[:n_frames - len(selected_indices)]\n            selected_indices.extend(additional_indices)\n            selected_indices = sorted(selected_indices)\n    \n    # Final check to ensure we have exactly n_frames\n    if len(selected_indices) > n_frames:\n        # Keep those closest to best_idx\n        selected_indices = sorted(selected_indices, key=lambda x: abs(x - best_idx))[:n_frames]\n        selected_indices = sorted(selected_indices)  # Sort back to ascending order\n    \n    # In the extremely unlikely case we still don't have enough, duplicate the last one\n    while len(selected_indices) < n_frames:\n        selected_indices.append(selected_indices[-1])\n    \n    # Now get the actual images\n    selected_images = [image_paths[i] for i in selected_indices]\n    \n    # Final verification\n    assert len(selected_images) == n_frames, f\"Failed to select exactly {n_frames} images, got {len(selected_images)}\"\n    \n    return selected_images","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:34.301341Z","iopub.execute_input":"2025-03-26T14:01:34.301564Z","iopub.status.idle":"2025-03-26T14:01:34.320653Z","shell.execute_reply.started":"2025-03-26T14:01:34.301545Z","shell.execute_reply":"2025-03-26T14:01:34.319962Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Testing the Selection aprraoch","metadata":{}},{"cell_type":"code","source":"# Test the selection function for a specific patient and MRI type\npatient_id = \"00002\"  # Change to any patient ID\nmri_type = \"T2w\"  # Change to FLAIR, T1w, T1wCE, or T2w\n\nif patient_id in patient_image_paths and mri_type in patient_image_paths[patient_id]:\n    image_paths = patient_image_paths[patient_id][mri_type]\n    \n    print(f\"\\n🔍 **Patient {patient_id}, MRI Type {mri_type}**\")\n    print(f\"📁 Found {len(image_paths)} images. Showing first 10:\")\n    for img in image_paths[:10]:\n        print(img)\n\n    # Call selection function with debug prints inside\n    selected_images = select_best_central_images(image_paths)\n\n    print(\"\\n✅ **Final Selected Images:**\")\n    for img in selected_images:\n        print(img)\n\nelse:\n    print(f\"❌ Patient {patient_id} or MRI type {mri_type} not found.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:34.321938Z","iopub.execute_input":"2025-03-26T14:01:34.322175Z","iopub.status.idle":"2025-03-26T14:01:37.143381Z","shell.execute_reply.started":"2025-03-26T14:01:34.322157Z","shell.execute_reply":"2025-03-26T14:01:37.142615Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Ploting the apraoch","metadata":{}},{"cell_type":"code","source":"\ndef plot_best_central(patient_id, mri_type, n_frames=14):\n    \"\"\"\n    Plots `n_frames` selected MRI images for a given patient and MRI type, displaying filenames below each image.\n    \n    Args:\n        patient_id (str): The ID of the patient.\n        mri_type (str): MRI type ('FLAIR', 'T1w', 'T1wCE', 'T2w').\n        n_frames (int): Number of images to display (default is 14).\n    \"\"\"\n    if patient_id not in patient_image_paths:\n        print(f\"❌ Patient {patient_id} not found.\")\n        return\n    \n    image_paths = patient_image_paths[patient_id].get(mri_type, [])\n    \n    if not image_paths:\n        print(f\"❌ No images found for Patient {patient_id} in {mri_type}.\")\n        return\n\n    # Select the best images based on the function\n    selected_images = select_best_central_images(image_paths, n_frames=n_frames)\n\n    if not selected_images:\n        print(\"❌ No images selected. Check selection function.\")\n        return\n\n    # Create a figure with 2 rows, 7 columns (14 images total)\n    fig, axes = plt.subplots(2, 7, figsize=(14, 6))\n    fig.suptitle(f\"Patient {patient_id} - {mri_type} Selected Slices\", fontsize=14)\n    \n    for idx, ax in enumerate(axes.flat):\n        if idx < len(selected_images):\n            img_path = selected_images[idx]\n            img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            ax.imshow(img, cmap='gray')\n\n            # Extract filename to display\n            filename = os.path.basename(img_path)\n            ax.set_title(filename, fontsize=8, pad=5)\n\n        ax.axis(\"off\")  # Hide axis\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:37.144206Z","iopub.execute_input":"2025-03-26T14:01:37.144425Z","iopub.status.idle":"2025-03-26T14:01:37.150761Z","shell.execute_reply.started":"2025-03-26T14:01:37.144405Z","shell.execute_reply":"2025-03-26T14:01:37.149978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example: Select a patient and an MRI type (\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\")\nplot_best_central(patient_id=\"00000\", mri_type=\"FLAIR\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:37.151507Z","iopub.execute_input":"2025-03-26T14:01:37.151708Z","iopub.status.idle":"2025-03-26T14:01:40.974423Z","shell.execute_reply.started":"2025-03-26T14:01:37.151690Z","shell.execute_reply":"2025-03-26T14:01:40.973575Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Another selction apraoch :  *uniform_temporal_subsample* & *select_patient_slices*","metadata":{}},{"cell_type":"code","source":"# ✅ Uniform subsampling function\ndef uniform_temporal_subsample(x, num_samples):\n    \"\"\"\n    Selects `num_samples` equispaced indices from `x`.\n\n    Args:\n        x (list): List of image paths.\n        num_samples (int): The number of equispaced samples to be selected.\n    Returns:\n        List: Selected image indices.\n    \"\"\"\n    t = len(x)\n    indices = torch.linspace(0, t - 1, num_samples)  # Generate evenly spaced indices\n    indices = torch.clamp(indices, 0, t - 1).long()  # Ensure indices are within bounds\n    return [i.item() for i in indices]\n\n# ✅ Function to get selected indices for each patient & MRI type\ndef select_patient_slices(patient_id, mri_type, num_samples=14):\n    \"\"\"\n    Given a patient ID and MRI type, selects 14 slices using uniform subsampling.\n\n    Args:\n        patient_id (str): Patient ID.\n        mri_type (str): MRI type (FLAIR, T1w, etc.).\n        num_samples (int): Number of slices to select.\n\n    Returns:\n        list: Selected image paths.\n    \"\"\"\n    image_paths = patient_image_paths.get(patient_id, {}).get(mri_type, [])\n    \n    if not image_paths:\n        return []  # No images available\n\n    selected_indices = uniform_temporal_subsample(image_paths, num_samples)\n    return [image_paths[i] for i in selected_indices]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:40.975522Z","iopub.execute_input":"2025-03-26T14:01:40.975847Z","iopub.status.idle":"2025-03-26T14:01:40.981551Z","shell.execute_reply.started":"2025-03-26T14:01:40.975818Z","shell.execute_reply":"2025-03-26T14:01:40.980772Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Plotting the UNiform aproach selction","metadata":{}},{"cell_type":"code","source":"def plot_uniform(patient_id, mri_type, num_samples=14):\n    \"\"\"\n    Plots `num_samples` MRI images in a 2x7 grid with filenames.\n\n    Args:\n        patient_id (str): Patient ID.\n        mri_type (str): MRI type ('FLAIR', 'T1w', 'T1wCE', 'T2w').\n        num_samples (int): Number of images to display (default is 14).\n    \"\"\"\n    if patient_id not in patient_image_paths:\n        print(f\"❌ Patient {patient_id} not found.\")\n        return\n    \n    image_paths = patient_image_paths[patient_id].get(mri_type, [])\n    \n    if not image_paths:\n        print(f\"❌ No images found for Patient {patient_id} in {mri_type}.\")\n        return\n\n    # Select images using uniform sampling\n    selected_indices = uniform_temporal_subsample(image_paths, num_samples)\n    selected_images = [image_paths[i] for i in selected_indices]\n\n    if not selected_images:\n        print(\"❌ No images selected. Check selection function.\")\n        return\n\n    # Create a figure with 2 rows, 7 columns (14 images total)\n    fig, axes = plt.subplots(2, 7, figsize=(14, 6))\n    fig.suptitle(f\"Patient {patient_id} - {mri_type} Uniform Selection\", fontsize=14)\n\n    for idx, ax in enumerate(axes.flat):\n        if idx < len(selected_images):\n            img_path = selected_images[idx]\n            img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            # img = load_image(img_path)\n            ax.imshow(img, cmap='gray')\n\n            # Extract filename to display\n            filename = os.path.basename(img_path)\n            ax.set_title(filename, fontsize=8, pad=5)\n\n        ax.axis(\"off\")  # Hide axis\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:40.982381Z","iopub.execute_input":"2025-03-26T14:01:40.982580Z","iopub.status.idle":"2025-03-26T14:01:41.002310Z","shell.execute_reply.started":"2025-03-26T14:01:40.982562Z","shell.execute_reply":"2025-03-26T14:01:41.001607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_uniform(patient_id=\"00000\", mri_type=\"FLAIR\", num_samples=14)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:41.003053Z","iopub.execute_input":"2025-03-26T14:01:41.003250Z","iopub.status.idle":"2025-03-26T14:01:42.065903Z","shell.execute_reply.started":"2025-03-26T14:01:41.003232Z","shell.execute_reply":"2025-03-26T14:01:42.064968Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Retriver","metadata":{}},{"cell_type":"code","source":"class DataRetriever(Dataset):\n    def __init__(self, paths, targets, mri_type=\"FLAIR\", selection_method=\"uniform\", transform=None):\n        \"\"\"\n        Dataset class for loading MRI images with selectable slice selection method.\n        Args:\n            paths (list): List of patient IDs.\n            targets (list): Labels for classification.\n            mri_type (str): MRI type to load (\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\").\n            selection_method (str): 'uniform' or 'best_central' (default: 'uniform').\n            transform (callable, optional): Image transformations.\n        \"\"\"\n        self.paths = paths\n        self.targets = targets\n        self.mri_type = mri_type\n        self.selection_method = selection_method  # Choose which method to use\n        self.transform = transform\n          \n    def __len__(self):\n        return len(self.paths)\n    \n    def read_video(self, vid_paths):\n        video = [load_image(path) for path in vid_paths]\n        if self.transform:\n            seed = random.randint(0, 99999)\n            for i in range(len(video)):\n                random.seed(seed)\n                video[i] = self.transform(image=video[i])[\"image\"]\n        \n        video = [torch.tensor(frame, dtype=torch.float32) for frame in video]\n        if len(video) == 0:\n            video = torch.zeros(CFG.n_frames, CFG.img_size, CFG.img_size)\n        else:\n            video = torch.stack(video)  # Shape: (T, H, W)\n        return video\n    \n    def __getitem__(self, index):\n        _id = self.paths[index]\n        # Fetch precomputed paths\n        image_paths = patient_image_paths.get(str(_id).zfill(5), {}).get(self.mri_type, [])\n        \n        # Select images based on method choice\n        num_samples = CFG.n_frames\n        if len(image_paths) < num_samples:\n            in_frames_path = image_paths\n        else:\n            if self.selection_method == \"uniform\":\n                sampled_ids = uniform_temporal_subsample(image_paths, num_samples)\n                in_frames_path = [image_paths[i] for i in sampled_ids]\n            elif self.selection_method == \"best_central\":\n                in_frames_path = select_best_central_images(image_paths, n_frames=num_samples)\n            else:\n                raise ValueError(f\"❌ Unknown selection method: {self.selection_method}\")\n                \n        # Get the video frames\n        channel = self.read_video(in_frames_path)  # Shape: (T, H, W)\n        \n        if channel.shape[0] == 0:\n            print(f\"⚠️ Empty MRI sequence for {self.mri_type}, filling with zeros.\")\n            channel = torch.zeros(num_samples, CFG.img_size, CFG.img_size)\n            \n        # Add channel dimension for CNN (T, 1, H, W)\n        channel = channel.unsqueeze(1)\n        \n        y = torch.tensor(self.targets[index], dtype=torch.float)\n        return {\"X\": channel.float(), \"y\": y}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:42.066912Z","iopub.execute_input":"2025-03-26T14:01:42.067167Z","iopub.status.idle":"2025-03-26T14:01:42.312302Z","shell.execute_reply.started":"2025-03-26T14:01:42.067146Z","shell.execute_reply":"2025-03-26T14:01:42.311247Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Testing the data reteriever ","metadata":{}},{"cell_type":"markdown","source":"#### basic infos","metadata":{}},{"cell_type":"code","source":"sample_dataset = DataRetriever(paths=patient_ids[:10], targets=[0] * 10, mri_type=\"FLAIR\", selection_method=\"uniform\")\nsample_data = sample_dataset[0]\n\nprint(f\"Shape of X: {sample_data['X'].shape}\")  # Should be (14, 1, 256, 256)\nprint(f\"Shape of y: {sample_data['y'].shape}\")  # Should be (1,)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:42.313114Z","iopub.execute_input":"2025-03-26T14:01:42.313337Z","iopub.status.idle":"2025-03-26T14:01:42.381054Z","shell.execute_reply.started":"2025-03-26T14:01:42.313318Z","shell.execute_reply":"2025-03-26T14:01:42.379982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ✅ Sample test for DataRetriever\nsample_dataset = DataRetriever(paths=patient_ids[:10], targets=[0]*10, mri_type=\"FLAIR\", selection_method=\"uniform\")\n\nprint(f\"Dataset length: {len(sample_dataset)}\")\n\n# Fetch a sample\nsample_data = sample_dataset[0]\nprint(f\"Input shape: {sample_data['X'].shape}\")  # Should be (14, H, W)\nprint(f\"Label: {sample_data['y'].shape}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:42.382012Z","iopub.execute_input":"2025-03-26T14:01:42.382316Z","iopub.status.idle":"2025-03-26T14:01:42.423392Z","shell.execute_reply.started":"2025-03-26T14:01:42.382284Z","shell.execute_reply":"2025-03-26T14:01:42.422717Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### comparison selction methods","metadata":{}},{"cell_type":"code","source":"# ✅ Load same patient with both selection methods\ndataset_uniform = DataRetriever(paths=patient_ids[:1], targets=[0], mri_type=\"FLAIR\", selection_method=\"uniform\")\ndataset_best = DataRetriever(paths=patient_ids[:1], targets=[0], mri_type=\"FLAIR\", selection_method=\"best_central\")\n\n# Fetch the MRI sequences\nsample_uniform = dataset_uniform[0][\"X\"]\nsample_best = dataset_best[0][\"X\"]\n\n# ✅ Fix subplot shape to 2 rows, 14 columns\nfig, axes = plt.subplots(2, 14, figsize=(20, 6))  # 2 rows (1 per method), 14 columns (one per image)\n\nfor i in range(14):\n    axes[0, i].imshow(sample_uniform[i,0].numpy(), cmap=\"gray\")\n    axes[1, i].imshow(sample_best[i,0].numpy(), cmap=\"gray\")\n    axes[0, i].axis(\"off\")\n    axes[1, i].axis(\"off\")\n\naxes[0, 0].set_ylabel(\"Uniform Sampling\")\naxes[1, 0].set_ylabel(\"Best Central\")\nplt.suptitle(\"Comparison: Uniform vs. Best Central Selection\")\nplt.show()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:42.424050Z","iopub.execute_input":"2025-03-26T14:01:42.424261Z","iopub.status.idle":"2025-03-26T14:01:44.051233Z","shell.execute_reply.started":"2025-03-26T14:01:42.424242Z","shell.execute_reply":"2025-03-26T14:01:44.050369Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Transformation *(For Data Augmentation)*","metadata":{}},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),  # Flip horizontally\n    A.Affine(\n        scale=(0.9, 1.1),       # Equivalent to scale_limit=0.1\n        translate_percent=(0.05, 0.05),  # Equivalent to shift_limit=0.05\n        rotate=(-10, 10),       # Equivalent to rotate_limit=10\n        p=0.5\n    ),\n    A.OneOf([                   # Either brightness or contrast adjustment\n        A.RandomBrightnessContrast(\n            brightness_limit=0.2, \n            contrast_limit=0.2, \n            p=0.5\n        ),\n        A.CLAHE(clip_limit=2.0, p=0.5)  # Contrast Limited Adaptive Histogram Equalization\n    ], p=0.3),  \n    A.GaussianBlur(blur_limit=3, p=0.2),  # Subtle blur to simulate noise\n])\n\nvalid_transform = A.Compose([\n    # nothing\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.052261Z","iopub.execute_input":"2025-03-26T14:01:44.052572Z","iopub.status.idle":"2025-03-26T14:01:44.062492Z","shell.execute_reply.started":"2025-03-26T14:01:44.052542Z","shell.execute_reply":"2025-03-26T14:01:44.061798Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" **we already have dataframe that has lavels from csv files**\n","metadata":{}},{"cell_type":"code","source":"print(labels_df.shape)  # (num_rows, num_columns)\nprint(labels_df.size)   # Total number of elements (num_rows * num_columns)\nprint(labels_df.head(101))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.063390Z","iopub.execute_input":"2025-03-26T14:01:44.063690Z","iopub.status.idle":"2025-03-26T14:01:44.082562Z","shell.execute_reply.started":"2025-03-26T14:01:44.063660Z","shell.execute_reply":"2025-03-26T14:01:44.081677Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**removing the bad ids from the dataframe and store it in *df***","metadata":{}},{"cell_type":"code","source":"# Remove excluded IDs\nexcluded_images = [ \"00109\",\"00123\", \"00709\"] \ndf = labels_df[~labels_df.BraTS21ID.isin(excluded_images)]\n\ndf.reset_index(drop=True, inplace=True)\n\nprint(df.head(101))  # Check first few rows\ndf.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.083519Z","iopub.execute_input":"2025-03-26T14:01:44.083727Z","iopub.status.idle":"2025-03-26T14:01:44.115204Z","shell.execute_reply.started":"2025-03-26T14:01:44.083703Z","shell.execute_reply":"2025-03-26T14:01:44.114501Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## testing the Data Retriver with the Transformations and with DF","metadata":{}},{"cell_type":"markdown","source":"**init**","metadata":{}},{"cell_type":"code","source":"# Define dataset using the filtered df\ndata = DataRetriever(\n    paths=df[\"BraTS21ID\"].values,  # Patient IDs\n    targets=df[\"MGMT_value\"],  # Labels (convert to int)\n    # mri_type=\"FLAIR\",  # Choose MRI type\n    # selection_method=\"best_central\",  # Selection strategy\n    transform=train_transform  # Apply transformations!\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.115986Z","iopub.execute_input":"2025-03-26T14:01:44.116246Z","iopub.status.idle":"2025-03-26T14:01:44.133113Z","shell.execute_reply.started":"2025-03-26T14:01:44.116227Z","shell.execute_reply":"2025-03-26T14:01:44.132226Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**fetching**","metadata":{}},{"cell_type":"code","source":"# Get a sample from the dataset\nsample = data[0]  # First patient\n\n# Extract X (image sequence) and y (label)\nX_transformed = sample[\"X\"]  # Transformed MRI slices (Tensor)\ny_label = sample[\"y\"]  # Label\n\nprint(f\"Transformed X Shape: {X_transformed.shape}\")  # Should be (14, H, W)\nprint(f\"Label: {y_label}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.133934Z","iopub.execute_input":"2025-03-26T14:01:44.134162Z","iopub.status.idle":"2025-03-26T14:01:44.204199Z","shell.execute_reply.started":"2025-03-26T14:01:44.134143Z","shell.execute_reply":"2025-03-26T14:01:44.203502Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Visuals**","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(2, 7, figsize=(24, 5))\nfor i, ax in enumerate(axes.flat):\n    if i < X_transformed.shape[0]:  # Ensure we don't exceed slice count\n        ax.imshow(X_transformed[i,0].numpy(), cmap=\"gray\")  # Access the single channel\n        ax.axis(\"off\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.204990Z","iopub.execute_input":"2025-03-26T14:01:44.205213Z","iopub.status.idle":"2025-03-26T14:01:44.629606Z","shell.execute_reply.started":"2025-03-26T14:01:44.205194Z","shell.execute_reply":"2025-03-26T14:01:44.628785Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"markdown","source":"## Loss and acc metters","metadata":{}},{"cell_type":"code","source":"class LossMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n\n    def update(self, val):\n        \"\"\"Updates the running loss average.\"\"\"\n        self.n += 1\n        self.avg = val / self.n + (self.n - 1) / self.n * self.avg\n\n    def reset(self):\n        \"\"\"Resets the loss meter.\"\"\"\n        self.avg = 0\n        self.n = 0\n\n\nclass AccMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n\n    def update(self, y_true, y_pred):\n        \"\"\"Updates the accuracy based on true vs predicted labels.\"\"\"\n        y_true = y_true.cpu().numpy().astype(int)\n        y_pred = (y_pred.cpu().numpy() >= 0).astype(int)  # Threshold at 0 for binary classification\n\n        last_n = self.n\n        self.n += len(y_true)  # Total samples processed\n        correct = np.sum(y_true == y_pred)  # Correct predictions\n\n        # Incremental accuracy update formula\n        self.avg = correct / self.n + last_n / self.n * self.avg\n\n    def reset(self):\n        \"\"\"Resets the accuracy meter.\"\"\"\n        self.avg = 0\n        self.n = 0\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.630459Z","iopub.execute_input":"2025-03-26T14:01:44.630777Z","iopub.status.idle":"2025-03-26T14:01:44.636752Z","shell.execute_reply.started":"2025-03-26T14:01:44.630745Z","shell.execute_reply":"2025-03-26T14:01:44.635961Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Trainer class","metadata":{}},{"cell_type":"code","source":"class Trainer: # class that manges train and val of dl model\n    # this func run when a trainer object is created\n    def __init__(self, model, device, optimizer, criterion, loss_meter, score_meter): \n        \n        self.model = model.to(device) # model move it to gpu if available\n        self.device = device # cpu or gpu\n        self.optimizer = optimizer # wheights updater\n        self.criterion = criterion # loss func\n        self.loss_meter = loss_meter  # tracks avg loss over batches\n        self.score_meter = score_meter  # tracks acc or other metric\n\n        #tracks perfomance for plotting\n        self.hist = {\n            'train_loss': [], 'train_score': [],\n            'val_loss': [], 'val_score': []\n        }\n        # best score tracking + early stoping \n        self.best_valid_score = -np.inf # best val acc start very low\n        self.best_valid_loss = np.inf # best val loss start very high\n        self.n_patience = 0 # counts epochs without improvement\n        self.messages = {\n            \"epoch\": \"[Epoch {}: {}] loss: {:.5f}, score: {:.5f}, time: {} s\",\n            \"checkpoint\": \"Score improved from {:.5f} to {:.5f}. Saving model...\",\n            \"patience\": \"\\nValidation score didn't improve for {} epochs.\"\n        }\n    \n    def fit(self, epochs, train_loader, valid_loader, save_path, patience=5):  \n        # training loop\n        for n_epoch in range(1, epochs + 1):\n            self.info_message(\"EPOCH: {}\", n_epoch)\n            \n            # train & vall Calls\n            train_loss, train_score, train_time = self.train_epoch(train_loader)\n            valid_loss, valid_score, valid_time = self.valid_epoch(valid_loader)\n            \n            # storting results for later plotting\n            self.hist['train_loss'].append(train_loss)\n            self.hist['train_score'].append(train_score)\n            self.hist['val_loss'].append(valid_loss)\n            self.hist['val_score'].append(valid_score)\n\n            self.info_message(self.messages[\"epoch\"], \"Train\", n_epoch, train_loss, train_score, train_time)\n            self.info_message(self.messages[\"epoch\"], \"Valid\", n_epoch, valid_loss, valid_score, valid_time)\n            \n            # 'earlystoping & model saving\n            # Save best model \n            if valid_score > self.best_valid_score:\n                self.info_message(self.messages[\"checkpoint\"], self.best_valid_score, valid_score)\n                self.best_valid_score = valid_score\n                self.best_valid_loss = valid_loss\n                self.save_model(n_epoch, save_path)\n                self.n_patience = 0  # Reset patience\n            else:\n                self.n_patience += 1 #if no imporovement patience increases\n            \n            # Early stopping check after 5 epochs\n            if self.n_patience >= patience:\n                self.info_message(self.messages[\"patience\"], patience)\n                break\n                \n        return self.best_valid_loss, self.best_valid_score\n        \n    # Goes through all training batches once and updates the model.     \n    def train_epoch(self, train_loader):\n        self.model.train() #training mode\n        t = time.time()\n        train_loss = self.loss_meter()\n        train_score = self.score_meter()\n        \n        # Loads a batch of training data.\n        for step, batch in enumerate(train_loader, 1):\n            X = batch[\"X\"].to(self.device)  # Shape: (B, T, C, H, W)\n            targets = batch[\"y\"].to(self.device)\n            \n            self.optimizer.zero_grad() #rest the gradients\n            outputs = self.model(X).squeeze(1) #pass the data throuhg the model \n            \n            loss = self.criterion(outputs, targets) # compute loss\n            loss.backward() # compute gradients\n\n            # update the avg loss and acc\n            train_loss.update(loss.detach().item()) \n            train_score.update(targets, outputs.detach())\n            \n            # update weights\n            self.optimizer.step() \n            \n            _loss, _score = train_loss.avg, train_score.avg\n            message = 'Train Step {}/{}, train_loss: {:.5f}, train_score: {:.5f}'\n            self.info_message(message, step, len(train_loader), _loss, _score, end=\"\\r\")\n        \n        return train_loss.avg, train_score.avg, int(time.time() - t)\n        \n    #validate the model (run without updating the model only evaluation)\n    def valid_epoch(self, valid_loader):\n        self.model.eval() #evaluationn mode\n        t = time.time()\n        valid_loss = self.loss_meter()\n        valid_score = self.score_meter()\n        \n        #turns off the gradient (faster)\n        with torch.no_grad():\n            for step, batch in enumerate(valid_loader, 1):\n                X = batch[\"X\"].to(self.device)\n                targets = batch[\"y\"].to(self.device)\n\n                outputs = self.model(X).squeeze(1)\n                loss = self.criterion(outputs, targets)\n\n                valid_loss.update(loss.detach().item())\n                valid_score.update(targets, outputs)\n                \n                _loss, _score = valid_loss.avg, valid_score.avg\n                message = 'Valid Step {}/{}, valid_loss: {:.5f}, valid_score: {:.5f}'\n                self.info_message(message, step, len(valid_loader), _loss, _score, end=\"\\r\")\n        \n        return valid_loss.avg, valid_score.avg, int(time.time() - t)\n    #ploting\n    def plot_loss(self):\n        plt.figure(figsize=(10, 5))\n        plt.plot(self.hist['train_loss'], label=\"Train Loss\")\n        plt.plot(self.hist['val_loss'], label=\"Validation Loss\")\n        plt.xlabel(\"Epochs\")\n        plt.ylabel(\"Loss\")\n        plt.legend()\n        plt.title(\"Loss Curve\")\n        plt.show()\n    \n    def plot_score(self):\n        plt.figure(figsize=(10, 5))\n        plt.plot(self.hist['train_score'], label=\"Train Accuracy\")\n        plt.plot(self.hist['val_score'], label=\"Validation Accuracy\")\n        plt.xlabel(\"Epochs\")\n        plt.ylabel(\"Score\")\n        plt.legend()\n        plt.title(\"Accuracy Curve\")\n        plt.show()\n    #saving func\n    def save_model(self, n_epoch, save_path):\n        torch.save(\n            {\n                \"model_state_dict\": self.model.state_dict(),\n                \"optimizer_state_dict\": self.optimizer.state_dict(),\n                \"best_valid_score\": self.best_valid_score,\n                \"n_epoch\": n_epoch,\n            },\n            save_path,\n        )\n    #no self using\n    @staticmethod\n    def info_message(message, *args, end=\"\\n\"):\n        print(message.format(*args), end=end)\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T14:01:44.637672Z","iopub.execute_input":"2025-03-26T14:01:44.637971Z","iopub.status.idle":"2025-03-26T14:01:44.656804Z","shell.execute_reply.started":"2025-03-26T14:01:44.637941Z","shell.execute_reply":"2025-03-26T14:01:44.656008Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cross-Validation Training","metadata":{}},{"cell_type":"code","source":"# Define Stratified K-Fold for Cross-Validation\nskf = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True)\nt = df['MGMT_value']  # Assuming 'MGMT_value' is the target column\n\nnum_workers = min(4, os.cpu_count() or 1)  # Use at most 4 workers\n\nstart_time = time.time()\n\n# To track performance across folds\nlosses = []\nscores = []\ntest_dfs = []\n\nfor fold, (train_idx, val_idx) in enumerate(skf.split(np.zeros(len(t)), t), 1):\n    print(f\"{'-'*30}\\nFold {fold}\")\n\n    # Split the data into training and validation for the current fold\n    train_df = df.loc[train_idx]\n    val_df = df.loc[val_idx]\n    test_dfs.append(val_df)\n\n    # Data retrieval for training and validation\n    train_retriever = DataRetriever(\n        train_df[\"BraTS21ID\"].values, \n        train_df[\"MGMT_value\"].values,\n        mri_type= \"T1wCE\", #  (FLAIR T1w  T1wCE T2w) choose one of these\n        selection_method=\"best_central\", \n        transform = train_transform  # Assuming data augmentations for training\n    )\n    val_retriever = DataRetriever(\n        val_df[\"BraTS21ID\"].values, \n        val_df[\"MGMT_value\"].values,\n        selection_method=\"best_central\", \n        mri_type= \"T1wCE\"\n    )\n\n    # Data loaders\n    train_loader = torch.utils.data.DataLoader(\n        train_retriever, \n        batch_size=6, \n        shuffle=True, \n        num_workers=num_workers\n    )\n    valid_loader = torch.utils.data.DataLoader(\n        val_retriever, \n        batch_size=6, \n        shuffle=False, \n        num_workers=num_workers\n    )\n\n    # Model, optimizer, and loss function\n    model = Model().to(device)\n    optimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\n    criterion = F.binary_cross_entropy_with_logits\n\n    # Trainer instance\n    trainer = Trainer(\n        model, \n        device, \n        optimizer, \n        criterion, \n        LossMeter, \n        AccMeter\n    )\n\n    # Fit model and track performance\n    loss, score = trainer.fit(\n        CFG.n_epochs,\n        train_loader,\n        valid_loader,\n        f\"best-model-fold{fold}.pth\",  # Save the best model for each fold\n        patience= 5 # Early stopping patience\n    )\n\n    losses.append(loss)\n    scores.append(score)\n\n    # Plotting performance for each fold\n    trainer.plot_loss()\n    trainer.plot_score()\n\n# Time tracking\nelapsed_time = time.time() - start_time\nprint(f\"\\nTraining complete in {elapsed_time // 60:.0f}m {elapsed_time % 60:.0f}s\")\nprint(f\"Average Loss: {np.mean(losses):.4f}\")\nprint(f\"Average Score: {np.mean(scores):.4f}\")\n","metadata":{"execution":{"iopub.status.busy":"2025-03-26T14:01:44.657686Z","iopub.execute_input":"2025-03-26T14:01:44.657938Z","iopub.status.idle":"2025-03-26T14:55:19.023273Z","shell.execute_reply.started":"2025-03-26T14:01:44.657907Z","shell.execute_reply":"2025-03-26T14:55:19.022203Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(test_dfs))  # Should be equal to CFG.n_fold\nprint(test_dfs[2]) ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T15:05:17.646771Z","iopub.execute_input":"2025-03-26T15:05:17.647072Z","iopub.status.idle":"2025-03-26T15:05:17.653974Z","shell.execute_reply.started":"2025-03-26T15:05:17.647050Z","shell.execute_reply":"2025-03-26T15:05:17.653005Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## metrics\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import (accuracy_score, precision_score, recall_score, \n                           f1_score, roc_auc_score, roc_curve, confusion_matrix)\n\n# Evaluation Cell with Confusion Matrix\nall_fold_metrics = []\nplt.style.use('ggplot')  # Better looking plots\n\nfor fold in range(1, CFG.n_fold + 1):\n    print(f\"\\n{'='*50}\")\n    print(f\"Evaluating Fold {fold}\")\n    print(f\"{'='*50}\")\n    \n    try:\n        # 1. Load the saved checkpoint\n        model_path = f\"best-model-fold{fold}.pth\"\n        checkpoint = torch.load(model_path, map_location=device)\n        \n        # Initialize and load model\n        model = Model().to(device)\n        model.load_state_dict(checkpoint['model_state_dict'])\n        model.eval()\n        \n        # 2. Get validation data for this fold\n        val_df = test_dfs[fold-1]\n        \n        # 3. Create validation dataloader\n        val_retriever = DataRetriever(\n            val_df[\"BraTS21ID\"].values, \n            val_df[\"MGMT_value\"].values,\n            selection_method=\"best_central\", \n            mri_type=\"T1wCE\"\n        )\n        \n        valid_loader = torch.utils.data.DataLoader(\n            val_retriever, \n            batch_size=6, \n            shuffle=False, \n            num_workers=num_workers\n        )\n        \n        # 4. Get predictions\n        true_labels = []\n        pred_probs = []\n        \n        with torch.no_grad():\n            for batch in valid_loader:\n                images = batch[\"X\"].to(device)\n                labels = batch[\"y\"].to(device)\n                \n                outputs = model(images).squeeze(1)\n                probabilities = torch.sigmoid(outputs)\n                \n                true_labels.extend(labels.cpu().numpy())\n                pred_probs.extend(probabilities.cpu().numpy())\n        \n        y_true = np.array(true_labels)\n        y_scores = np.array(pred_probs)\n        y_pred = (y_scores >= 0.5).astype(int)\n        \n        # 5. Calculate metrics only if we have data\n        if len(y_true) > 0 and len(y_pred) > 0:\n            # Standard metrics\n            fold_metrics = {\n                'fold': fold,\n                'accuracy': accuracy_score(y_true, y_pred),\n                'precision': precision_score(y_true, y_pred, zero_division=1),\n                'recall': recall_score(y_true, y_pred, zero_division=1),\n                'f1': f1_score(y_true, y_pred, zero_division=1),\n                'roc_auc': roc_auc_score(y_true, y_scores),\n                'support': len(y_true)\n            }\n            all_fold_metrics.append(fold_metrics)\n            \n            # Print metrics\n            print(f\"\\nFold {fold} Validation Metrics:\")\n            print(f\"- Accuracy:  {fold_metrics['accuracy']:.4f}\")\n            print(f\"- Precision: {fold_metrics['precision']:.4f}\")\n            print(f\"- Recall:    {fold_metrics['recall']:.4f}\")\n            print(f\"- F1 Score:  {fold_metrics['f1']:.4f}\")\n            print(f\"- ROC AUC:   {fold_metrics['roc_auc']:.4f}\")\n            print(f\"- Samples:   {fold_metrics['support']}\")\n            \n            # ROC Curve\n            fpr, tpr, _ = roc_curve(y_true, y_scores)\n            plt.figure(figsize=(6, 5))\n            plt.plot(fpr, tpr, color='darkorange', lw=2, \n                    label=f'ROC (AUC = {fold_metrics[\"roc_auc\"]:.2f})')\n            plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n            plt.xlabel('False Positive Rate')\n            plt.ylabel('True Positive Rate')\n            plt.title(f'ROC Curve - Fold {fold}')\n            plt.legend(loc=\"lower right\")\n            plt.grid(True)\n            plt.tight_layout()\n            plt.savefig(f\"ROC_curve_fold{fold}.png\", dpi=300, bbox_inches='tight')\n            plt.show()\n            \n            # Confusion Matrix\n            cm = confusion_matrix(y_true, y_pred)\n            plt.figure(figsize=(4, 4))\n            sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", cbar=False,\n                      annot_kws={\"size\": 14}, square=True)\n            \n            plt.xlabel('Predicted', fontsize=12)\n            plt.ylabel('Actual', fontsize=12)\n            plt.title(f'Confusion Matrix - Fold {fold}', fontsize=14)\n            \n            # Handle binary or multi-class labels\n            if len(np.unique(y_true)) == 2:  # Binary case\n                class_names = [\"Negative\", \"Positive\"]\n                tick_pos = [0.5, 1.5]\n            else:  # Multi-class case\n                class_names = [f\"Class {i}\" for i in range(len(np.unique(y_true)))]\n                tick_pos = np.arange(len(class_names)) + 0.5\n            \n            plt.xticks(tick_pos, class_names, rotation=45, ha='right')\n            plt.yticks(tick_pos, class_names, va='center')\n            \n            plt.tight_layout()\n            plt.savefig(f\"Confusion_matrix_fold{fold}.png\", dpi=300, bbox_inches='tight')\n            plt.show()\n            \n        else:\n            print(f\"Warning: No valid predictions for Fold {fold}\")\n            \n    except Exception as e:\n        print(f\"Error evaluating Fold {fold}: {str(e)}\")\n        continue\n\n# Print comprehensive summary\nif len(all_fold_metrics) > 0:\n    print(\"\\n\\nSUMMARY ACROSS ALL FOLDS:\")\n    print(\"{:<6} {:<9} {:<9} {:<9} {:<9} {:<9} {:<9}\".format(\n        'Fold', 'Accuracy', 'Precision', 'Recall', 'F1', 'ROC AUC', 'Samples'))\n\n    for metrics in all_fold_metrics:\n        print(\"{:<6} {:<9.4f} {:<9.4f} {:<9.4f} {:<9.4f} {:<9.4f} {:<9}\".format(\n            metrics['fold'],\n            metrics['accuracy'],\n            metrics['precision'],\n            metrics['recall'],\n            metrics['f1'],\n            metrics['roc_auc'],\n            metrics['support']))\n\n    # Calculate averages\n    avg_metrics = {\n        'accuracy': np.mean([m['accuracy'] for m in all_fold_metrics]),\n        'precision': np.mean([m['precision'] for m in all_fold_metrics]),\n        'recall': np.mean([m['recall'] for m in all_fold_metrics]),\n        'f1': np.mean([m['f1'] for m in all_fold_metrics]),\n        'roc_auc': np.mean([m['roc_auc'] for m in all_fold_metrics])\n    }\n\n    print(\"\\nAverage Metrics Across All Folds:\")\n    print(f\"- Accuracy:  {avg_metrics['accuracy']:.4f}\")\n    print(f\"- Precision: {avg_metrics['precision']:.4f}\")\n    print(f\"- Recall:    {avg_metrics['recall']:.4f}\")\n    print(f\"- F1 Score:  {avg_metrics['f1']:.4f}\")\n    print(f\"- ROC AUC:   {avg_metrics['roc_auc']:.4f}\")\nelse:\n    print(\"No valid folds were evaluated.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-26T15:21:11.991596Z","iopub.execute_input":"2025-03-26T15:21:11.991952Z","iopub.status.idle":"2025-03-26T15:22:33.878310Z","shell.execute_reply.started":"2025-03-26T15:21:11.991915Z","shell.execute_reply":"2025-03-26T15:22:33.877106Z"}},"outputs":[],"execution_count":null}]}