{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q pydicom nibabel scikit-image monai tqdm scikit-learn\n\nimport pandas as pd\nimport numpy as np\nimport nibabel as nib\nimport os\nimport ast\nimport pydicom\nfrom scipy.ndimage import zoom\nfrom skimage.exposure import equalize_adapthist\nfrom sklearn.model_selection import train_test_split\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset\nfrom monai.data import DataLoader\nfrom monai.networks.nets import UNet\nfrom monai.losses import DiceCELoss\nfrom monai.transforms import Compose, RandCropByPosNegLabeld, SpatialPadd\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\nprint(\"Setup complete.\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-03T14:26:49.922825Z","iopub.execute_input":"2025-10-03T14:26:49.923181Z","iopub.status.idle":"2025-10-03T14:29:20.858842Z","shell.execute_reply.started":"2025-10-03T14:26:49.923154Z","shell.execute_reply":"2025-10-03T14:29:20.857416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- PATHS ---\nBASE_INPUT_DIR = '/kaggle/input/rsna-intracranial-aneurysm-detection/'\nORIGINAL_DICOM_DIR = os.path.join(BASE_INPUT_DIR, 'series')\nLOCALIZER_CSV_PATH = os.path.join(BASE_INPUT_DIR, 'train_localizers.csv')\nVESSEL_SEG_DIR = os.path.join(BASE_INPUT_DIR, 'segmentations') # <-- NEW PATH ADDED\n\nPROCESSED_IMAGES_DIR = '/kaggle/working/processed_images/'\nMASKS_DIR = '/kaggle/working/processed_masks/'\n\n# --- DATA & PREPROCESSING PARAMETERS ---\nTARGET_SPACING = (1.0, 1.0, 1.0)\nSUBSET_SIZE = 200\n\n# --- TRAINING HYPERPARAMETERS ---\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nLEARNING_RATE = 1e-5\nBATCH_SIZE = 4\nNUM_EPOCHS = 25\nVAL_SPLIT = 0.2\nPATCH_SIZE = (96, 96, 96)\n\n# --- MAPPING ---\nLOCATION_TO_CHANNEL = { 'Other Posterior Circulation': 0, 'Basilar Tip': 1, 'Right Posterior Communicating Artery': 2, 'Left Posterior Communicating Artery': 3, 'Right Infraclinoid Internal Carotid Artery': 4, 'Left Infraclinoid Internal Carotid Artery': 5, 'Right Supraclinoid Internal Carotid Artery': 6, 'Left Supraclinoid Internal Carotid Artery': 7, 'Right Middle Cerebral Artery': 8, 'Left Middle Cerebral Artery': 9, 'Right Anterior Cerebral Artery': 10, 'Left Anterior Cerebral Artery': 11, 'Anterior Communicating Artery': 12 }\n\nprint(f\"Configuration complete. Using device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T14:29:20.860895Z","iopub.execute_input":"2025-10-03T14:29:20.86164Z","iopub.status.idle":"2025-10-03T14:29:20.871085Z","shell.execute_reply.started":"2025-10-03T14:29:20.861581Z","shell.execute_reply":"2025-10-03T14:29:20.869794Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def resample_volume(volume, o_spacing, t_spacing):\n    zoom_factor = [o / t for o, t in zip(o_spacing, t_spacing)]\n    return zoom(volume, zoom_factor, order=1)\n\ndef normalize_intensity(volume, clip_percentiles=(0.5, 99.5)):\n    if np.all(volume == 0): return volume\n    lower, upper = np.percentile(volume, clip_percentiles)\n    clipped = np.clip(volume, lower, upper)\n    min_val, max_val = np.min(clipped), np.max(clipped)\n    if max_val - min_val == 0: return np.zeros_like(clipped)\n    return (clipped - min_val) / (max_val - min_val)\n\ndef enhance_contrast_clahe(volume):\n    enhanced_volume = np.zeros_like(volume)\n    for i in range(volume.shape[2]):\n        enhanced_volume[:, :, i] = equalize_adapthist(volume[:, :, i], clip_limit=0.01)\n    return enhanced_volume\n\ndef reconstruct_from_dicom(series_path):\n    if not os.path.isdir(series_path): return None, None\n    dicom_files = [pydicom.dcmread(os.path.join(series_path, f)) for f in os.listdir(series_path) if f.endswith('.dcm')]\n    if not dicom_files: return None, None\n    def get_slice_pos(dcm):\n        if 'ImagePositionPatient' in dcm: return float(dcm.ImagePositionPatient[2])\n        elif 'SliceLocation' in dcm: return float(dcm.SliceLocation)\n        else: return float(dcm.get('InstanceNumber', 0))\n    dicom_files.sort(key=get_slice_pos)\n    p_space = dicom_files[0].get('PixelSpacing', [1.0, 1.0])\n    s_thick = abs(get_slice_pos(dicom_files[1]) - get_slice_pos(dicom_files[0])) if len(dicom_files) > 1 else 1.0\n    if s_thick < 1e-3: return None, None\n    o_space = [float(p_space[0]), float(p_space[1]), float(s_thick)]\n    volume = np.stack([s.pixel_array for s in dicom_files], axis=-1)\n    return volume, o_space\n\ndef parse_coords(s):\n    try: return ast.literal_eval(s)\n    except: return [np.nan, np.nan]\n\ndef get_original_dicom_info(series_uid):\n    s_path = os.path.join(ORIGINAL_DICOM_DIR, series_uid)\n    if not os.path.isdir(s_path): return None, None\n    dcms = [pydicom.dcmread(os.path.join(s_path, f), stop_before_pixels=True) for f in os.listdir(s_path) if f.endswith('.dcm')]\n    if not dcms: return None, None\n    def get_slice_pos(dcm):\n        if 'ImagePositionPatient' in dcm: return float(dcm.ImagePositionPatient[2])\n        elif 'SliceLocation' in dcm: return float(dcm.SliceLocation)\n        else: return float(dcm.get('InstanceNumber', 0))\n    dcms.sort(key=get_slice_pos)\n    sops = [d.SOPInstanceUID for d in dcms]\n    return sops\n\ndef draw_sphere(mask, center, radius):\n    center = [int(round(c)) for c in center]\n    x_c, y_c, z_c = center\n    for i in range(max(0, x_c - radius), min(mask.shape[0], x_c + radius + 1)):\n        for j in range(max(0, y_c - radius), min(mask.shape[1], y_c + radius + 1)):\n            for k in range(max(0, z_c - radius), min(mask.shape[2], z_c + radius + 1)):\n                if (i - x_c)**2 + (j - y_c)**2 + (k - z_c)**2 <= radius**2:\n                    mask[i, j, k] = 1\n    return mask\n\nprint(\"All helper functions defined.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T14:29:20.872267Z","iopub.execute_input":"2025-10-03T14:29:20.872675Z","iopub.status.idle":"2025-10-03T14:29:20.916659Z","shell.execute_reply.started":"2025-10-03T14:29:20.872644Z","shell.execute_reply":"2025-10-03T14:29:20.915474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- FINAL, CORRECTED Cell for Data Preparation ---\n\ndef process_dicom_series(series_uid):\n    \"\"\"\n    This single function handles the full pipeline for one series:\n    1. Reconstructs from DICOM.\n    2. Preprocesses the image (resample, normalize, CLAHE).\n    3. Creates the corresponding multi-channel mask.\n    Returns the processed image and mask, or None if an error occurs.\n    \"\"\"\n    series_path = os.path.join(ORIGINAL_DICOM_DIR, series_uid)\n    if not os.path.isdir(series_path): return None, None\n\n    # --- Reconstruction and Metadata Gathering (robust version) ---\n    dicom_files = [pydicom.dcmread(os.path.join(series_path, f)) for f in os.listdir(series_path) if f.endswith('.dcm')]\n    if not dicom_files: return None, None\n    \n    def get_slice_pos(dcm):\n        if 'ImagePositionPatient' in dcm: return float(dcm.ImagePositionPatient[2])\n        elif 'SliceLocation' in dcm: return float(dcm.SliceLocation)\n        else: return float(dcm.get('InstanceNumber', 0))\n    dicom_files.sort(key=get_slice_pos)\n\n    p_space = dicom_files[0].get('PixelSpacing', [1.0, 1.0])\n    s_thick = abs(get_slice_pos(dicom_files[1]) - get_slice_pos(dicom_files[0])) if len(dicom_files) > 1 else 1.0\n    if s_thick < 1e-3: return None, None\n    \n    original_spacing = [float(p_space[0]), float(p_space[1]), float(s_thick)]\n    volume = np.stack([s.pixel_array for s in dicom_files], axis=-1)\n    if volume.ndim > 3: volume = np.squeeze(volume)\n    if volume.ndim != 3: return None, None\n    \n    # --- Image Preprocessing ---\n    resampled = resample_volume(volume, original_spacing, TARGET_SPACING)\n    normalized = normalize_intensity(resampled)\n    enhanced_image = enhance_contrast_clahe(normalized)\n\n    # --- Mask Generation ---\n    sops = [d.SOPInstanceUID for d in dicom_files]\n    zoom_factor = [o / t for o, t in zip(original_spacing, TARGET_SPACING)]\n    \n    mask = np.zeros((13,) + enhanced_image.shape, dtype=np.uint8)\n    aneurysms_in_series = localizer_df[localizer_df['SeriesInstanceUID'] == series_uid]\n\n    for _, row in aneurysms_in_series.iterrows():\n        try:\n            z = sops.index(row['SOPInstanceUID'])\n            center = [row['x'] * zoom_factor[0], row['y'] * zoom_factor[1], z * zoom_factor[2]]\n            ch_idx = LOCATION_TO_CHANNEL.get(row['location'])\n            if ch_idx is not None:\n                mask[ch_idx] = draw_sphere(mask[ch_idx], center, radius=3)\n        except (ValueError, AttributeError):\n            pass\n\n    if np.sum(mask) == 0:\n        return None, None # Skip series if no aneurysm was successfully drawn\n        \n    final_mask = np.transpose(mask, (0, 3, 1, 2))\n    \n    return enhanced_image, final_mask\n\n# --- MAIN EXECUTION ---\n!rm -rf /kaggle/working/*\nos.makedirs(PROCESSED_IMAGES_DIR, exist_ok=True)\nos.makedirs(MASKS_DIR, exist_ok=True)\n\nlocalizer_df = pd.read_csv(LOCALIZER_CSV_PATH)\nlocalizer_df['coordinates'] = localizer_df['coordinates'].apply(parse_coords)\ncoords_df = pd.DataFrame(localizer_df['coordinates'].tolist(), columns=['x', 'y'], index=localizer_df.index)\nlocalizer_df = pd.concat([localizer_df, coords_df], axis=1)\nlocalizer_df.dropna(subset=['x', 'y'], inplace=True)\nlocalizer_df['x'] = localizer_df['x'].astype(int); localizer_df['y'] = localizer_df['y'].astype(int)\n\npositive_series_uids = localizer_df['SeriesInstanceUID'].unique().tolist()\nseries_to_process = positive_series_uids[:SUBSET_SIZE]\n\nprint(f\"--- Starting Data Preparation for {len(series_to_process)} series ---\")\nfor series_uid in tqdm(series_to_process, desc=\"Preparing Data\"):\n    image, mask = process_dicom_series(series_uid)\n    \n    if image is not None and mask is not None:\n        affine = np.eye(4)\n        nib.save(nib.Nifti1Image(image.astype(np.float32), affine), os.path.join(PROCESSED_IMAGES_DIR, f\"{series_uid}.nii.gz\"))\n        nib.save(nib.Nifti1Image(mask.astype(np.uint8), affine), os.path.join(MASKS_DIR, f\"{series_uid}.nii.gz\"))\n\nprint(\"\\n--- Data Preparation Complete! ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T14:29:20.917816Z","iopub.execute_input":"2025-10-03T14:29:20.918209Z","iopub.status.idle":"2025-10-03T15:14:21.486755Z","shell.execute_reply.started":"2025-10-03T14:29:20.918176Z","shell.execute_reply":"2025-10-03T15:14:21.485545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## 👁️ VISUALIZATION: Before vs. After (Same Aneurysm Slice)\n\nprint(\"--- Finding and displaying the same aneurysm slice before and after processing ---\")\n\n# --- 1. Find the location of an aneurysm in the raw data ---\nlocalizer_df = pd.read_csv(LOCALIZER_CSV_PATH)\nif not localizer_df.empty:\n    # Pick the first aneurysm from the list\n    first_aneurysm_info = localizer_df.iloc[0]\n    sample_series_uid = first_aneurysm_info['SeriesInstanceUID']\n    aneurysm_sop_uid = first_aneurysm_info['SOPInstanceUID']\n\n    # Get original DICOM info to find the raw slice index\n    # We need to call the full reconstruct function to get original_spacing\n    raw_volume_for_spacing, original_spacing = reconstruct_from_dicom(os.path.join(ORIGINAL_DICOM_DIR, sample_series_uid))\n    sops = get_original_dicom_info(sample_series_uid)\n    \n    aneurysm_slice_idx_raw = -1\n    if sops:\n        try:\n            aneurysm_slice_idx_raw = sops.index(aneurysm_sop_uid)\n        except ValueError:\n            print(f\"Warning: Aneurysm slice {aneurysm_sop_uid} not found in series {sample_series_uid}.\")\n\n    if aneurysm_slice_idx_raw != -1 and raw_volume_for_spacing is not None:\n        raw_volume = raw_volume_for_spacing\n        if raw_volume.ndim > 3: raw_volume = np.squeeze(raw_volume)\n        \n        # --- 2. Display the RAW DICOM slice ---\n        plt.figure(figsize=(12, 6))\n        plt.subplot(1, 2, 1)\n        plt.imshow(raw_volume[:, :, aneurysm_slice_idx_raw], cmap='gray')\n        plt.title(f\"Raw DICOM\\nSlice index: {aneurysm_slice_idx_raw}\")\n        plt.axis('off')\n\n        # --- 3. Display the PROCESSED NIfTI slice ---\n        processed_path = os.path.join(PROCESSED_IMAGES_DIR, f\"{sample_series_uid}.nii.gz\")\n        if os.path.exists(processed_path):\n            # Calculate the corresponding slice index in the resampled volume\n            zoom_factor_z = original_spacing[2] / TARGET_SPACING[2]\n            aneurysm_slice_idx_processed = int(round(aneurysm_slice_idx_raw * zoom_factor_z))\n\n            processed_data = nib.load(processed_path).get_fdata()\n            \n            plt.subplot(1, 2, 2)\n            # Ensure the processed slice index is within bounds\n            aneurysm_slice_idx_processed = min(aneurysm_slice_idx_processed, processed_data.shape[2] - 1)\n            plt.imshow(processed_data[:, :, aneurysm_slice_idx_processed], cmap='gray')\n            plt.title(f\"Processed Scan\\nSlice index: {aneurysm_slice_idx_processed}\")\n            plt.axis('off')\n            plt.show()\n        else:\n            print(\"Processed file not found. Please run the Data Preparation cell first.\")\nelse:\n    print(\"Localizer CSV is empty.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T15:14:21.489688Z","iopub.execute_input":"2025-10-03T15:14:21.490519Z","iopub.status.idle":"2025-10-03T15:14:28.961415Z","shell.execute_reply.started":"2025-10-03T15:14:21.490484Z","shell.execute_reply":"2025-10-03T15:14:28.96029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport nibabel as nib\nimport os\nimport numpy as np\n\nprint(\"--- Verifying a sample image and mask ---\")\n\n# 1. Get a list of the masks you just created\nmask_files = [f for f in os.listdir(MASKS_DIR) if f.endswith('.nii.gz')]\n\nif not mask_files:\n    print(\"❌ Verification failed: No mask files were found in the output directory.\")\nelse:\n    # 2. Pick one sample to check\n    sample_filename = mask_files[0]\n    print(f\"Checking sample: {sample_filename}\")\n\n    image_path = os.path.join(PROCESSED_IMAGES_DIR, sample_filename)\n    mask_path = os.path.join(MASKS_DIR, sample_filename)\n\n    # 3. Load the image and the 13-channel mask\n    image_nii = nib.load(image_path)\n    mask_nii = nib.load(mask_path)\n    image_data = image_nii.get_fdata() # Expected shape: (H, W, D)\n    mask_data = mask_nii.get_fdata()   # Expected shape: (C, D, H, W)\n\n    print(f\"Loaded image shape: {image_data.shape}\")\n    print(f\"Loaded mask shape: {mask_data.shape}\")\n\n    # 4. Find a slice that contains the aneurysm\n    # To do this, we \"flatten\" the 13 channels by taking the max value at each voxel\n    combined_mask_3d = np.max(mask_data, axis=0) # Shape: (D, H, W)\n    \n    # Find the 3D coordinates of any voxel that is not zero\n    aneurysm_coords = np.argwhere(combined_mask_3d > 0)\n    \n    if len(aneurysm_coords) > 0:\n        # Get the Z-slice index (depth) from the first found coordinate\n        slice_idx_d = aneurysm_coords[0][0]\n        print(f\"✅ Aneurysm found on slice (depth index): {slice_idx_d}\")\n\n        # 5. Plot the corresponding slices side-by-side\n        plt.figure(figsize=(12, 6))\n\n        # Plot the image slice. Image is (H, W, D), so we get the slice with [:, :, slice_idx]\n        plt.subplot(1, 2, 1)\n        plt.imshow(image_data[:, :, slice_idx_d], cmap='gray', origin='lower')\n        plt.title(f'Processed Image (Slice {slice_idx_d})')\n        plt.axis('off')\n\n        # Plot the mask slice. Combined mask is (D, H, W), so we get the slice with [slice_idx, :, :]\n        plt.subplot(1, 2, 2)\n        plt.imshow(combined_mask_3d[slice_idx_d, :, :], cmap='gray', origin='lower')\n        plt.title(f'Generated Mask (Slice {slice_idx_d})')\n        plt.axis('off')\n\n        plt.show()\n    else:\n        # This shouldn't happen with the new data prep script, but it's a good safety check\n        print(\"⚠️ Verification warning: A mask file was found, but it appears to be all black.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T15:14:28.962829Z","iopub.execute_input":"2025-10-03T15:14:28.963218Z","iopub.status.idle":"2025-10-03T15:14:41.168153Z","shell.execute_reply.started":"2025-10-03T15:14:28.963188Z","shell.execute_reply":"2025-10-03T15:14:41.166533Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom monai.transforms import Compose, SpatialPadd, RandCropByPosNegLabeld\nfrom torch.utils.data import Dataset\nfrom monai.data import DataLoader\n\n# --- Transform Pipelines (No changes needed) ---\ntrain_transforms = Compose([\n    SpatialPadd(keys=['image', 'mask'], spatial_size=PATCH_SIZE, method='end'),\n    RandCropByPosNegLabeld(keys=['image', 'mask'], label_key='mask', spatial_size=PATCH_SIZE, pos=1, neg=1, num_samples=2, image_key='image', image_threshold=0)\n])\nval_transforms = Compose([\n    SpatialPadd(keys=['image', 'mask'], spatial_size=PATCH_SIZE, method='end'),\n    RandCropByPosNegLabeld(keys=['image', 'mask'], label_key='mask', spatial_size=PATCH_SIZE, pos=1, neg=0, num_samples=1)\n])\n\n# --- UPGRADED Dataset Class for 2-Channel Input ---\nclass AneurysmDataset(Dataset):\n    def __init__(self, image_dir, mask_dir, vessel_dir, filenames, transform=None):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.vessel_dir = vessel_dir # <-- Store path to vessel masks\n        self.filenames = filenames\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.filenames)\n\n    def __getitem__(self, idx):\n        fname = self.filenames[idx]\n        try:\n            # 1. Load preprocessed image and aneurysm mask\n            img_path = os.path.join(self.image_dir, fname)\n            mask_path = os.path.join(self.mask_dir, fname)\n            img = nib.load(img_path).get_fdata().astype(np.float32)\n            aneurysm_mask = nib.load(mask_path).get_fdata().astype(np.float32)\n\n            # 2. Load corresponding vessel segmentation mask\n            vessel_path = os.path.join(self.vessel_dir, fname)\n            if os.path.exists(vessel_path):\n                vessel_mask = nib.load(vessel_path).get_fdata().astype(np.float32)\n                # Resample vessel mask to match the processed image's shape\n                zoom_factor = np.array(img.shape) / np.array(vessel_mask.shape)\n                vessel_mask = zoom(vessel_mask, zoom_factor, order=0)\n            else:\n                # If no vessel mask exists, create a blank one\n                vessel_mask = np.zeros_like(img)\n\n            # 3. Stack image and vessel mask into a 2-channel input\n            # The final shape will be (2, H, W, D)\n            img = np.stack([img, vessel_mask], axis=0)\n            \n            # Transpose to PyTorch's expected format: (C, D, H, W)\n            img = np.transpose(img, (0, 3, 1, 2))\n            \n            sample = {'image': img, 'mask': aneurysm_mask}\n            \n            if self.transform:\n                sample = self.transform(sample)\n                \n            return sample\n        except Exception as e:\n            print(f\"\\n  - Error loading file {fname}: {e}. Skipping.\")\n            return self.__getitem__((idx + 1) % len(self))\n\n# --- Create Datasets and DataLoaders ---\nlabeled_files = sorted([f for f in os.listdir(MASKS_DIR) if f.endswith('.nii.gz')])\ntrain_files, val_files = train_test_split(labeled_files, test_size=VAL_SPLIT, random_state=42)\n\n# Pass the VESSEL_SEG_DIR path when creating the datasets\ntrain_dataset = AneurysmDataset(PROCESSED_IMAGES_DIR, MASKS_DIR, VESSEL_SEG_DIR, train_files, transform=train_transforms)\nval_dataset = AneurysmDataset(PROCESSED_IMAGES_DIR, MASKS_DIR, VESSEL_SEG_DIR, val_files, transform=val_transforms)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)\n\nprint(f\"✅ DataLoaders are ready with 2-channel input.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T15:14:41.169572Z","iopub.execute_input":"2025-10-03T15:14:41.170274Z","iopub.status.idle":"2025-10-03T15:14:41.230817Z","shell.execute_reply.started":"2025-10-03T15:14:41.170239Z","shell.execute_reply":"2025-10-03T15:14:41.229438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nprint(\"--- Verifying a batch from the DataLoader ---\")\n\ntry:\n    # Get one batch of data from the training loader\n    batch = next(iter(train_loader))\n    img_tensor = batch['image'][0]  # Get the first image in the batch\n    mask_tensor = batch['mask'][0] # Get the corresponding mask\n\n    print(f\"Image tensor shape: {img_tensor.shape}\")\n    print(f\"Mask tensor shape: {mask_tensor.shape}\")\n\n    # Visualize a slice from the middle of the 3D patch\n    slice_idx = img_tensor.shape[1] // 2 # Middle slice (dim 1 is Depth)\n\n    plt.figure(figsize=(10, 5))\n    plt.subplot(1, 2, 1)\n    plt.imshow(img_tensor[0, slice_idx, :, :].cpu(), cmap='gray')\n    plt.title(f'Image Patch (Slice {slice_idx})')\n    plt.axis('off')\n\n    plt.subplot(1, 2, 2)\n    # Take the max across the 13 channels to see the combined mask\n    plt.imshow(torch.max(mask_tensor, dim=0)[0][slice_idx, :, :].cpu(), cmap='gray')\n    plt.title(f'Mask Patch (Slice {slice_idx})')\n    plt.axis('off')\n    plt.show()\n\nexcept Exception as e:\n    print(f\"An error occurred while fetching a batch: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T15:28:57.279321Z","iopub.execute_input":"2025-10-03T15:28:57.280434Z","iopub.status.idle":"2025-10-03T15:28:57.311014Z","shell.execute_reply.started":"2025-10-03T15:28:57.280391Z","shell.execute_reply":"2025-10-03T15:28:57.309699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet(\n    spatial_dims=3,\n    in_channels=2, # <-- EDIT: Changed from 1 to 2 to accept the new input\n    out_channels=13,\n    channels=(16, 32, 64, 128, 256), strides=(2, 2, 2, 2), num_res_units=2\n).to(DEVICE)\n\nloss_function = DiceCELoss(to_onehot_y=False, sigmoid=True)\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)\n\nprint(f\"Model upgraded for 2-channel input.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T15:14:50.167158Z","iopub.execute_input":"2025-10-03T15:14:50.167497Z","iopub.status.idle":"2025-10-03T15:14:50.261197Z","shell.execute_reply.started":"2025-10-03T15:14:50.167463Z","shell.execute_reply":"2025-10-03T15:14:50.259939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- FINAL CORRECTED Training & Validation Loop ---\n\nbest_val_loss = float('inf')\ntrain_losses = []\nval_losses = []\n\nfor epoch in range(NUM_EPOCHS):\n    print(f\"\\n--- Epoch {epoch + 1}/{NUM_EPOCHS} ---\")\n    \n    # --- Training Phase ---\n    model.train()\n    epoch_loss = 0\n    for batch_data in tqdm(train_loader, desc=\"Training\"):\n        # --- FIX IS HERE ---\n        # The DataLoader has already prepared the batch correctly.\n        # Access the tensors directly from the dictionary.\n        inputs = batch_data['image'].to(DEVICE)\n        labels = batch_data['mask'].to(DEVICE)\n        \n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = loss_function(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        epoch_loss += loss.item()\n    \n    avg_train_loss = epoch_loss / len(train_loader)\n    train_losses.append(avg_train_loss)\n    print(f\"Epoch {epoch + 1} Average Training Loss: {avg_train_loss:.4f}\")\n\n    # --- Validation Phase ---\n    model.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for batch_data in tqdm(val_loader, desc=\"Validating\"):\n            # --- APPLY THE SAME FIX HERE ---\n            inputs = batch_data['image'].to(DEVICE)\n            labels = batch_data['mask'].to(DEVICE)\n            \n            outputs = model(inputs)\n            val_loss += loss_function(outputs, labels).item()\n    \n    avg_val_loss = val_loss / len(val_loader)\n    val_losses.append(avg_val_loss)\n    print(f\"Epoch {epoch + 1} Average Validation Loss: {avg_val_loss:.4f}\")\n\n    # --- Save the Best Model ---\n    if avg_val_loss < best_val_loss:\n        best_val_loss = avg_val_loss\n        torch.save(model.state_dict(), '/kaggle/working/best_model.pth')\n        print(f\"🎉 New best model saved with validation loss: {best_val_loss:.4f}\")\n\nprint(\"\\n--- Training Complete! ---\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-03T15:14:50.262282Z","iopub.execute_input":"2025-10-03T15:14:50.263077Z","execution_failed":"2025-10-03T15:28:49.886Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(10, 5))\nplt.plot(train_losses, label='Training Loss')\nplt.plot(val_losses, label='Validation Loss')\nplt.title('Training & Validation Loss Over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2025-10-03T15:28:49.886Z"}},"outputs":[],"execution_count":null}]}