{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13851420},{"sourceType":"datasetVersion","sourceId":13258161,"datasetId":8401368,"databundleVersionId":13956035},{"sourceType":"datasetVersion","sourceId":13266765,"datasetId":8407116,"databundleVersionId":13965553}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA Aneurysm Detection - V3 Training (Attention-Guided Segmentation)\n\n## 🚀 Major Improvements in V3:\n1. **Attention-guided training** using segmentation masks\n2. **Dual-loss optimization** (classification + attention)\n3. **Segmentation-aware model** that learns WHERE to look\n4. **Proper train/val/test split** (70/20/10) with positive + negative samples\n5. **5-fold cross-validation** on training set\n6. **Mixed precision training** for speed\n\n## 🎯 How Attention-Guided Training Works:\n- For samples WITH segmentation (117): Model learns to focus on aneurysm regions\n- For samples WITHOUT segmentation (2,883): Regular classification training\n- During inference: Model has learned attention patterns, applies to ALL samples\n\n## 📊 Expected Performance:\n- V1 (baseline): ~65-70% accuracy\n- V2 (no segmentation): ~70% accuracy\n- **V3 (with attention): ~73-75% accuracy (+3-5% gain)**","metadata":{}},{"cell_type":"markdown","source":"## 1. Install Dependencies","metadata":{}},{"cell_type":"code","source":"%%capture\n!pip install timm albumentations nibabel","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.443Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Import Libraries","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport os\nimport glob\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\nfrom tqdm.notebook import tqdm\nimport warnings\nimport gc\nimport nibabel as nib\nfrom collections import Counter\nwarnings.filterwarnings('ignore')\n\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.445Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Configuration","metadata":{}},{"cell_type":"code","source":"class Config:\n    # Model settings\n    model_name = \"tf_efficientnetv2_s.in21k_ft_in1k\"\n    size = 384\n    in_chans = 32\n    \n    # Training settings\n    n_fold = 5\n    trn_fold = [0, 1, 2, 3, 4]\n    epochs = 10  # Increased for attention learning\n    batch_size = 6  # Reduced due to attention maps\n    lr = 2e-4\n    weight_decay = 1e-6\n    \n    # Attention-specific settings\n    use_attention = True\n    attention_loss_weight = 0.3  # Weight for attention loss\n    classification_loss_weight = 0.7  # Weight for classification loss\n    \n    # Data paths - UPDATE THESE\n    data_dirs = [\n        \"/kaggle/input/datasets/khanramshaayub/rsna-preprocessed-images\",\n        \"/kaggle/input/datasets/khanramshaayub/rsna-processed-images-2\"\n    ]\n    csv_path = \"/kaggle/input/competitions/rsna-intracranial-aneurysm-detection/train_localizers.csv\"\n    segmentation_dir = \"/kaggle/input/competitions/rsna-intracranial-aneurysm-detection/segmentations\"  \n    \n    # Data split ratios\n    train_ratio = 0.70\n    val_ratio = 0.20\n    test_ratio = 0.10\n    \n    # Output settings\n    output_dir = \"/kaggle/working\"\n    \n    # Target columns\n    target_cols = [\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 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        'Aneurysm Present',\n    ]\n    num_classes = len(target_cols)\n    \n    # Device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Debug mode\n    debug = False\n    debug_samples = 100\n\nCFG = Config()\nos.makedirs(CFG.output_dir, exist_ok=True)\n\nprint(f\"Device: {CFG.device}\")\nprint(f\"Model: {CFG.model_name}\")\nprint(f\"Image size: {CFG.size}x{CFG.size}\")\nprint(f\"Input channels: {CFG.in_chans}\")\nprint(f\"Output classes: {CFG.num_classes}\")\nprint(f\"\\n🎯 V3 Features:\")\nprint(f\"  - Attention-guided training: {CFG.use_attention}\")\nprint(f\"  - Attention loss weight: {CFG.attention_loss_weight}\")\nprint(f\"  - Classification loss weight: {CFG.classification_loss_weight}\")\nprint(f\"  - Segmentation directory: {CFG.segmentation_dir}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.445Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Build Label Map + Load Segmentations","metadata":{}},{"cell_type":"code","source":"def build_complete_label_map(csv_path, all_series_ids):\n    \"\"\"\n    Build label map for ALL series (positive AND negative)\n    Same as V2\n    \"\"\"\n    print(\"=\"*60)\n    print(\"Building Complete Label Map\")\n    print(\"=\"*60)\n    \n    df = pd.read_csv(csv_path)\n    print(f\"\\nCSV shape: {df.shape}\")\n    \n    label_map = {}\n    positive_series = set()\n    \n    # Positive samples from CSV\n    for sid in df['SeriesInstanceUID'].unique():\n        labels = np.zeros(CFG.num_classes, dtype=np.float32)\n        series_data = df[df['SeriesInstanceUID'] == sid]\n        \n        for _, row in series_data.iterrows():\n            location = row['location']\n            if location in CFG.target_cols[:-1]:\n                idx = CFG.target_cols.index(location)\n                labels[idx] = 1.0\n        \n        if labels[:-1].sum() > 0:\n            labels[-1] = 1.0\n        \n        label_map[sid] = labels\n        positive_series.add(sid)\n    \n    print(f\"\\n✅ POSITIVE series: {len(positive_series)}\")\n    \n    # Negative samples\n    negative_series = all_series_ids - positive_series\n    for sid in negative_series:\n        labels = np.zeros(CFG.num_classes, dtype=np.float32)\n        label_map[sid] = labels\n    \n    print(f\"✅ NEGATIVE series: {len(negative_series)}\")\n    print(f\"✅ TOTAL: {len(label_map)}\")\n    \n    return label_map, positive_series, negative_series\n\ndef load_segmentation_map(seg_dir):\n    \"\"\"\n    NEW in V3: Load all segmentation files\n    Returns: {series_id: segmentation_file_path}\n    \"\"\"\n    print(\"\\n\" + \"=\"*60)\n    print(\"Loading Segmentation Files\")\n    print(\"=\"*60)\n    \n    if not os.path.exists(seg_dir):\n        print(f\"⚠️  Segmentation directory not found: {seg_dir}\")\n        print(\"   Training will proceed WITHOUT attention guidance\")\n        return {}\n    \n    seg_files = glob.glob(f\"{seg_dir}/*.nii\")\n    seg_map = {}\n    \n    for seg_file in seg_files:\n        series_id = os.path.basename(seg_file).replace('.nii', '')\n        seg_map[series_id] = seg_file\n    \n    print(f\"\\n✅ Found {len(seg_map)} segmentation files\")\n    return seg_map","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.445Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Gather Files and Build Maps","metadata":{}},{"cell_type":"code","source":"# Gather all image files\nprint(\"=\"*60)\nprint(\"Gathering Image Files\")\nprint(\"=\"*60)\n\nall_files = []\nfor data_dir in CFG.data_dirs:\n    if os.path.exists(data_dir):\n        files = glob.glob(os.path.join(data_dir, \"*.npy\"))\n        all_files.extend(files)\n        print(f\"\\n{data_dir}: {len(files)} files\")\n\nprint(f\"\\nTOTAL FILES: {len(all_files)}\")\n\n# Extract series IDs\nall_series_ids = set([os.path.basename(f).replace('.npy', '') for f in all_files])\nseries_to_file = {os.path.basename(f).replace('.npy', ''): f for f in all_files}\n\n# Build label map\nlabel_map, positive_series, negative_series = build_complete_label_map(\n    CFG.csv_path, all_series_ids\n)\n\n# Load segmentation map\nseg_map = load_segmentation_map(CFG.segmentation_dir)\n\n# Find overlap\nseries_with_seg = set(seg_map.keys())\nseries_with_both = all_series_ids & series_with_seg\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"Data Summary\")\nprint(\"=\"*60)\nprint(f\"Images: {len(all_series_ids)}\")\nprint(f\"Segmentations: {len(seg_map)}\")\nprint(f\"Images WITH segmentation: {len(series_with_both)} ({len(series_with_both)/len(all_series_ids)*100:.1f}%)\")\nprint(f\"Images WITHOUT segmentation: {len(all_series_ids - series_with_seg)}\")\nprint(\"\\n💡 Attention guidance will be used for {:.1f}% of data\".format(\n    len(series_with_both)/len(all_series_ids)*100\n))","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Create Train/Val/Test Split","metadata":{}},{"cell_type":"code","source":"# Same stratified split as V2\ndef create_stratified_split(series_ids, label_map, train_ratio=0.7, val_ratio=0.2, test_ratio=0.1, random_state=42):\n    print(\"\\n\" + \"=\"*60)\n    print(\"Creating Stratified Train/Val/Test Split\")\n    print(\"=\"*60)\n    \n    series_list = list(series_ids)\n    y_stratify = np.array([label_map[sid][-1] for sid in series_list])\n    \n    # Split: test set first\n    train_val_series, test_series, train_val_labels, test_labels = train_test_split(\n        series_list, y_stratify,\n        test_size=test_ratio,\n        stratify=y_stratify,\n        random_state=random_state\n    )\n    \n    # Split: train and val\n    adjusted_val_ratio = val_ratio / (train_ratio + val_ratio)\n    train_series, val_series, train_labels, val_labels = train_test_split(\n        train_val_series, train_val_labels,\n        test_size=adjusted_val_ratio,\n        stratify=train_val_labels,\n        random_state=random_state\n    )\n    \n    def print_split_info(name, series, labels):\n        n_total = len(series)\n        n_pos = labels.sum()\n        n_neg = n_total - n_pos\n        print(f\"\\n{name}:\")\n        print(f\"  Total: {n_total:5d}\")\n        print(f\"  Positive: {n_pos:5.0f} ({n_pos/n_total*100:5.2f}%)\")\n        print(f\"  Negative: {n_neg:5.0f} ({n_neg/n_total*100:5.2f}%)\")\n    \n    print_split_info(\"TRAIN\", train_series, train_labels)\n    print_split_info(\"VAL\", val_series, val_labels)\n    print_split_info(\"TEST\", test_series, test_labels)\n    \n    return train_series, val_series, test_series\n\ntrain_series, val_series, test_series = create_stratified_split(\n    all_series_ids, label_map,\n    train_ratio=CFG.train_ratio,\n    val_ratio=CFG.val_ratio,\n    test_ratio=CFG.test_ratio\n)\n\ntrain_files = [series_to_file[sid] for sid in train_series]\nval_files = [series_to_file[sid] for sid in val_series]\ntest_files = [series_to_file[sid] for sid in test_series]","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.448Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Segmentation Processing Functions\n\n### 🆕 NEW IN V3: Functions to extract attention maps from segmentations","metadata":{}},{"cell_type":"code","source":"def extract_aneurysm_labels(seg_data, size_threshold_min=100, size_threshold_max=100000):\n    \"\"\"\n    NEW in V3: Identify which labels in segmentation are aneurysms\n    \n    Args:\n        seg_data: Segmentation volume (D, H, W)\n        size_threshold_min: Minimum voxels for aneurysm\n        size_threshold_max: Maximum voxels for aneurysm\n    \n    Returns:\n        List of label IDs that are aneurysms\n    \"\"\"\n    unique_vals, counts = np.unique(seg_data, return_counts=True)\n    \n    # Aneurysms are small structures (100 - 100,000 voxels)\n    aneurysm_labels = [\n        val for val, count in zip(unique_vals, counts)\n        if size_threshold_min < count < size_threshold_max and val > 0\n    ]\n    \n    return aneurysm_labels\n\ndef create_attention_map_from_segmentation(seg_data, target_shape=(32, 384, 384)):\n    \"\"\"\n    NEW in V3: Convert 3D segmentation to attention map\n    \n    Args:\n        seg_data: Segmentation volume (D, H, W) - e.g., (296, 512, 512)\n        target_shape: Target shape (32, 384, 384) to match model input\n    \n    Returns:\n        Attention map (32, 384, 384) with values 0-1\n    \"\"\"\n    # Step 1: Extract aneurysm labels\n    aneurysm_labels = extract_aneurysm_labels(seg_data)\n    \n    if len(aneurysm_labels) == 0:\n        # No aneurysm found - return zeros\n        return np.zeros(target_shape, dtype=np.float32)\n    \n    # Step 2: Create binary aneurysm mask\n    aneurysm_mask = np.isin(seg_data, aneurysm_labels).astype(np.float32)\n    \n    # Step 3: Resize to target shape\n    depth_orig, height_orig, width_orig = seg_data.shape\n    target_depth, target_height, target_width = target_shape\n    \n    # Select slices (depth dimension)\n    if depth_orig > target_depth:\n        indices = np.linspace(0, depth_orig - 1, target_depth).astype(int)\n        aneurysm_mask = aneurysm_mask[indices]\n    elif depth_orig < target_depth:\n        pad_size = target_depth - depth_orig\n        aneurysm_mask = np.pad(aneurysm_mask, ((0, pad_size), (0, 0), (0, 0)), mode='edge')\n    \n    # Resize spatial dimensions (height, width)\n    attention_map = []\n    for i in range(target_depth):\n        slice_2d = aneurysm_mask[i]\n        slice_resized = cv2.resize(slice_2d, (target_width, target_height))\n        attention_map.append(slice_resized)\n    \n    attention_map = np.array(attention_map, dtype=np.float32)\n    \n    # Step 4: Dilate attention map (expand focus region)\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (15, 15))\n    for i in range(target_depth):\n        if attention_map[i].max() > 0:\n            attention_map[i] = cv2.dilate(attention_map[i], kernel, iterations=1)\n    \n    # Step 5: Normalize to 0-1\n    if attention_map.max() > 0:\n        attention_map = attention_map / attention_map.max()\n    \n    return attention_map\n\nprint(\"✅ Segmentation processing functions loaded\")\nprint(\"   - extract_aneurysm_labels: Identifies aneurysm structures\")\nprint(\"   - create_attention_map_from_segmentation: Converts seg → attention map\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.448Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Dataset Class with Attention Maps\n\n### 🆕 NEW IN V3: Dataset now returns attention maps alongside images","metadata":{}},{"cell_type":"code","source":"class AneurysmDataset3D_WithAttention(Dataset):\n    \"\"\"\n    V3 Dataset: Returns (image, labels, attention_map, has_attention)\n    \n    NEW: Loads segmentation and creates attention maps\n    \"\"\"\n    def __init__(self, file_paths, label_map, seg_map, transform=None, is_training=True):\n        self.files = file_paths\n        self.label_map = label_map\n        self.seg_map = seg_map\n        self.transform = transform\n        self.is_training = is_training\n        \n        # Count how many samples have segmentation\n        self.n_with_seg = sum(1 for f in self.files \n                              if os.path.basename(f).replace('.npy', '') in self.seg_map)\n        \n        print(f\"Dataset initialized with {len(self.files)} files\")\n        print(f\"  With segmentation: {self.n_with_seg} ({self.n_with_seg/len(self.files)*100:.1f}%)\")\n        print(f\"  Without segmentation: {len(self.files) - self.n_with_seg}\")\n    \n    def __len__(self):\n        return len(self.files)\n    \n    def __getitem__(self, idx):\n        path = self.files[idx]\n        sid = os.path.basename(path).replace(\".npy\", \"\")\n        labels = self.label_map[sid]\n        \n        # Load image\n        img = np.load(path)  # (512, 512, 3)\n        \n        # NEW: Load segmentation and create attention map if available\n        has_attention = False\n        attention_map = np.zeros((CFG.in_chans, CFG.size, CFG.size), dtype=np.float32)\n        \n        if sid in self.seg_map:\n            try:\n                # Load segmentation\n                seg_path = self.seg_map[sid]\n                seg = nib.load(seg_path).get_fdata()\n                \n                # Create attention map\n                attention_map = create_attention_map_from_segmentation(\n                    seg, target_shape=(CFG.in_chans, CFG.size, CFG.size)\n                )\n                \n                has_attention = True\n            except Exception as e:\n                # If error loading segmentation, use zero attention\n                pass\n        \n        # Convert image to 32 channels\n        img = self.convert_to_32_channels(img)\n        \n        # Apply transforms to IMAGE only (not attention map)\n        if self.transform:\n            # Transform image\n            augmented = self.transform(image=img)\n            img = augmented['image']\n            \n            # Convert attention map to tensor manually\n            attention_map = torch.from_numpy(attention_map).float()\n        else:\n            img = img.transpose(2, 0, 1)\n            img = torch.from_numpy(img).float()\n            if img.max() > 1.0:\n                img /= 255.0\n            attention_map = torch.from_numpy(attention_map).float()\n        \n        # Return: image, labels, attention_map, has_attention\n        return img, torch.from_numpy(labels).float(), attention_map, torch.tensor(has_attention).float()\n    \n    def convert_to_32_channels(self, img):\n        \"\"\"Convert (512, 512, 3) to (384, 384, 32)\"\"\"\n        img = cv2.resize(img, (CFG.size, CFG.size))\n        \n        if img.max() > 1.0:\n            img = img.astype(np.float32) / 255.0\n        \n        channels = []\n        for i in range(CFG.in_chans):\n            ratio = (i / (CFG.in_chans - 1)) * 2\n            if ratio < 1:\n                ch = img[:, :, 0] * (1 - ratio) + img[:, :, 1] * ratio\n            else:\n                ratio = ratio - 1\n                ch = img[:, :, 1] * (1 - ratio) + img[:, :, 2] * ratio\n            channels.append(ch)\n        \n        volume = np.stack(channels, axis=-1).astype(np.float32)\n        return volume","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.45Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Data Augmentation","metadata":{}},{"cell_type":"code","source":"def get_train_transform():\n    return A.Compose([\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5),\n        A.Normalize(mean=[0.485] * CFG.in_chans, std=[0.229] * CFG.in_chans),\n        ToTensorV2(),\n    ])\n\ndef get_valid_transform():\n    return A.Compose([\n        A.Normalize(mean=[0.485] * CFG.in_chans, std=[0.229] * CFG.in_chans),\n        ToTensorV2(),\n    ])","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.45Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Model with Attention Output\n\n### 🆕 NEW IN V3: Model outputs both classification AND attention maps","metadata":{}},{"cell_type":"code","source":"class AneurysmClassifierWithAttention(nn.Module):\n    \"\"\"\n    V3 Model: EfficientNetV2 with attention head\n    \n    NEW: Additional attention head that learns to focus on aneurysm regions\n    \"\"\"\n    def __init__(self, model_name=CFG.model_name, num_classes=CFG.num_classes, \n                 in_chans=CFG.in_chans, pretrained=True):\n        super().__init__()\n        \n        # Backbone: EfficientNetV2 (same as V2)\n        self.backbone = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            features_only=True,  # NEW: Get intermediate features\n            in_chans=in_chans\n        )\n        \n        # Get number of output features from last layer\n        with torch.no_grad():\n            dummy_input = torch.randn(1, in_chans, CFG.size, CFG.size)\n            features = self.backbone(dummy_input)\n            self.n_features = features[-1].shape[1]  # Last feature map channels\n            self.feature_size = features[-1].shape[2]  # Spatial size\n        \n        # Classification head (same as V2)\n        self.classifier = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(self.n_features, num_classes)\n        )\n        \n        # NEW: Attention head\n        # Takes feature maps and outputs attention map (same size as input)\n        self.attention_head = nn.Sequential(\n            nn.Conv2d(self.n_features, 256, kernel_size=3, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(256, 128, kernel_size=3, padding=1),\n            nn.BatchNorm2d(128),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(128, in_chans, kernel_size=1),  # Output: (B, 32, H, W)\n            nn.Sigmoid()  # Attention values 0-1\n        )\n        \n    def forward(self, x, return_attention=False):\n        \"\"\"\n        Args:\n            x: Input images (B, 32, 384, 384)\n            return_attention: Whether to return attention map\n        \n        Returns:\n            If return_attention=False: classification outputs (B, 14)\n            If return_attention=True: (classification, attention_map)\n        \"\"\"\n        # Extract features\n        features = self.backbone(x)\n        last_features = features[-1]  # (B, n_features, H_small, W_small)\n        \n        # Classification\n        cls_output = self.classifier(last_features)\n        \n        if return_attention:\n            # Generate attention map\n            attention = self.attention_head(last_features)  # (B, 32, H_small, W_small)\n            \n            # Upsample to match input size\n            attention = F.interpolate(\n                attention, \n                size=(CFG.size, CFG.size),\n                mode='bilinear', \n                align_corners=False\n            )\n            \n            return cls_output, attention\n        \n        return cls_output\n\n# Test model\nprint(\"\\nCreating V3 model with attention head...\")\ntest_model = AneurysmClassifierWithAttention(pretrained=True).to(CFG.device)\nprint(f\"✅ Model created!\")\nprint(f\"Parameters: {sum(p.numel() for p in test_model.parameters()):,}\")\nprint(f\"Backbone features: {test_model.n_features}\")\nprint(f\"Feature spatial size: {test_model.feature_size}x{test_model.feature_size}\")\n\n# Test forward pass\ndummy_input = torch.randn(2, CFG.in_chans, CFG.size, CFG.size).to(CFG.device)\ncls_out, att_out = test_model(dummy_input, return_attention=True)\nprint(f\"\\nTest forward pass:\")\nprint(f\"  Input shape: {dummy_input.shape}\")\nprint(f\"  Classification output: {cls_out.shape}\")\nprint(f\"  Attention output: {att_out.shape}\")\n\ndel test_model, dummy_input, cls_out, att_out\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.451Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Dual Loss Function\n\n### 🆕 NEW IN V3: Combined classification + attention loss","metadata":{}},{"cell_type":"code","source":"class DualLoss(nn.Module):\n    \"\"\"\n    V3 Loss: Combines classification loss + attention loss\n    \n    Total Loss = α × Classification Loss + β × Attention Loss\n    \n    Only applies attention loss to samples that have segmentation\n    \"\"\"\n    def __init__(self, classification_weight=0.7, attention_weight=0.3):\n        super().__init__()\n        self.classification_weight = classification_weight\n        self.attention_weight = attention_weight\n        \n        # Classification loss (same as V2)\n        self.cls_criterion = nn.BCEWithLogitsLoss()\n        \n        # Attention loss (MSE between predicted and ground truth attention)\n        self.att_criterion = nn.MSELoss()\n    \n    def forward(self, cls_outputs, attention_outputs, labels, \n                attention_targets, has_attention):\n        \"\"\"\n        Args:\n            cls_outputs: Classification predictions (B, 14)\n            attention_outputs: Attention maps (B, 32, 384, 384)\n            labels: Ground truth labels (B, 14)\n            attention_targets: Ground truth attention (B, 32, 384, 384)\n            has_attention: Mask indicating which samples have segmentation (B,)\n        \n        Returns:\n            total_loss, cls_loss, att_loss\n        \"\"\"\n        # Classification loss (apply to ALL samples)\n        cls_loss = self.cls_criterion(cls_outputs, labels)\n        \n        # Attention loss (apply ONLY to samples with segmentation)\n        if has_attention.sum() > 0:\n            # Mask to select samples with segmentation\n            mask = has_attention.bool()\n            \n            # Calculate attention loss only for these samples\n            att_loss = self.att_criterion(\n                attention_outputs[mask],\n                attention_targets[mask]\n            )\n        else:\n            # No samples with segmentation in this batch\n            att_loss = torch.tensor(0.0, device=cls_outputs.device)\n        \n        # Combined loss\n        total_loss = (self.classification_weight * cls_loss + \n                     self.attention_weight * att_loss)\n        \n        return total_loss, cls_loss, att_loss\n\nprint(\"✅ Dual Loss initialized\")\nprint(f\"   Classification weight: {CFG.classification_loss_weight}\")\nprint(f\"   Attention weight: {CFG.attention_loss_weight}\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.453Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Training Functions\n\n### 🆕 MODIFIED IN V3: Training loop now handles attention","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, train_loader, optimizer, criterion, scaler, device, epoch):\n    \"\"\"V3 Training: Now processes attention maps\"\"\"\n    model.train()\n    running_loss = 0.0\n    running_cls_loss = 0.0\n    running_att_loss = 0.0\n    n_with_attention = 0\n    \n    pbar = tqdm(train_loader, desc=f'Epoch {epoch+1} [TRAIN]')\n    for images, labels, attention_targets, has_attention in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n        attention_targets = attention_targets.to(device)\n        has_attention = has_attention.to(device)\n        \n        optimizer.zero_grad()\n        \n        with autocast():\n            # Forward pass with attention\n            cls_outputs, attention_outputs = model(images, return_attention=True)\n            \n            # Calculate dual loss\n            loss, cls_loss, att_loss = criterion(\n                cls_outputs, attention_outputs, labels, \n                attention_targets, has_attention\n            )\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        running_loss += loss.item()\n        running_cls_loss += cls_loss.item()\n        running_att_loss += att_loss.item()\n        n_with_attention += has_attention.sum().item()\n        \n        pbar.set_postfix({\n            'loss': running_loss / (pbar.n + 1),\n            'cls': running_cls_loss / (pbar.n + 1),\n            'att': running_att_loss / (pbar.n + 1)\n        })\n    \n    n_batches = len(train_loader)\n    print(f\"  Samples with attention guidance: {n_with_attention}/{len(train_loader.dataset)}\")\n    \n    return (running_loss / n_batches, \n            running_cls_loss / n_batches, \n            running_att_loss / n_batches)\n\n@torch.no_grad()\ndef validate(model, valid_loader, criterion, device, epoch):\n    \"\"\"V3 Validation: Now evaluates attention accuracy\"\"\"\n    model.eval()\n    running_loss = 0.0\n    running_cls_loss = 0.0\n    running_att_loss = 0.0\n    all_preds = []\n    all_labels = []\n    n_with_attention = 0\n    \n    pbar = tqdm(valid_loader, desc=f'Epoch {epoch+1} [VALID]')\n    for images, labels, attention_targets, has_attention in pbar:\n        images = images.to(device)\n        labels = labels.to(device)\n        attention_targets = attention_targets.to(device)\n        has_attention = has_attention.to(device)\n        \n        with autocast():\n            cls_outputs, attention_outputs = model(images, return_attention=True)\n            loss, cls_loss, att_loss = criterion(\n                cls_outputs, attention_outputs, labels,\n                attention_targets, has_attention\n            )\n        \n        running_loss += loss.item()\n        running_cls_loss += cls_loss.item()\n        running_att_loss += att_loss.item()\n        n_with_attention += has_attention.sum().item()\n        \n        probs = torch.sigmoid(cls_outputs).cpu().numpy()\n        all_preds.append(probs)\n        all_labels.append(labels.cpu().numpy())\n        \n        pbar.set_postfix({\n            'loss': running_loss / (pbar.n + 1),\n            'cls': running_cls_loss / (pbar.n + 1),\n            'att': running_att_loss / (pbar.n + 1)\n        })\n    \n    all_preds = np.concatenate(all_preds, axis=0)\n    all_labels = np.concatenate(all_labels, axis=0)\n    \n    pred_binary = (all_preds > 0.5).astype(int)\n    accuracy = (pred_binary == all_labels).mean()\n    \n    n_batches = len(valid_loader)\n    return (running_loss / n_batches,\n            running_cls_loss / n_batches,\n            running_att_loss / n_batches,\n            accuracy, all_preds, all_labels)","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.454Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13. K-Fold Training Loop","metadata":{}},{"cell_type":"code","source":"def train_fold(fold, fold_train_files, fold_val_files, label_map, seg_map):\n    print(f\"\\n{'='*60}\")\n    print(f\"Training Fold {fold} (V3 - Attention-Guided)\")\n    print(f\"{'='*60}\")\n    \n    # Create datasets\n    print(\"\\nCreating datasets...\")\n    train_dataset = AneurysmDataset3D_WithAttention(\n        fold_train_files, label_map, seg_map,\n        transform=get_train_transform(),\n        is_training=True\n    )\n    valid_dataset = AneurysmDataset3D_WithAttention(\n        fold_val_files, label_map, seg_map,\n        transform=get_valid_transform(),\n        is_training=False\n    )\n    \n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=CFG.batch_size,\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True,\n        drop_last=True\n    )\n    valid_loader = DataLoader(\n        valid_dataset,\n        batch_size=CFG.batch_size,\n        shuffle=False,\n        num_workers=2,\n        pin_memory=True\n    )\n    \n    print(f\"\\nDataLoader Info:\")\n    print(f\"  Train batches: {len(train_loader)}\")\n    print(f\"  Valid batches: {len(valid_loader)}\")\n    \n    # Create model\n    model = AneurysmClassifierWithAttention(pretrained=True).to(CFG.device)\n    \n    # Loss and optimizer\n    criterion = DualLoss(\n        classification_weight=CFG.classification_loss_weight,\n        attention_weight=CFG.attention_loss_weight\n    )\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=CFG.lr,\n        weight_decay=CFG.weight_decay\n    )\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=CFG.epochs, eta_min=1e-6\n    )\n    scaler = GradScaler()\n    \n    # Training loop\n    best_accuracy = 0.0\n    best_epoch = 0\n    \n    for epoch in range(CFG.epochs):\n        train_loss, train_cls_loss, train_att_loss = train_one_epoch(\n            model, train_loader, optimizer, criterion, scaler, CFG.device, epoch\n        )\n        \n        val_loss, val_cls_loss, val_att_loss, accuracy, preds, labels = validate(\n            model, valid_loader, criterion, CFG.device, epoch\n        )\n        \n        scheduler.step()\n        \n        print(f\"\\nEpoch {epoch+1}/{CFG.epochs}\")\n        print(f\"Train - Total: {train_loss:.4f} | Cls: {train_cls_loss:.4f} | Att: {train_att_loss:.4f}\")\n        print(f\"Valid - Total: {val_loss:.4f} | Cls: {val_cls_loss:.4f} | Att: {val_att_loss:.4f}\")\n        print(f\"Accuracy: {accuracy:.4f}\")\n        print(f\"LR: {optimizer.param_groups[0]['lr']:.6f}\")\n        \n        if accuracy > best_accuracy:\n            best_accuracy = accuracy\n            best_epoch = epoch\n            \n            checkpoint = {\n                'model': model.state_dict(),\n                'optimizer': optimizer.state_dict(),\n                'epoch': epoch,\n                'accuracy': accuracy,\n                'fold': fold,\n                'config': {\n                    'attention_loss_weight': CFG.attention_loss_weight,\n                    'classification_loss_weight': CFG.classification_loss_weight\n                }\n            }\n            \n            save_path = f\"{CFG.output_dir}/{CFG.model_name}_v3_fold{fold}_best.pth\"\n            torch.save(checkpoint, save_path)\n            print(f\"✅ Saved V3 model! Accuracy: {accuracy:.4f}\")\n    \n    print(f\"\\n🏆 Best Accuracy for Fold {fold}: {best_accuracy:.4f} at epoch {best_epoch+1}\")\n    \n    del model, train_loader, valid_loader, train_dataset, valid_dataset\n    gc.collect()\n    torch.cuda.empty_cache()\n    \n    return best_accuracy","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.455Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 14. Start K-Fold Training","metadata":{}},{"cell_type":"code","source":"# K-Fold Cross Validation on training set\ntrain_labels = np.array([label_map[os.path.basename(f).replace('.npy', '')][-1] for f in train_files])\n\nskf = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=42)\nfold_accuracies = []\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"Starting K-Fold Cross-Validation (V3 - Attention-Guided)\")\nprint(\"=\"*60)\n\nfor fold, (fold_train_idx, fold_val_idx) in enumerate(skf.split(train_files, train_labels)):\n    if fold not in CFG.trn_fold:\n        continue\n    \n    fold_train_files = [train_files[i] for i in fold_train_idx]\n    fold_val_files = [train_files[i] for i in fold_val_idx]\n    \n    print(f\"\\nFold {fold}:\")\n    print(f\"  Fold Train: {len(fold_train_files)} files\")\n    print(f\"  Fold Val:   {len(fold_val_files)} files\")\n    \n    fold_acc = train_fold(fold, fold_train_files, fold_val_files, label_map, seg_map)\n    fold_accuracies.append(fold_acc)\n    \n    if CFG.debug:\n        break","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.457Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 15. Training Summary","metadata":{}},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"V3 Training Complete!\")\nprint(\"=\"*60)\nprint(f\"Average Accuracy: {np.mean(fold_accuracies):.4f}\")\nprint(f\"Std Dev: {np.std(fold_accuracies):.4f}\")\nprint(\"\\nPer-fold results:\")\nfor i, acc in enumerate(fold_accuracies):\n    print(f\"Fold {CFG.trn_fold[i]}: {acc:.4f}\")\n\nprint(f\"\\n✅ V3 Models saved in: {CFG.output_dir}\")\nprint(f\"   Model naming: *_v3_fold*_best.pth\")\n\nsaved_models = glob.glob(f\"{CFG.output_dir}/*_v3_*.pth\")\nprint(f\"\\nSaved V3 models ({len(saved_models)}):\")\nfor model_path in saved_models:\n    size_mb = os.path.getsize(model_path) / (1024*1024)\n    print(f\"  - {os.path.basename(model_path)} ({size_mb:.1f} MB)\")\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"🎯 V3 Features Used:\")\nprint(\"=\"*60)\nprint(f\"✅ Attention-guided training with segmentation masks\")\nprint(f\"✅ Dual-loss optimization (classification + attention)\")\nprint(f\"✅ Model learns WHERE to look for aneurysms\")\nprint(f\"✅ Attention patterns transfer to samples without segmentation\")\nprint(f\"\\n📊 Expected improvement over V2: +3-5% accuracy\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-02-20T20:27:54.457Z"}},"outputs":[],"execution_count":null}]}