{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":"nvidiaTeslaT4","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"sourceType":"competition"}],"dockerImageVersionId":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA 2025 Intracranial Aneurysm Detection - Complete Pipeline\n## Stage 1: ROI Extraction + Stage 2: Multi-Task Learning","metadata":{}},{"cell_type":"markdown","source":"## Dependencies","metadata":{}},{"cell_type":"code","source":"!pip install -q nibabel dicom2nifti pydicom","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:35.698423Z","iopub.execute_input":"2025-10-28T08:34:35.698755Z","iopub.status.idle":"2025-10-28T08:34:39.224258Z","shell.execute_reply.started":"2025-10-28T08:34:35.69873Z","shell.execute_reply":"2025-10-28T08:34:39.223391Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score, accuracy_score\n\nimport nibabel as nib\nimport nibabel.orientations as nio\nimport pydicom\nfrom scipy.ndimage import zoom, binary_dilation, rotate\nfrom scipy.ndimage.filters import gaussian_filter\nimport time","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.226577Z","iopub.execute_input":"2025-10-28T08:34:39.226879Z","iopub.status.idle":"2025-10-28T08:34:39.233465Z","shell.execute_reply.started":"2025-10-28T08:34:39.226857Z","shell.execute_reply":"2025-10-28T08:34:39.232686Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    \"\"\"Global configuration\"\"\"\n    # Paths\n    data_dir = '/kaggle/input/rsna-intracranial-aneurysm-detection'\n    train_csv = f'{data_dir}/train.csv'\n    localizers_csv = f'{data_dir}/train_localizers.csv'\n    series_dir = f'{data_dir}/series'\n    segmentations_dir = f'{data_dir}/segmentations'\n    \n    work_dir = '/kaggle/working'\n    nifti_cache_dir = f'{work_dir}/nifti_cache'\n    roi_crops_dir = f'{work_dir}/roi_crops'\n    \n    # Subset for testing\n    use_subset = True\n    n_subset_samples = 1000\n    \n    # Model\n    img_size = (128, 128, 128)\n    base_channels = 16\n    num_vessel_classes = 14  # 0=bg + 13 vessels\n    num_aneurysm_classes = 2\n    num_location_classes = 13\n    num_modality_classes = 4\n    \n    # Training\n    batch_size = 1\n    num_epochs = 100\n    learning_rate = 1e-3\n    weight_decay = 3e-5\n    train_split = 0.7\n    val_split = 0.15\n    \n    # Loss weights\n    vessel_seg_weight = 1.0\n    aneurysm_seg_weight = 2.0\n    binary_cls_weight = 3.0\n    location_cls_weight = 2.5\n    coord_reg_weight = 1.5\n    modality_cls_weight = 0.2\n    \n    # Augmentation\n    use_flip_augmentation = True\n    flip_probability = 0.5\n    use_rotation = True\n    use_scaling = True\n    use_noise = True\n    \n    # TTA\n    use_tta = True\n    tta_n_flips = 8\n    \n    # Optimization\n    mixed_precision = True\n    \n    # Other\n    seed = 42\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.234329Z","iopub.execute_input":"2025-10-28T08:34:39.234663Z","iopub.status.idle":"2025-10-28T08:34:39.251669Z","shell.execute_reply.started":"2025-10-28T08:34:39.234637Z","shell.execute_reply":"2025-10-28T08:34:39.250911Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"np.random.seed(CFG.seed)\ntorch.manual_seed(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.25262Z","iopub.execute_input":"2025-10-28T08:34:39.252987Z","iopub.status.idle":"2025-10-28T08:34:39.270856Z","shell.execute_reply.started":"2025-10-28T08:34:39.252964Z","shell.execute_reply":"2025-10-28T08:34:39.270228Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Constants","metadata":{}},{"cell_type":"code","source":"LOCATION_NAMES = [\n    'Left Infraclinoid Internal Carotid Artery',\n    'Right Infraclinoid Internal Carotid Artery',\n    'Left Supraclinoid Internal Carotid Artery',\n    'Right Supraclinoid Internal Carotid Artery',\n    'Left Middle Cerebral Artery',\n    'Right Middle Cerebral lo Artery',\n    'Anterior Communicating Artery',\n    'Left Anterior Cerebral Artery',\n    'Right Anterior Cerebral Artery',\n    'Left Posterior Communicating Artery',\n    'Right Posterior Communicating Artery',\n    'Basilar Tip',\n    'Other Posterior Circulation'\n]\n\nLEFT_RIGHT_VESSEL_PAIRS = {\n    3: 4, 4: 3,    # PComm\n    5: 6, 6: 5,    # Infraclinoid ICA\n    7: 8, 8: 7,    # Supraclinoid ICA\n    9: 10, 10: 9,  # MCA\n    11: 12, 12: 11 # ACA\n}\n\nLEFT_RIGHT_LOCATION_PAIRS = {\n    0: 1, 1: 0,    # Infraclinoid ICA\n    2: 3, 3: 2,    # Supraclinoid ICA\n    4: 5, 5: 4,    # MCA\n    7: 8, 8: 7,    # ACA\n    9: 10, 10: 9   # PComm\n}\n\nMODALITY_MAP = {'CTA': 0, 'MRA': 1, 'MRI_T1': 2, 'MRI_T2': 3}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.272824Z","iopub.execute_input":"2025-10-28T08:34:39.273231Z","iopub.status.idle":"2025-10-28T08:34:39.279592Z","shell.execute_reply.started":"2025-10-28T08:34:39.273216Z","shell.execute_reply":"2025-10-28T08:34:39.278784Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Loading","metadata":{}},{"cell_type":"code","source":"os.makedirs(CFG.nifti_cache_dir, exist_ok=True)\nos.makedirs(CFG.roi_crops_dir, exist_ok=True)\n\ntrain_df = pd.read_csv(CFG.train_csv)\nlocalizers_df = pd.read_csv(CFG.localizers_csv)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.280258Z","iopub.execute_input":"2025-10-28T08:34:39.280496Z","iopub.status.idle":"2025-10-28T08:34:39.322113Z","shell.execute_reply.started":"2025-10-28T08:34:39.280472Z","shell.execute_reply":"2025-10-28T08:34:39.321269Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 1: ROI Extraction Functions","metadata":{}},{"cell_type":"code","source":"def reorient_to_lps(nifti_img):\n    \"\"\"Reorient NIfTI image to LPS orientation\"\"\"\n    orig_ornt = nio.io_orientation(nifti_img.affine)\n    lps_ornt = nio.axcodes2ornt(('L', 'P', 'S'))\n    transform = nio.ornt_transform(orig_ornt, lps_ornt)\n    data_reoriented = nio.apply_orientation(nifti_img.get_fdata(), transform)\n    affine_reoriented = nifti_img.affine @ nio.inv_ornt_aff(transform, nifti_img.shape)\n    return nib.Nifti1Image(data_reoriented, affine_reoriented)\n\n\ndef load_dicom_to_nifti(series_uid, cache_dir):\n    \"\"\"Load DICOM series and convert to NIfTI with caching\"\"\"\n    cache_path = Path(cache_dir) / f\"{series_uid}.nii.gz\"\n    \n    if cache_path.exists():\n        return nib.load(str(cache_path))\n    \n    series_path = Path(CFG.series_dir) / series_uid\n    if not series_path.exists():\n        return None\n    \n    dicom_files = sorted(list(series_path.glob('*.dcm')))\n    if len(dicom_files) == 0:\n        return None\n    \n    slices = []\n    for dcm_file in dicom_files:\n        try:\n            dcm = pydicom.dcmread(str(dcm_file))\n            _ = dcm.pixel_array\n            slices.append(dcm)\n        except:\n            continue\n    \n    if len(slices) == 0:\n        return None\n    \n    try:\n        slices.sort(key=lambda x: float(x.ImagePositionPatient[2]))\n    except:\n        try:\n            slices.sort(key=lambda x: int(x.InstanceNumber))\n        except:\n            pass\n    \n    first_array = slices[0].pixel_array\n    if len(first_array.shape) == 3:\n        first_array = first_array[0]\n    \n    img_shape = first_array.shape\n    volume = np.zeros((img_shape[0], img_shape[1], len(slices)), dtype=np.float32)\n    \n    for i, dcm_slice in enumerate(slices):\n        try:\n            arr = dcm_slice.pixel_array.astype(np.float32)\n            if len(arr.shape) == 3:\n                arr = arr[0]\n            if arr.shape != img_shape:\n                arr = zoom(arr, (img_shape[0]/arr.shape[0], img_shape[1]/arr.shape[1]), order=1)\n            volume[:, :, i] = arr\n        except:\n            volume[:, :, i] = np.zeros(img_shape, dtype=np.float32)\n    \n    affine = np.eye(4)\n    nifti_img = nib.Nifti1Image(volume, affine)\n    nifti_img = reorient_to_lps(nifti_img)\n    \n    nib.save(nifti_img, str(cache_path))\n    return nifti_img\n\n\ndef load_vessel_segmentation(series_uid, segmentations_dir):\n    \"\"\"Load vessel segmentation from segmentations folder\"\"\"\n    seg_path = Path(segmentations_dir) / f\"{series_uid}_cowseg.nii\"\n    \n    if seg_path.exists():\n        seg_nii = nib.load(str(seg_path))\n        seg_nii = reorient_to_lps(seg_nii)\n        return seg_nii.get_fdata().astype(np.int64)\n    \n    return None\n\n\ndef create_bbox_from_mask(mask, padding_ratio=0.15):\n    \"\"\"Create 3D bounding box from binary/multi-class mask\"\"\"\n    coords = np.where(mask > 0)\n    \n    if len(coords[0]) == 0:\n        return None\n    \n    x_min, x_max = coords[0].min(), coords[0].max()\n    y_min, y_max = coords[1].min(), coords[1].max()\n    z_min, z_max = coords[2].min(), coords[2].max()\n    \n    x_range = max(x_max - x_min, 1)  # Mínimo 1\n    y_range = max(y_max - y_min, 1)\n    z_range = max(z_max - z_min, 1)\n    \n    x_pad = int(x_range * padding_ratio)\n    y_pad = int(y_range * padding_ratio)\n    z_pad = int(z_range * padding_ratio)\n    \n    x_min = max(0, x_min - x_pad)\n    x_max = min(mask.shape[0], x_max + x_pad + 1)\n    y_min = max(0, y_min - y_pad)\n    y_max = min(mask.shape[1], y_max + y_pad + 1)\n    z_min = max(0, z_min - z_pad)\n    z_max = min(mask.shape[2], z_max + z_pad + 1)\n    \n    # Validar bbox\n    if x_min >= x_max or y_min >= y_max or z_min >= z_max:\n        return None\n    \n    return (x_min, x_max, y_min, y_max, z_min, z_max)\n\n\ndef create_bbox_from_localizers(localizers_df_subset, volume_shape):\n    \"\"\"Create bbox from localizer coordinates when no segmentation available\"\"\"\n    if len(localizers_df_subset) == 0:\n        return None\n    \n    all_coords = []\n    for _, row in localizers_df_subset.iterrows():\n        try:\n            coords = eval(row['coordinates'])\n            if isinstance(coords, dict):\n                x = coords.get('x', volume_shape[0] / 2)\n                y = coords.get('y', volume_shape[1] / 2)\n                z = coords.get('z', volume_shape[2] / 2)\n                all_coords.append([x, y, z])\n            else:\n                if len(coords) >= 2:\n                    x = coords[0]\n                    y = coords[1]\n                    z = coords[2] if len(coords) > 2 else volume_shape[2] / 2\n                    all_coords.append([x, y, z])\n        except:\n            continue\n    \n    if len(all_coords) == 0:\n        return None\n    \n    all_coords = np.array(all_coords)\n    \n    x_center = all_coords[:, 0].mean()\n    y_center = all_coords[:, 1].mean()\n    z_center = all_coords[:, 2].mean()\n    \n    box_size = 180\n    half_size = box_size // 2\n    \n    x_min = max(0, int(x_center - half_size))\n    x_max = min(volume_shape[0], int(x_center + half_size))\n    y_min = max(0, int(y_center - half_size))\n    y_max = min(volume_shape[1], int(y_center + half_size))\n    z_min = max(0, int(z_center - half_size))\n    z_max = min(volume_shape[2], int(z_center + half_size))\n    \n    # Validar bbox\n    if x_min >= x_max or y_min >= y_max or z_min >= z_max:\n        return None\n    \n    return (x_min, x_max, y_min, y_max, z_min, z_max)\n\n\ndef extract_roi(volume, bbox):\n    \"\"\"Extract ROI crop from volume using bbox\"\"\"\n    if bbox is None:\n        return volume\n    \n    x_min, x_max, y_min, y_max, z_min, z_max = bbox\n    \n    # Validar bbox\n    if x_min >= x_max or y_min >= y_max or z_min >= z_max:\n        return volume\n    \n    roi = volume[x_min:x_max, y_min:y_max, z_min:z_max]\n    \n    # Verificar que ROI no esté vacío\n    if roi.size == 0 or 0 in roi.shape:\n        return volume\n    \n    return roi\n\n\ndef resize_volume(volume, target_size):\n    \"\"\"Resize volume to target size\"\"\"\n    # Verificar volumen válido\n    if volume.size == 0 or 0 in volume.shape:\n        return np.zeros(target_size, dtype=volume.dtype)\n    \n    zoom_factors = [target_size[i] / volume.shape[i] for i in range(3)]\n    resized = zoom(volume, zoom_factors, order=1)\n    \n    resized = resized[:target_size[0], :target_size[1], :target_size[2]]\n    \n    if resized.shape != target_size:\n        padded = np.zeros(target_size, dtype=resized.dtype)\n        padded[:resized.shape[0], :resized.shape[1], :resized.shape[2]] = resized\n        resized = padded\n    \n    return resized\n\n\ndef process_stage1_roi_extraction(series_uid, localizers_df, target_size):\n    \"\"\"Complete Stage 1 pipeline: load, extract ROI, resize\"\"\"\n    nifti_img = load_dicom_to_nifti(series_uid, CFG.nifti_cache_dir)\n    if nifti_img is None:\n        return None, None, None\n    \n    volume = nifti_img.get_fdata().astype(np.float32)\n    \n    # Verificar volumen válido\n    if volume.size == 0 or 0 in volume.shape:\n        return None, None, None\n    \n    vessel_seg = load_vessel_segmentation(series_uid, CFG.segmentations_dir)\n    \n    if vessel_seg is not None and vessel_seg.size > 0:\n        bbox = create_bbox_from_mask(vessel_seg)\n        if bbox is not None:\n            vessel_seg_roi = extract_roi(vessel_seg, bbox)\n            vessel_seg_resized = resize_volume(vessel_seg_roi, target_size).astype(np.int64)\n        else:\n            vessel_seg_resized = None\n    else:\n        loc_subset = localizers_df[localizers_df['SeriesInstanceUID'] == series_uid]\n        bbox = create_bbox_from_localizers(loc_subset, volume.shape)\n        vessel_seg_resized = None\n    \n    volume_roi = extract_roi(volume, bbox)\n    \n    # Verificar que ROI sea válido\n    if volume_roi.size == 0 or 0 in volume_roi.shape:\n        volume_resized = resize_volume(volume, target_size)\n    else:\n        volume_resized = resize_volume(volume_roi, target_size)\n    \n    return volume_resized, vessel_seg_resized, bbox","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:48:59.970744Z","iopub.execute_input":"2025-10-28T08:48:59.971325Z","iopub.status.idle":"2025-10-28T08:48:59.997702Z","shell.execute_reply.started":"2025-10-28T08:48:59.971297Z","shell.execute_reply":"2025-10-28T08:48:59.996919Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 1: Visualization","metadata":{}},{"cell_type":"code","source":"def visualize_stage1_pipeline(series_uid, localizers_df):\n    \"\"\"Visualize Stage 1 ROI extraction pipeline\"\"\"\n    volume, vessel_seg, bbox = process_stage1_roi_extraction(\n        series_uid, localizers_df, CFG.img_size\n    )\n    \n    if volume is None:\n        return\n    \n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    slices_idx = [volume.shape[2]//4, volume.shape[2]//2, 3*volume.shape[2]//4]\n    \n    for i, slice_idx in enumerate(slices_idx):\n        axes[0, i].imshow(volume[:, :, slice_idx], cmap='gray')\n        axes[0, i].set_title(f'Volume Slice {slice_idx}')\n        axes[0, i].axis('off')\n        \n        if vessel_seg is not None:\n            axes[1, i].imshow(volume[:, :, slice_idx], cmap='gray', alpha=0.7)\n            axes[1, i].imshow(vessel_seg[:, :, slice_idx], cmap='jet', alpha=0.3, vmin=0, vmax=13)\n            axes[1, i].set_title(f'Vessel Seg Overlay {slice_idx}')\n        else:\n            axes[1, i].imshow(volume[:, :, slice_idx], cmap='gray')\n            axes[1, i].set_title(f'No Segmentation {slice_idx}')\n        axes[1, i].axis('off')\n    \n    plt.suptitle(f'Stage 1 ROI: {series_uid[:30]}...', fontsize=14)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.347099Z","iopub.execute_input":"2025-10-28T08:34:39.347732Z","iopub.status.idle":"2025-10-28T08:34:39.360456Z","shell.execute_reply.started":"2025-10-28T08:34:39.347714Z","shell.execute_reply":"2025-10-28T08:34:39.35971Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 2: Dataset","metadata":{}},{"cell_type":"code","source":"def generate_pseudo_vessel_mask(volume, target_size):\n    \"\"\"Generate pseudo vessel segmentation when real segmentation not available\"\"\"\n    threshold = np.percentile(volume, 85)\n    binary_mask = (volume > threshold).astype(np.int64)\n    \n    struct = np.ones((3, 3, 3))\n    binary_mask = binary_dilation(binary_mask, structure=struct, iterations=2).astype(np.int64)\n    \n    vessel_mask = np.zeros_like(binary_mask, dtype=np.int64)\n    \n    vessel_coords = np.where(binary_mask > 0)\n    if len(vessel_coords[0]) > 0:\n        weights = np.array([0.05, 0.08, 0.08, 0.08, 0.08, 0.12, 0.12, 0.12, 0.12, 0.05, 0.05, 0.03, 0.02])\n        labels = np.random.choice(13, size=len(vessel_coords[0]), p=weights) + 1\n        vessel_mask[vessel_coords] = labels\n    \n    return vessel_mask\n\n\ndef create_aneurysm_mask(series_uid, localizers_df, target_size, radius=12):\n    \"\"\"Create aneurysm segmentation mask from localizer coordinates\"\"\"\n    mask = np.zeros(target_size, dtype=np.int64)\n    \n    loc_subset = localizers_df[localizers_df['SeriesInstanceUID'] == series_uid]\n    \n    for _, row in loc_subset.iterrows():\n        try:\n            coords = eval(row['coordinates'])\n            if len(coords) >= 3:\n                x = int(coords[0] * target_size[0] / 512)\n                y = int(coords[1] * target_size[1] / 512)\n                z = int(coords[2] * target_size[2] / 512)\n                \n                x = np.clip(x, 0, target_size[0]-1)\n                y = np.clip(y, 0, target_size[1]-1)\n                z = np.clip(z, 0, target_size[2]-1)\n                \n                for i in range(max(0, x-radius), min(target_size[0], x+radius)):\n                    for j in range(max(0, y-radius), min(target_size[1], y+radius)):\n                        for k in range(max(0, z-radius), min(target_size[2], z+radius)):\n                            dist = np.sqrt((i-x)**2 + (j-y)**2 + (k-z)**2)\n                            if dist <= radius:\n                                mask[i, j, k] = 1\n        except:\n            continue\n    \n    return mask\n\n\ndef create_aneurysm_heatmap(series_uid, localizers_df, target_size, sigma=15):\n    \"\"\"Create gaussian heatmap centered on aneurysm for loss weighting\"\"\"\n    heatmap = np.zeros(target_size, dtype=np.float32)\n    \n    loc_subset = localizers_df[localizers_df['SeriesInstanceUID'] == series_uid]\n    \n    for _, row in loc_subset.iterrows():\n        try:\n            coords = eval(row['coordinates'])\n            if len(coords) >= 3:\n                x = int(coords[0] * target_size[0] / 512)\n                y = int(coords[1] * target_size[1] / 512)\n                z = int(coords[2] * target_size[2] / 512)\n                \n                x = np.clip(x, 0, target_size[0]-1)\n                y = np.clip(y, 0, target_size[1]-1)\n                z = np.clip(z, 0, target_size[2]-1)\n                \n                point_heatmap = np.zeros(target_size, dtype=np.float32)\n                point_heatmap[x, y, z] = 1.0\n                point_heatmap = gaussian_filter(point_heatmap, sigma=sigma)\n                heatmap += point_heatmap\n        except:\n            continue\n    \n    if heatmap.max() > 0:\n        heatmap = heatmap / heatmap.max()\n    \n    heatmap = heatmap * 5.0 + 1.0\n    \n    return heatmap\n\n\ndef extract_location_labels(row):\n    \"\"\"Extract 13-location binary labels from dataframe row\"\"\"\n    labels = np.zeros(13, dtype=np.float32)\n    for i, loc_name in enumerate(LOCATION_NAMES):\n        labels[i] = float(row.get(loc_name, 0))\n    return labels\n\n\ndef swap_vessel_labels(mask):\n    \"\"\"Swap left-right vessel labels after horizontal flip\"\"\"\n    swapped = mask.copy()\n    for left_label, right_label in LEFT_RIGHT_VESSEL_PAIRS.items():\n        swapped[mask == left_label] = right_label\n    return swapped\n\n\ndef swap_location_labels(labels):\n    \"\"\"Swap left-right location labels after horizontal flip\"\"\"\n    swapped = labels.copy()\n    for left_idx, right_idx in LEFT_RIGHT_LOCATION_PAIRS.items():\n        swapped[left_idx] = labels[right_idx]\n        swapped[right_idx] = labels[left_idx]\n    return swapped\n\n\nclass AneurysmDataset(Dataset):\n    \"\"\"Multi-task dataset for aneurysm detection\"\"\"\n    \n    def __init__(self, df, localizers_df, target_size, augment):\n        self.df = df.reset_index(drop=True)\n        self.localizers_df = localizers_df\n        self.target_size = target_size\n        self.augment = augment\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        series_uid = row['SeriesInstanceUID']\n        \n        volume, vessel_seg, _ = process_stage1_roi_extraction(\n            series_uid, self.localizers_df, self.target_size\n        )\n        \n        if volume is None:\n            volume = np.random.randn(*self.target_size).astype(np.float32)\n        \n        p1, p99 = np.percentile(volume, 1), np.percentile(volume, 99)\n        volume = np.clip(volume, p1, p99)\n        v_min, v_max = volume.min(), volume.max()\n        if v_max > v_min:\n            volume = (volume - v_min) / (v_max - v_min)\n        \n        if vessel_seg is None:\n            vessel_seg = generate_pseudo_vessel_mask(volume, self.target_size)\n        \n        aneurysm_mask = create_aneurysm_mask(series_uid, self.localizers_df, self.target_size)\n        aneurysm_heatmap = create_aneurysm_heatmap(series_uid, self.localizers_df, self.target_size)\n        \n        location_labels = extract_location_labels(row)\n        binary_label = float(row['Aneurysm Present'])\n        \n        modality = row.get('Modality', 'CTA')\n        if 'T1' in modality:\n            modality_label = MODALITY_MAP['MRI_T1']\n        elif 'T2' in modality:\n            modality_label = MODALITY_MAP['MRI_T2']\n        else:\n            modality_label = MODALITY_MAP.get(modality, 0)\n        \n        if self.augment:\n            if CFG.use_flip_augmentation and np.random.random() > 0.5:\n                volume = np.flip(volume, axis=0).copy()\n                vessel_seg = swap_vessel_labels(np.flip(vessel_seg, axis=0).copy())\n                aneurysm_mask = np.flip(aneurysm_mask, axis=0).copy()\n                aneurysm_heatmap = np.flip(aneurysm_heatmap, axis=0).copy()\n                location_labels = swap_location_labels(location_labels)\n            \n            if CFG.use_rotation and np.random.random() > 0.7:\n                angle = np.random.uniform(-15, 15)\n                axes = np.random.choice([0, 1, 2], size=2, replace=False)\n                volume = rotate(volume, angle, axes=axes, reshape=False, order=1)\n                vessel_seg = rotate(vessel_seg, angle, axes=axes, reshape=False, order=0)\n                aneurysm_mask = rotate(aneurysm_mask, angle, axes=axes, reshape=False, order=0)\n            \n            if CFG.use_noise and np.random.random() > 0.7:\n                volume = volume + np.random.randn(*volume.shape) * 0.05\n                volume = np.clip(volume, 0, 1)\n        \n        volume = torch.from_numpy(volume).unsqueeze(0).float()\n        vessel_seg = torch.from_numpy(vessel_seg).long()\n        aneurysm_mask = torch.from_numpy(aneurysm_mask).long()\n        aneurysm_heatmap = torch.from_numpy(aneurysm_heatmap).float()\n        location_labels = torch.from_numpy(location_labels).float()\n        binary_label = torch.tensor([binary_label], dtype=torch.float32)\n        modality_label = torch.tensor(modality_label, dtype=torch.long)\n        \n        return {\n            'image': volume,\n            'vessel_mask': vessel_seg,\n            'aneurysm_mask': aneurysm_mask,\n            'aneurysm_heatmap': aneurysm_heatmap,\n            'location_labels': location_labels,\n            'binary_label': binary_label,\n            'modality_label': modality_label,\n            'series_uid': series_uid\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.361102Z","iopub.execute_input":"2025-10-28T08:34:39.361318Z","iopub.status.idle":"2025-10-28T08:34:39.386518Z","shell.execute_reply.started":"2025-10-28T08:34:39.361303Z","shell.execute_reply":"2025-10-28T08:34:39.385742Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 2: Model Architecture","metadata":{}},{"cell_type":"code","source":"class CrossAttentionPooling(nn.Module):\n    \"\"\"Cross-attention pooling for global feature aggregation\"\"\"\n    \n    def __init__(self, dim, num_heads=8):\n        super().__init__()\n        self.num_heads = num_heads\n        self.scale = (dim // num_heads) ** -0.5\n        self.query = nn.Parameter(torch.randn(1, 1, dim))\n        self.kv = nn.Linear(dim, dim * 2, bias=False)\n        self.proj = nn.Linear(dim, dim)\n    \n    def forward(self, x):\n        B, C, D, H, W = x.shape\n        x_flat = x.flatten(2).transpose(1, 2)\n        \n        q = self.query.expand(B, -1, -1)\n        kv = self.kv(x_flat).reshape(B, -1, 2, self.num_heads, C // self.num_heads)\n        k, v = kv.unbind(2)\n        \n        q = q.reshape(B, 1, self.num_heads, C // self.num_heads).transpose(1, 2)\n        k = k.transpose(1, 2)\n        v = v.transpose(1, 2)\n        \n        attn = (q @ k.transpose(-2, -1)) * self.scale\n        attn = attn.softmax(dim=-1)\n        \n        out = (attn @ v).transpose(1, 2).reshape(B, 1, C)\n        out = self.proj(out)\n        \n        return out.squeeze(1)\n\n\nclass ResidualBlock3D(nn.Module):\n    \"\"\"3D Residual block\"\"\"\n    \n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv3d(in_channels, out_channels, 3, stride=stride, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm3d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = nn.Conv3d(out_channels, out_channels, 3, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm3d(out_channels)\n        \n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv3d(in_channels, out_channels, 1, stride=stride, bias=False),\n                nn.BatchNorm3d(out_channels)\n            )\n        else:\n            self.shortcut = nn.Identity()\n    \n    def forward(self, x):\n        identity = self.shortcut(x)\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out += identity\n        out = self.relu(out)\n        return out\n\n\nclass DecoderBlock3D(nn.Module):\n    \"\"\"3D Decoder block with skip connections\"\"\"\n    \n    def __init__(self, in_channels, skip_channels, out_channels):\n        super().__init__()\n        self.upconv = nn.ConvTranspose3d(in_channels, in_channels // 2, 2, stride=2)\n        self.conv = nn.Sequential(\n            nn.Conv3d(in_channels // 2 + skip_channels, out_channels, 3, padding=1, bias=False),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(out_channels, out_channels, 3, padding=1, bias=False),\n            nn.BatchNorm3d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x, skip):\n        x = self.upconv(x)\n        x = torch.cat([x, skip], dim=1)\n        x = self.conv(x)\n        return x\n\n\nclass nnXNetLite(nn.Module):\n    \"\"\"Lightweight nnXNet-inspired multi-task architecture\"\"\"\n    \n    def __init__(self, in_channels, num_vessel_classes, num_aneurysm_classes,\n                 num_location_classes, num_modality_classes, base_channels):\n        super().__init__()\n        \n        self.enc1 = ResidualBlock3D(in_channels, base_channels)\n        self.enc2 = ResidualBlock3D(base_channels, base_channels * 2, stride=2)\n        self.enc3 = ResidualBlock3D(base_channels * 2, base_channels * 4, stride=2)\n        self.enc4 = ResidualBlock3D(base_channels * 4, base_channels * 8, stride=2)\n        \n        self.bottleneck = ResidualBlock3D(base_channels * 8, base_channels * 16, stride=2)\n        \n        self.cross_attn = CrossAttentionPooling(base_channels * 16)\n        \n        self.dec4 = DecoderBlock3D(base_channels * 16, base_channels * 8, base_channels * 8)\n        self.dec3 = DecoderBlock3D(base_channels * 8, base_channels * 4, base_channels * 4)\n        self.dec2 = DecoderBlock3D(base_channels * 4, base_channels * 2, base_channels * 2)\n        self.dec1 = DecoderBlock3D(base_channels * 2, base_channels, base_channels)\n        \n        self.vessel_seg_head = nn.Conv3d(base_channels, num_vessel_classes, 1)\n        self.aneurysm_seg_head = nn.Conv3d(base_channels, num_aneurysm_classes, 1)\n        \n        self.binary_cls_head = nn.Sequential(\n            nn.Linear(base_channels * 16, 256),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, 1)\n        )\n        \n        self.location_cls_head = nn.Sequential(\n            nn.Linear(base_channels * 16, 256),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, num_location_classes)\n        )\n        \n        self.coord_reg_head = nn.Sequential(\n            nn.Linear(base_channels * 16, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3)\n        )\n        \n        self.modality_cls_head = nn.Sequential(\n            nn.Linear(base_channels * 16, 64),\n            nn.ReLU(),\n            nn.Linear(64, num_modality_classes)\n        )\n    \n    def forward(self, x):\n        e1 = self.enc1(x)\n        e2 = self.enc2(e1)\n        e3 = self.enc3(e2)\n        e4 = self.enc4(e3)\n        \n        bottleneck = self.bottleneck(e4)\n        \n        global_feat = self.cross_attn(bottleneck)\n        \n        d4 = self.dec4(bottleneck, e4)\n        d3 = self.dec3(d4, e3)\n        d2 = self.dec2(d3, e2)\n        d1 = self.dec1(d2, e1)\n        \n        vessel_seg = self.vessel_seg_head(d1)\n        aneurysm_seg = self.aneurysm_seg_head(d1)\n        \n        binary_cls = self.binary_cls_head(global_feat)\n        location_cls = self.location_cls_head(global_feat)\n        coord_reg = self.coord_reg_head(global_feat)\n        modality_cls = self.modality_cls_head(global_feat)\n        \n        return {\n            'vessel_seg': vessel_seg,\n            'aneurysm_seg': aneurysm_seg,\n            'binary_cls': binary_cls,\n            'location_cls': location_cls,\n            'coord_reg': coord_reg,\n            'modality_cls': modality_cls\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.387465Z","iopub.execute_input":"2025-10-28T08:34:39.387688Z","iopub.status.idle":"2025-10-28T08:34:39.410431Z","shell.execute_reply.started":"2025-10-28T08:34:39.387672Z","shell.execute_reply":"2025-10-28T08:34:39.409521Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Stage 2: Loss Functions","metadata":{}},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    \"\"\"Dice loss for segmentation\"\"\"\n    \n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n    \n    def forward(self, pred, target):\n        pred = F.softmax(pred, dim=1)\n        target_one_hot = F.one_hot(target, num_classes=pred.shape[1])\n        target_one_hot = target_one_hot.permute(0, 4, 1, 2, 3).float()\n        \n        intersection = (pred * target_one_hot).sum(dim=(2, 3, 4))\n        union = pred.sum(dim=(2, 3, 4)) + target_one_hot.sum(dim=(2, 3, 4))\n        \n        dice = (2.0 * intersection + self.smooth) / (union + self.smooth)\n        return 1.0 - dice.mean()\n\n\nclass HeatmapWeightedCE(nn.Module):\n    \"\"\"Cross-entropy with spatial heatmap weighting\"\"\"\n    \n    def __init__(self):\n        super().__init__()\n    \n    def forward(self, pred, target, heatmap):\n        ce_loss = F.cross_entropy(pred, target, reduction='none')\n        weighted_loss = ce_loss * heatmap\n        return weighted_loss.mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.411518Z","iopub.execute_input":"2025-10-28T08:34:39.412092Z","iopub.status.idle":"2025-10-28T08:34:39.426947Z","shell.execute_reply.started":"2025-10-28T08:34:39.412064Z","shell.execute_reply":"2025-10-28T08:34:39.426073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiTaskLoss(nn.Module):\n    def __init__(self, vessel_seg_weight, aneurysm_seg_weight, binary_cls_weight,\n                 location_cls_weight, coord_reg_weight, modality_cls_weight):\n        super().__init__()\n        self.vessel_seg_weight = vessel_seg_weight\n        self.aneurysm_seg_weight = aneurysm_seg_weight\n        self.binary_cls_weight = binary_cls_weight\n        self.location_cls_weight = location_cls_weight\n        self.coord_reg_weight = coord_reg_weight\n        self.modality_cls_weight = modality_cls_weight\n        \n        self.dice_loss = DiceLoss()\n        self.heatmap_ce = HeatmapWeightedCE()\n    \n    def forward(self, outputs, targets):\n        losses = {}\n        \n        vessel_ce = F.cross_entropy(outputs['vessel_seg'], targets['vessel_mask'])\n        vessel_dice = self.dice_loss(outputs['vessel_seg'], targets['vessel_mask'])\n        losses['vessel_seg'] = (vessel_ce + vessel_dice) * self.vessel_seg_weight\n        \n        aneurysm_ce = self.heatmap_ce(\n            outputs['aneurysm_seg'],\n            targets['aneurysm_mask'],\n            targets['aneurysm_heatmap']\n        )\n        aneurysm_dice = self.dice_loss(outputs['aneurysm_seg'], targets['aneurysm_mask'])\n        losses['aneurysm_seg'] = (aneurysm_ce + aneurysm_dice) * self.aneurysm_seg_weight\n        \n        losses['binary_cls'] = F.binary_cross_entropy_with_logits(\n            outputs['binary_cls'],\n            targets['binary_label']\n        ) * self.binary_cls_weight\n        \n        losses['location_cls'] = F.binary_cross_entropy_with_logits(\n            outputs['location_cls'],\n            targets['location_labels']\n        ) * self.location_cls_weight\n        \n        # FIX: Solo calcular coord loss si hay aneurismas reales\n        has_aneurysm = targets['binary_label'] > 0.5\n        if has_aneurysm.any():\n            aneurysm_coords = []\n            for i in range(len(targets['aneurysm_mask'])):\n                mask = targets['aneurysm_mask'][i].cpu().numpy()\n                coords_idx = np.where(mask > 0)\n                \n                # FIX: Verificar si hay píxeles antes de calcular mean\n                if len(coords_idx[0]) > 0:\n                    coords = np.array([coords_idx[0].mean(), coords_idx[1].mean(), coords_idx[2].mean()])\n                else:\n                    # Default al centro si máscara vacía\n                    coords = np.array([mask.shape[0]//2, mask.shape[1]//2, mask.shape[2]//2], dtype=np.float32)\n                \n                aneurysm_coords.append(coords)\n            \n            aneurysm_coords = torch.tensor(aneurysm_coords, device=outputs['coord_reg'].device).float()\n            \n            losses['coord_reg'] = F.mse_loss(\n                outputs['coord_reg'][has_aneurysm.squeeze()],\n                aneurysm_coords[has_aneurysm.squeeze()]\n            ) * self.coord_reg_weight\n        else:\n            losses['coord_reg'] = torch.tensor(0.0, device=outputs['coord_reg'].device)\n        \n        losses['modality_cls'] = F.cross_entropy(\n            outputs['modality_cls'],\n            targets['modality_label']\n        ) * self.modality_cls_weight\n        \n        losses['total'] = sum(losses.values())\n        \n        return losses","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.427827Z","iopub.execute_input":"2025-10-28T08:34:39.428118Z","iopub.status.idle":"2025-10-28T08:34:39.44072Z","shell.execute_reply.started":"2025-10-28T08:34:39.428094Z","shell.execute_reply":"2025-10-28T08:34:39.439895Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Functions","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, loader, criterion, optimizer, scaler, device):\n    \"\"\"Train one epoch\"\"\"\n    model.train()\n    total_losses = {}\n    \n    for batch in loader:\n        image = batch['image'].to(device)\n        targets = {\n            'vessel_mask': batch['vessel_mask'].to(device),\n            'aneurysm_mask': batch['aneurysm_mask'].to(device),\n            'aneurysm_heatmap': batch['aneurysm_heatmap'].to(device),\n            'location_labels': batch['location_labels'].to(device),\n            'binary_label': batch['binary_label'].to(device),\n            'modality_label': batch['modality_label'].to(device)\n        }\n        \n        optimizer.zero_grad()\n        \n        with autocast(enabled=CFG.mixed_precision):\n            outputs = model(image)\n            losses = criterion(outputs, targets)\n        \n        scaler.scale(losses['total']).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        for key, val in losses.items():\n            if key not in total_losses:\n                total_losses[key] = 0\n            total_losses[key] += val.item()\n    \n    for key in total_losses:\n        total_losses[key] /= len(loader)\n    \n    return total_losses\n\n\ndef validate(model, loader, criterion, device):\n    \"\"\"Validate model\"\"\"\n    model.eval()\n    total_losses = {}\n    all_binary_preds = []\n    all_binary_labels = []\n    \n    with torch.no_grad():\n        for batch in loader:\n            image = batch['image'].to(device)\n            targets = {\n                'vessel_mask': batch['vessel_mask'].to(device),\n                'aneurysm_mask': batch['aneurysm_mask'].to(device),\n                'aneurysm_heatmap': batch['aneurysm_heatmap'].to(device),\n                'location_labels': batch['location_labels'].to(device),\n                'binary_label': batch['binary_label'].to(device),\n                'modality_label': batch['modality_label'].to(device)\n            }\n            \n            outputs = model(image)\n            losses = criterion(outputs, targets)\n            \n            for key, val in losses.items():\n                if key not in total_losses:\n                    total_losses[key] = 0\n                total_losses[key] += val.item()\n            \n            binary_pred = torch.sigmoid(outputs['binary_cls']).cpu().numpy()\n            binary_label = targets['binary_label'].cpu().numpy()\n            \n            all_binary_preds.extend(binary_pred)\n            all_binary_labels.extend(binary_label)\n    \n    for key in total_losses:\n        total_losses[key] /= len(loader)\n    \n    all_binary_preds = np.array(all_binary_preds)\n    all_binary_labels = np.array(all_binary_labels)\n    \n    if len(np.unique(all_binary_labels)) > 1:\n        total_losses['auc'] = roc_auc_score(all_binary_labels, all_binary_preds)\n    else:\n        total_losses['auc'] = 0.0\n    \n    total_losses['acc'] = accuracy_score(all_binary_labels, all_binary_preds > 0.5)\n    \n    return total_losses","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.441531Z","iopub.execute_input":"2025-10-28T08:34:39.441788Z","iopub.status.idle":"2025-10-28T08:34:39.458229Z","shell.execute_reply.started":"2025-10-28T08:34:39.44177Z","shell.execute_reply":"2025-10-28T08:34:39.457435Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TTA Inference","metadata":{}},{"cell_type":"code","source":"def predict_with_tta(model, volume, device, n_flips=8):\n    \"\"\"Test-time augmentation with 8 flips and label swapping\"\"\"\n    model.eval()\n    \n    transforms = [\n        lambda x: x,\n        lambda x: torch.flip(x, [2]),  # X flip\n        lambda x: torch.flip(x, [3]),  # Y flip\n        lambda x: torch.flip(x, [4]),  # Z flip\n        lambda x: torch.flip(x, [2, 3]),  # XY flip\n        lambda x: torch.flip(x, [2, 4]),  # XZ flip\n        lambda x: torch.flip(x, [3, 4]),  # YZ flip\n        lambda x: torch.flip(x, [2, 3, 4])  # XYZ flip\n    ]\n    \n    x_flip_indices = [1, 4, 5, 7]\n    \n    all_predictions = []\n    \n    with torch.no_grad():\n        for i, transform in enumerate(transforms[:n_flips]):\n            vol_transformed = transform(volume)\n            outputs = model(vol_transformed.to(device))\n            \n            if i in x_flip_indices:\n                location_probs = torch.sigmoid(outputs['location_cls'])\n                location_probs_np = location_probs.cpu().numpy()[0]\n                location_probs_swapped = swap_location_labels(location_probs_np)\n                outputs['location_cls'] = torch.from_numpy(location_probs_swapped).unsqueeze(0).to(device)\n            \n            all_predictions.append({\n                'binary_cls': torch.sigmoid(outputs['binary_cls']).cpu().numpy(),\n                'location_cls': torch.sigmoid(outputs['location_cls']).cpu().numpy(),\n                'aneurysm_seg': torch.softmax(outputs['aneurysm_seg'], dim=1).cpu().numpy()\n            })\n    \n    final_pred = {\n        'binary_cls': np.mean([p['binary_cls'] for p in all_predictions], axis=0),\n        'location_cls': np.mean([p['location_cls'] for p in all_predictions], axis=0),\n        'aneurysm_seg': np.mean([p['aneurysm_seg'] for p in all_predictions], axis=0)\n    }\n    \n    return final_pred\n\n\ndef extract_aneurysm_info(aneurysm_seg, location_probs, img_size):\n    \"\"\"Extract aneurysm location info from predictions\"\"\"\n    binary_mask = (aneurysm_seg[0, 1] > 0.5).astype(np.uint8)\n    \n    if binary_mask.sum() > 0:\n        coords = np.array(np.where(binary_mask > 0)).mean(axis=1)\n        location_idx = location_probs[0].argmax()\n        location_name = LOCATION_NAMES[location_idx]\n        confidence = location_probs[0].max()\n    else:\n        coords = np.array([img_size[0]//2, img_size[1]//2, img_size[2]//2])\n        location_idx = location_probs[0].argmax()\n        location_name = LOCATION_NAMES[location_idx]\n        confidence = location_probs[0].max()\n    \n    return {\n        'coordinates': coords,\n        'location': location_name,\n        'confidence': confidence,\n        'all_location_probs': {LOCATION_NAMES[i]: location_probs[0, i] for i in range(13)}\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.4615Z","iopub.execute_input":"2025-10-28T08:34:39.461807Z","iopub.status.idle":"2025-10-28T08:34:39.47507Z","shell.execute_reply.started":"2025-10-28T08:34:39.461789Z","shell.execute_reply":"2025-10-28T08:34:39.474242Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Subset Creation","metadata":{}},{"cell_type":"code","source":"def create_balanced_subset(train_df, localizers_df, segmentations_dir, n_samples):\n    \"\"\"Create balanced subset with preference for cases with segmentation\"\"\"\n    available_segs = [f.replace('_cowseg.nii', '') for f in os.listdir(segmentations_dir) if f.endswith('.nii')]\n    \n    train_with_seg = train_df[train_df['SeriesInstanceUID'].isin(available_segs)]\n    \n    pos = train_with_seg[train_with_seg['Aneurysm Present'] == 1]\n    neg = train_with_seg[train_with_seg['Aneurysm Present'] == 0]\n    \n    n_pos = min(n_samples // 2, len(pos))\n    n_neg = min(n_samples - n_pos, len(neg))\n    \n    if n_pos > 0:\n        pos_sample = pos.sample(n_pos, random_state=CFG.seed)\n    else:\n        pos_sample = pd.DataFrame()\n    \n    if n_neg > 0:\n        neg_sample = neg.sample(n_neg, random_state=CFG.seed)\n    else:\n        neg_sample = pd.DataFrame()\n    \n    subset = pd.concat([pos_sample, neg_sample]).reset_index(drop=True)\n    \n    return subset\n\n\nif CFG.use_subset:\n    train_subset = create_balanced_subset(\n        train_df, localizers_df, CFG.segmentations_dir, CFG.n_subset_samples\n    )\nelse:\n    train_subset = train_df\n\ntrain_split_df, temp_df = train_test_split(\n    train_subset, test_size=(1 - CFG.train_split), random_state=CFG.seed, stratify=train_subset['Aneurysm Present']\n)\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.5, random_state=CFG.seed, stratify=temp_df['Aneurysm Present']\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.475997Z","iopub.execute_input":"2025-10-28T08:34:39.476304Z","iopub.status.idle":"2025-10-28T08:34:39.502165Z","shell.execute_reply.started":"2025-10-28T08:34:39.476277Z","shell.execute_reply":"2025-10-28T08:34:39.501601Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualize Stage 1 Samples","metadata":{}},{"cell_type":"code","source":"for i in range(min(2, len(train_subset))):\n    series_uid = train_subset.iloc[i]['SeriesInstanceUID']\n    visualize_stage1_pipeline(series_uid, localizers_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:34:39.502975Z","iopub.execute_input":"2025-10-28T08:34:39.503245Z","iopub.status.idle":"2025-10-28T08:34:52.436088Z","shell.execute_reply.started":"2025-10-28T08:34:39.503227Z","shell.execute_reply":"2025-10-28T08:34:52.435017Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Create Datasets and Loaders","metadata":{}},{"cell_type":"code","source":"train_dataset = AneurysmDataset(train_split_df, localizers_df, CFG.img_size, augment=True)\nval_dataset = AneurysmDataset(val_df, localizers_df, CFG.img_size, augment=False)\ntest_dataset = AneurysmDataset(test_df, localizers_df, CFG.img_size, augment=False)\n\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=2)\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:49:33.894083Z","iopub.execute_input":"2025-10-28T08:49:33.894657Z","iopub.status.idle":"2025-10-28T08:49:33.901573Z","shell.execute_reply.started":"2025-10-28T08:49:33.894634Z","shell.execute_reply":"2025-10-28T08:49:33.900723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Initialize Model and Training","metadata":{}},{"cell_type":"code","source":"model = nnXNetLite(\n    in_channels=1,\n    num_vessel_classes=CFG.num_vessel_classes,\n    num_aneurysm_classes=CFG.num_aneurysm_classes,\n    num_location_classes=CFG.num_location_classes,\n    num_modality_classes=CFG.num_modality_classes,\n    base_channels=CFG.base_channels\n).to(CFG.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:49:45.776311Z","iopub.execute_input":"2025-10-28T08:49:45.776839Z","iopub.status.idle":"2025-10-28T08:49:45.858696Z","shell.execute_reply.started":"2025-10-28T08:49:45.776807Z","shell.execute_reply":"2025-10-28T08:49:45.85772Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:49:47.567434Z","iopub.execute_input":"2025-10-28T08:49:47.567859Z","iopub.status.idle":"2025-10-28T08:49:47.575833Z","shell.execute_reply.started":"2025-10-28T08:49:47.56783Z","shell.execute_reply":"2025-10-28T08:49:47.574968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = MultiTaskLoss(\n    vessel_seg_weight=CFG.vessel_seg_weight,\n    aneurysm_seg_weight=CFG.aneurysm_seg_weight,\n    binary_cls_weight=CFG.binary_cls_weight,\n    location_cls_weight=CFG.location_cls_weight,\n    coord_reg_weight=CFG.coord_reg_weight,\n    modality_cls_weight=CFG.modality_cls_weight\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:49:51.666565Z","iopub.execute_input":"2025-10-28T08:49:51.667061Z","iopub.status.idle":"2025-10-28T08:49:51.671322Z","shell.execute_reply.started":"2025-10-28T08:49:51.667038Z","shell.execute_reply":"2025-10-28T08:49:51.670423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = optim.Adam(model.parameters(), lr=CFG.learning_rate, weight_decay=CFG.weight_decay)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.num_epochs)\nscaler = GradScaler(enabled=CFG.mixed_precision)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:49:53.733316Z","iopub.execute_input":"2025-10-28T08:49:53.734068Z","iopub.status.idle":"2025-10-28T08:49:53.739107Z","shell.execute_reply.started":"2025-10-28T08:49:53.734043Z","shell.execute_reply":"2025-10-28T08:49:53.738215Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"history = {'train_loss': [], 'val_loss': [], 'val_auc': [], 'val_acc': []}\nbest_val_loss = float('inf')\npatience = 15\npatience_counter = 0\n\nfor epoch in range(CFG.num_epochs):\n    train_losses = train_epoch(model, train_loader, criterion, optimizer, scaler, CFG.device)\n    val_losses = validate(model, val_loader, criterion, CFG.device)\n    scheduler.step()\n    \n    history['train_loss'].append(train_losses['total'])\n    history['val_loss'].append(val_losses['total'])\n    history['val_auc'].append(val_losses['auc'])\n    history['val_acc'].append(val_losses['acc'])\n    \n    print(f\"Epoch {epoch+1}/{CFG.num_epochs}\")\n    print(f\"  Train Loss: {train_losses['total']:.4f} | Val Loss: {val_losses['total']:.4f}\")\n    print(f\"  Val AUC: {val_losses['auc']:.4f} | Val Acc: {val_losses['acc']:.4f}\")\n    \n    if val_losses['total'] < best_val_loss:\n        best_val_loss = val_losses['total']\n        torch.save(model.state_dict(), f'{CFG.work_dir}/best_model.pth')\n        patience_counter = 0\n    else:\n        patience_counter += 1\n    \n    if patience_counter >= patience:\n        print(f\"Early stopping at epoch {epoch+1}\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:49:58.831885Z","iopub.execute_input":"2025-10-28T08:49:58.832221Z","execution_failed":"2025-10-28T09:04:24.999Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Curves","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(history['train_loss'], label='Train')\naxes[0].plot(history['val_loss'], label='Val')\naxes[0].set_title('Loss')\naxes[0].legend()\naxes[0].grid(True)\n\naxes[1].plot(history['val_auc'], label='AUC', color='green')\naxes[1].set_title('Validation AUC')\naxes[1].legend()\naxes[1].grid(True)\n\naxes[2].plot(history['val_acc'], label='Accuracy', color='orange')\naxes[2].set_title('Validation Accuracy')\naxes[2].legend()\naxes[2].grid(True)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:38:41.140457Z","iopub.status.idle":"2025-10-28T08:38:41.141352Z","shell.execute_reply.started":"2025-10-28T08:38:41.141054Z","shell.execute_reply":"2025-10-28T08:38:41.141078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Test Evaluation","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load(f'{CFG.work_dir}/best_model.pth'))\ntest_losses = validate(model, test_loader, criterion, CFG.device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:38:41.142843Z","iopub.status.idle":"2025-10-28T08:38:41.143105Z","shell.execute_reply.started":"2025-10-28T08:38:41.142975Z","shell.execute_reply":"2025-10-28T08:38:41.142986Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Test Results:\")\nprint(f\"  Total Loss: {test_losses['total']:.4f}\")\nprint(f\"  Binary Loss: {test_losses['binary_cls']:.4f}\")\nprint(f\"  AUC: {test_losses['auc']:.4f}\")\nprint(f\"  Accuracy: {test_losses['acc']:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:38:41.144543Z","iopub.status.idle":"2025-10-28T08:38:41.144887Z","shell.execute_reply.started":"2025-10-28T08:38:41.144684Z","shell.execute_reply":"2025-10-28T08:38:41.1447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## TTA Inference Example","metadata":{}},{"cell_type":"code","source":"if CFG.use_tta:\n    sample_idx = 0\n    sample_batch = test_dataset[sample_idx]\n    volume = sample_batch['image'].unsqueeze(0)\n    \n    tta_pred = predict_with_tta(model, volume, CFG.device, n_flips=CFG.tta_n_flips)\n    \n    aneurysm_info = extract_aneurysm_info(\n        tta_pred['aneurysm_seg'],\n        tta_pred['location_cls'],\n        CFG.img_size\n    )\n    \n    print(\"\\nTTA Prediction Results:\")\n    print(f\"  Aneurysm Present: {tta_pred['binary_cls'][0][0]:.4f}\")\n    print(f\"  Location: {aneurysm_info['location']}\")\n    print(f\"  Confidence: {aneurysm_info['confidence']:.4f}\")\n    print(f\"  Coordinates: {aneurysm_info['coordinates']}\")\n    print(f\"\\n  All Location Probabilities:\")\n    for loc, prob in aneurysm_info['all_location_probs'].items():\n        print(f\"    {loc}: {prob:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:38:41.145925Z","iopub.status.idle":"2025-10-28T08:38:41.146228Z","shell.execute_reply.started":"2025-10-28T08:38:41.146106Z","shell.execute_reply":"2025-10-28T08:38:41.14612Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Visualization of Predictions","metadata":{}},{"cell_type":"code","source":"def visualize_with_ground_truth(model, dataset, idx, device, localizers_df):\n    \"\"\"Visualize with explicit ground truth location\"\"\"\n    model.eval()\n    batch = dataset[idx]\n    series_uid = batch['series_uid']\n    \n    loc_gt = localizers_df[localizers_df['SeriesInstanceUID'] == series_uid]\n    \n    print(f\"\\n=== Ground Truth Info ===\")\n    print(f\"Series: {series_uid}\")\n    print(f\"Binary label (has aneurysm): {batch['binary_label'].item()}\")\n    print(f\"GT mask sum: {batch['aneurysm_mask'].sum().item()}\")\n    print(f\"Localizers entries: {len(loc_gt)}\")\n    \n    if len(loc_gt) > 0:\n        print(\"\\nLocalizer data:\")\n        for _, row in loc_gt.iterrows():\n            print(f\"  Location: {row['location']}\")\n            print(f\"  Coordinates: {row['coordinates']}\")\n    else:\n        print(\"  No localizer data (negative case)\")\n    \n    image = batch['image'].unsqueeze(0).to(device)\n    with torch.no_grad():\n        outputs = model(image)\n    \n    volume = batch['image'].squeeze().cpu().numpy()\n    aneurysm_pred = torch.softmax(outputs['aneurysm_seg'], dim=1)[0, 1].cpu().numpy()\n    aneurysm_gt = batch['aneurysm_mask'].cpu().numpy()\n    \n    binary_pred = torch.sigmoid(outputs['binary_cls']).item()\n    location_probs = torch.sigmoid(outputs['location_cls']).cpu().numpy()[0]\n    location_pred = LOCATION_NAMES[location_probs.argmax()]\n    coord_pred = outputs['coord_reg'].cpu().numpy()[0]\n    \n    # Parse GT centroid\n    if aneurysm_gt.sum() > 0:\n        gt_coords = np.where(aneurysm_gt > 0)\n        gt_centroid = [gt_coords[0].mean(), gt_coords[1].mean(), gt_coords[2].mean()]\n        use_centroid = gt_centroid\n        has_gt = True\n    elif len(loc_gt) > 0:\n        try:\n            coords = eval(loc_gt.iloc[0]['coordinates'])\n            # Handle dict format {'x': ..., 'y': ..., 'z': ...} or list/tuple\n            if isinstance(coords, dict):\n                x = coords.get('x', coords.get('y', volume.shape[0] / 2))\n                y = coords.get('y', coords.get('x', volume.shape[1] / 2))\n                z = coords.get('z', volume.shape[2] / 2)  # Default to middle if no z\n            else:\n                x = coords[0] if len(coords) > 0 else volume.shape[0] / 2\n                y = coords[1] if len(coords) > 1 else volume.shape[1] / 2\n                z = coords[2] if len(coords) > 2 else volume.shape[2] / 2\n            \n            # Scale to current volume\n            gt_centroid = [\n                x * volume.shape[0] / 512,\n                y * volume.shape[1] / 512,\n                z * volume.shape[2] / 512\n            ]\n            use_centroid = gt_centroid\n            has_gt = True\n        except:\n            use_centroid = coord_pred\n            has_gt = False\n    else:\n        use_centroid = coord_pred\n        has_gt = False\n    \n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    slices = [\n        max(0, int(use_centroid[2]) - 16),\n        int(use_centroid[2]),\n        min(volume.shape[2] - 1, int(use_centroid[2]) + 16)\n    ]\n    \n    for i, s in enumerate(slices):\n        axes[0, i].imshow(volume[:, :, s], cmap='gray')\n        axes[0, i].set_title(f'Volume {s}')\n        axes[0, i].axis('off')\n        \n        axes[1, i].imshow(volume[:, :, s], cmap='gray', alpha=0.7)\n        \n        pred_slice = aneurysm_pred[:, :, s]\n        if pred_slice.max() > 0.01:\n            axes[1, i].imshow(pred_slice, cmap='Reds', alpha=0.4, vmin=0, vmax=0.5)\n        \n        if aneurysm_gt[:, :, s].sum() > 0:\n            axes[1, i].contour(aneurysm_gt[:, :, s], colors='lime', linewidths=3, levels=[0.5])\n        \n        if s == int(use_centroid[2]):\n            if has_gt:\n                axes[1, i].plot(use_centroid[1], use_centroid[0], 'g*', markersize=25, \n                               markeredgewidth=2, markeredgecolor='white', label='GT')\n            axes[1, i].plot(coord_pred[1], coord_pred[0], 'r*', markersize=20, \n                           markeredgewidth=2, markeredgecolor='yellow', label='Pred')\n            axes[1, i].legend(loc='upper right')\n        \n        axes[1, i].set_title(f'Pred {s} (max={pred_slice.max():.3f})')\n        axes[1, i].axis('off')\n    \n    title = f'GT: {batch[\"binary_label\"].item():.0f} | Pred: {binary_pred:.3f} | {location_pred[:30]}'\n    plt.suptitle(title, fontsize=11)\n    plt.tight_layout()\n    plt.show()\n    \n    print(f\"\\n=== Prediction ===\")\n    print(f\"Binary: {binary_pred:.4f}\")\n    print(f\"Location: {location_pred} (conf: {location_probs.max():.4f})\")\n    print(f\"Pred coords: {coord_pred}\")\n    if has_gt:\n        print(f\"GT coords: {use_centroid}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:38:41.147891Z","iopub.status.idle":"2025-10-28T08:38:41.14817Z","shell.execute_reply.started":"2025-10-28T08:38:41.14803Z","shell.execute_reply":"2025-10-28T08:38:41.148041Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"- Estrella VERDE: GT real (ground truth)\n- Estrella ROJA: Predicción del modelo\n- Contorno VERDE: GT mask si existe\n- Heatmap ROJO: Predicción del modelo\n- Print: Info de localizers + coordenadas reales","metadata":{}},{"cell_type":"code","source":"for i in range(len(test_dataset)):\n    visualize_with_ground_truth(model, test_dataset, i, CFG.device, localizers_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-28T08:38:41.148923Z","iopub.status.idle":"2025-10-28T08:38:41.149145Z","shell.execute_reply.started":"2025-10-28T08:38:41.149034Z","shell.execute_reply":"2025-10-28T08:38:41.149045Z"}},"outputs":[],"execution_count":null}]}