{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":13451,"datasetId":654585,"databundleVersionId":1188070}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# BrainHemNet: Automated Intracranial Hemorrhage Detection and Segmentation using DINOv3 and SAM2","metadata":{}},{"cell_type":"code","source":"# First, install all required packages\n!pip install torch torchvision monai scikit-learn opencv-python matplotlib tqdm pandas\n!pip install transformers  # For DINOv3\n!pip install git+https://github.com/facebookresearch/sam2.git  # Install SAM2 directly","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-07T16:04:29.65658Z","iopub.execute_input":"2025-09-07T16:04:29.656877Z","iopub.status.idle":"2025-09-07T16:10:02.638407Z","shell.execute_reply.started":"2025-09-07T16:04:29.656854Z","shell.execute_reply":"2025-09-07T16:10:02.637701Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import all necessary libraries\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nimport pydicom  # For DICOM file reading\nimport monai\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.decomposition import PCA\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Check for GPU availability\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n\n# Import DINOv3 from transformers\ntry:\n    from transformers import AutoImageProcessor, AutoModel\n    dinov3_available = True\n    print(\"DINOv3 transformers imported successfully\")\nexcept ImportError:\n    print(\"Transformers library not available.\")\n    dinov3_available = False\n\n# Import SAM2 \ntry:\n    from sam2.build_sam import build_sam2\n    from sam2.sam2_image_predictor import SAM2ImagePredictor\n    sam2_available = True\n    print(\"SAM2 imported successfully\")\nexcept ImportError as e:\n    print(f\"SAM2 import failed: {e}\")\n    sam2_available = False\n\n# Configuration\nclass Config:\n    # Data paths (update these based on your Kaggle dataset structure)\n    train_csv_path = \"/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection/stage_2_train.csv\"\n    train_image_dir = \"/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection/stage_2_train\"\n    test_image_dir = \"/kaggle/input/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection/stage_2_test\"\n    \n    # Model parameters\n    image_size = 512\n    dinov3_model_name = \"facebook/dinov2-base\"  # Using DINOv2 which works better with transformers\n    \n    # Training parameters\n    batch_size = 4  # Reduced for Kaggle memory constraints\n    learning_rate = 1e-4\n    num_epochs = 5   # Reduced for Kaggle time constraints\n    \n    # Anomaly detection parameters\n    anomaly_threshold = 0.7\n    num_components = 5  # Reduced for Kaggle memory constraints\n    \n    # Medical image specific parameters\n    window_center = 40    # Typical brain window center\n    window_width = 80     # Typical brain window width\n    use_windowing = True  # Apply medical image windowing\n\nconfig = Config()\n\n# Enhanced DICOM reading functions\ndef apply_windowing(image, window_center, window_width):\n    \"\"\"Apply medical image windowing to 16-bit DICOM images\"\"\"\n    window_min = window_center - window_width / 2\n    window_max = window_center + window_width / 2\n    \n    # Clip values to window range\n    image = np.clip(image, window_min, window_max)\n    \n    # Scale to 0-255\n    image = ((image - window_min) / (window_max - window_min + 1e-8) * 255).astype(np.uint8)\n    \n    return image\n\ndef read_dicom_image(dicom_path):\n    try:\n        # Read DICOM file\n        dicom = pydicom.dcmread(dicom_path)\n        \n        # Extract pixel array\n        img_array = dicom.pixel_array\n        \n        # Handle specific problematic format: (1, 1, 3), <i2\n        if img_array.shape == (1, 1, 3) and (img_array.dtype == np.int16 or img_array.dtype == np.uint16):\n            # This appears to be a single pixel with RGB values\n            # Extract the RGB values and create a solid color image\n            rgb_values = img_array[0, 0]\n            # Create a 256x256 image with the same color\n            img_array = np.full((256, 256, 3), rgb_values, dtype=np.uint8)\n            image = Image.fromarray(img_array)\n            return image\n        \n        # Handle other unusual shapes\n        if len(img_array.shape) == 3 and img_array.shape[2] == 3 and img_array.shape[0] == 1 and img_array.shape[1] == 1:\n            # Single pixel RGB image - expand to reasonable size\n            rgb_values = img_array[0, 0]\n            img_array = np.full((256, 256, 3), rgb_values, dtype=np.uint8)\n            image = Image.fromarray(img_array)\n            return image\n        \n        # Convert 16-bit to 8-bit if necessary with medical windowing\n        if img_array.dtype == np.uint16 or img_array.dtype == np.int16:\n            # Use windowing for medical images\n            window_center = getattr(dicom, 'WindowCenter', config.window_center)\n            window_width = getattr(dicom, 'WindowWidth', config.window_width)\n            \n            if hasattr(window_center, '__iter__') and not isinstance(window_center, str):\n                window_center = window_center[0] if len(window_center) > 0 else config.window_center\n            if hasattr(window_width, '__iter__') and not isinstance(window_width, str):\n                window_width = window_width[0] if len(window_width) > 0 else config.window_width\n            \n            # Apply windowing\n            img_array = apply_windowing(img_array.astype(np.float32), float(window_center), float(window_width))\n        \n        # Handle single channel images\n        if len(img_array.shape) == 2:\n            # Convert to 3-channel\n            img_array = np.stack([img_array] * 3, axis=-1)\n        elif len(img_array.shape) == 3 and img_array.shape[2] == 1:\n            # Remove singleton dimension and convert to 3-channel\n            img_array = np.repeat(img_array, 3, axis=2)\n        elif len(img_array.shape) == 3 and img_array.shape[2] > 3:\n            # Take first 3 channels if more than 3\n            img_array = img_array[:, :, :3]\n        \n        # Normalize to 0-255 if not already\n        if img_array.dtype != np.uint8:\n            img_min = img_array.min()\n            img_max = img_array.max()\n            if img_max > img_min:  # Avoid division by zero\n                img_array = ((img_array - img_min) / (img_max - img_min) * 255).astype(np.uint8)\n            else:\n                img_array = np.zeros_like(img_array, dtype=np.uint8)\n        \n        # Convert to PIL Image\n        image = Image.fromarray(img_array)\n        return image\n        \n    except Exception as e:\n        print(f\"Error reading DICOM file {dicom_path}: {e}\")\n        print(f\"Array shape: {getattr(img_array, 'shape', 'unknown')}\")\n        print(f\"Array dtype: {getattr(img_array, 'dtype', 'unknown')}\")\n        # Return a blank image if loading fails\n        return Image.new('RGB', (256, 256), color='black')\n\ndef analyze_dicom_file(dicom_path):\n    \"\"\"Analyze DICOM file metadata to understand its structure\"\"\"\n    try:\n        dicom = pydicom.dcmread(dicom_path)\n        print(f\"File: {dicom_path}\")\n        print(f\"Shape: {dicom.pixel_array.shape}\")\n        print(f\"Data type: {dicom.pixel_array.dtype}\")\n        print(f\"Photometric Interpretation: {getattr(dicom, 'PhotometricInterpretation', 'Unknown')}\")\n        print(f\"Samples per Pixel: {getattr(dicom, 'SamplesPerPixel', 'Unknown')}\")\n        print(f\"Bits Stored: {getattr(dicom, 'BitsStored', 'Unknown')}\")\n        print(f\"Window Center: {getattr(dicom, 'WindowCenter', 'Unknown')}\")\n        print(f\"Window Width: {getattr(dicom, 'WindowWidth', 'Unknown')}\")\n        print(\"-\" * 50)\n        return True\n    except Exception as e:\n        print(f\"Error analyzing {dicom_path}: {e}\")\n        return False\n\n# Enhanced transform for medical images\ndef make_medical_transform(resize_size=512):\n    import torchvision.transforms as transforms\n    \n    transform_list = [\n        transforms.ToTensor(),\n        transforms.Resize((resize_size, resize_size)),\n        # Medical images often benefit from different normalization\n        transforms.Normalize(\n            mean=[0.5, 0.5, 0.5],\n            std=[0.5, 0.5, 0.5],\n        )\n    ]\n    \n    return transforms.Compose(transform_list)\n\ntrain_transform = make_medical_transform(resize_size=config.image_size)\ntest_transform = make_medical_transform(resize_size=config.image_size)\n\n# Data preprocessing\nclass IntracranialHemorrhageDataset(Dataset):\n    def __init__(self, csv_file, image_dir, transform=None, is_test=False, sample_size=None):\n        self.image_dir = image_dir\n        self.transform = transform\n        self.is_test = is_test\n        self.valid_indices = []\n        \n        if not is_test:\n            self.data = pd.read_csv(csv_file)\n            # Extract labels and image IDs - FIXED: use proper string splitting\n            self.data['Image'] = self.data['ID'].apply(lambda x: x.rsplit('_', 1)[0])\n            self.data['Subtype'] = self.data['ID'].apply(lambda x: x.rsplit('_', 1)[1] if '_' in x else 'any')\n            \n            # FIXED: Handle duplicate entries by aggregating (taking max value)\n            self.labels_df = self.data.groupby(['Image', 'Subtype'])['Label'].max().unstack(fill_value=0)\n            self.image_ids = self.labels_df.index.tolist()\n            self.labels = self.labels_df.values\n            \n            # Limit sample size for Kaggle\n            if sample_size is not None and sample_size < len(self.image_ids):\n                self.image_ids = self.image_ids[:sample_size]\n                self.labels = self.labels[:sample_size]\n        else:\n            # Look for DICOM files instead of PNG\n            self.image_ids = [f.split('.')[0] for f in os.listdir(image_dir) \n                             if f.endswith('.dcm')]\n            # Limit sample size for Kaggle\n            if sample_size is not None and sample_size < len(self.image_ids):\n                self.image_ids = self.image_ids[:sample_size]\n        \n        # Pre-validate images and create valid indices\n        for idx in range(len(self.image_ids)):\n            img_id = self.image_ids[idx]\n            img_path = os.path.join(self.image_dir, f\"{img_id}.dcm\")\n            try:\n                # Test if we can read the image\n                test_image = read_dicom_image(img_path)\n                if test_image.size[0] > 0 and test_image.size[1] > 0:\n                    self.valid_indices.append(idx)\n                else:\n                    print(f\"Skipping invalid image: {img_path}\")\n            except Exception as e:\n                print(f\"Skipping problematic file {img_path}: {e}\")\n    \n    def __len__(self):\n        return len(self.valid_indices)\n    \n    def __getitem__(self, idx):\n        actual_idx = self.valid_indices[idx]\n        img_id = self.image_ids[actual_idx]\n        img_path = os.path.join(self.image_dir, f\"{img_id}.dcm\")\n        \n        # Load DICOM image\n        image = read_dicom_image(img_path)\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        if self.is_test:\n            return image, img_id\n        else:\n            labels = self.labels[actual_idx]\n            return image, torch.FloatTensor(labels)\n\n# Initialize models\ndef setup_models():\n    models = {}\n    \n    # Setup DINOv3 using transformers\n    if dinov3_available:\n        print(\"Loading DINOv3 model...\")\n        try:\n            dinov3_processor = AutoImageProcessor.from_pretrained(config.dinov3_model_name)\n            dinov3_model = AutoModel.from_pretrained(config.dinov3_model_name).to(device)\n            dinov3_model.eval()\n            models['dinov3'] = dinov3_model\n            models['dinov3_processor'] = dinov3_processor\n            print(\"DINOv3 model loaded successfully\")\n        except Exception as e:\n            print(f\"Error loading DINOv3 model: {e}\")\n            # Create a dummy model for demonstration\n            class DummyDINOv3(nn.Module):\n                def __init__(self):\n                    super().__init__()\n                    self.patch_embed = nn.Conv2d(3, 768, kernel_size=16, stride=16)\n                    \n                def forward(self, x):\n                    return self.patch_embed(x).flatten(2).transpose(1, 2)\n            \n            models['dinov3'] = DummyDINOv3().to(device)\n            models['dinov3_processor'] = None\n    else:\n        print(\"DINOv3 not available. Using random features for demonstration.\")\n        # Create a dummy model for demonstration\n        class DummyDINOv3(nn.Module):\n            def __init__(self):\n                super().__init__()\n                self.patch_embed = nn.Conv2d(3, 768, kernel_size=16, stride=16)\n                \n            def forward(self, x):\n                return self.patch_embed(x).flatten(2).transpose(1, 2)\n        \n        models['dinov3'] = DummyDINOv3().to(device)\n        models['dinov3_processor'] = None\n    \n    # Setup SAM2 - Simplified approach without config file\n    if sam2_available:\n        print(\"Loading SAM2 model...\")\n        try:\n            # Use Hugging Face integration instead of direct loading\n            from sam2.sam2_image_predictor import SAM2ImagePredictor\n            sam2_predictor = SAM2ImagePredictor.from_pretrained(\"facebook/sam2-hiera-tiny\")\n            models['sam2'] = sam2_predictor\n            print(\"SAM2 model loaded successfully from Hugging Face\")\n        except Exception as e:\n            print(f\"Error loading SAM2 model: {e}\")\n            models['sam2'] = None\n    else:\n        print(\"SAM2 not available. Using dummy segmentation for demonstration.\")\n        models['sam2'] = None\n    \n    return models\n\nprint(\"Setting up models...\")\nmodels = setup_models()\ndinov3_model = models['dinov3']\ndinov3_processor = models.get('dinov3_processor')\nsam2_predictor = models['sam2']\n\n# Create datasets with limited samples for Kaggle\ntrain_dataset = IntracranialHemorrhageDataset(\n    config.train_csv_path, config.train_image_dir, \n    transform=train_transform, sample_size=1000  # Limited for Kaggle\n)\n\nprint(f\"Found {len(train_dataset)} valid training images out of {len(train_dataset.image_ids)} total\")\n\ntrain_loader = DataLoader(\n    train_dataset, batch_size=config.batch_size, shuffle=True, num_workers=2\n)\n\n# Feature extraction with DINOv3\ndef extract_features(model, processor, dataloader):\n    model.eval()\n    features = []\n    labels_list = []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(dataloader, desc=\"Extracting features\"):\n            images = images.to(device)\n            \n            if processor is not None:\n                # Use processor if available (transformers approach)\n                try:\n                    inputs = processor(images=images, return_tensors=\"pt\").to(device)\n                    outputs = model(**inputs)\n                    features_batch = outputs.last_hidden_state\n                except:\n                    # Fallback if processor fails\n                    features_batch = model(images).last_hidden_state\n            else:\n                # Fallback for dummy model\n                features_batch = model(images)\n            \n            features.append(features_batch.cpu())\n            labels_list.append(labels)\n    \n    return torch.cat(features, dim=0), torch.cat(labels_list, dim=0)\n\nprint(\"Extracting features from training data...\")\ntrain_features, train_labels = extract_features(dinov3_model, dinov3_processor, train_loader)\n\n# Reshape features for anomaly detection\nif len(train_features.shape) == 3:\n    # Use [CLS] token or average pooling\n    train_features = train_features.mean(dim=1)  # Average pooling\n\n# Anomaly detection module\nclass AnomalyDetector:\n    def __init__(self, n_components=5):\n        self.n_components = n_components\n        self.pca = None\n        self.normal_features = None\n        self.normal_mean = None\n        self.normal_std = None\n        \n    def fit(self, features, labels):\n        # Use only normal samples for training\n        if len(labels.shape) > 1:\n            # Multi-label case\n            normal_idx = (labels.sum(axis=1) == 0).nonzero().squeeze()\n        else:\n            # Single label case\n            normal_idx = (labels == 0).nonzero().squeeze()\n            \n        if normal_idx.numel() == 0:\n            print(\"No normal samples found for training anomaly detector\")\n            # Use all samples as fallback\n            normal_idx = torch.arange(len(features))\n            \n        if len(normal_idx.shape) == 0:  # Handle single sample case\n            normal_idx = normal_idx.unsqueeze(0)\n            \n        self.normal_features = features[normal_idx].numpy()\n        \n        # Normalize features\n        self.normal_mean = np.mean(self.normal_features, axis=0)\n        self.normal_std = np.std(self.normal_features, axis=0) + 1e-8\n        self.normal_features = (self.normal_features - self.normal_mean) / self.normal_std\n        \n        # Apply PCA to normal features\n        self.pca = PCA(n_components=self.n_components)\n        self.pca.fit(self.normal_features)\n        \n    def compute_anomaly_score(self, features):\n        # Normalize features\n        features_norm = (features - self.normal_mean) / self.normal_std\n        \n        # Reconstruct features using PCA\n        reduced = self.pca.transform(features_norm)\n        reconstructed = self.pca.inverse_transform(reduced)\n        \n        # Compute reconstruction error as anomaly score\n        error = np.mean((features_norm - reconstructed) ** 2, axis=1)\n        return error\n    \n    def create_anomaly_heatmap(self, image_features, original_size):\n        # For DINOv3, we need to handle the patch-based structure\n        h, w = original_size\n        \n        # If features are already flattened, create a simple heatmap\n        if len(image_features.shape) == 1:\n            # Create a uniform heatmap based on the overall anomaly score\n            score = self.compute_anomaly_score(image_features.reshape(1, -1))[0]\n            heatmap = np.ones((h, w)) * score\n            return heatmap\n        \n        seq_len = image_features.shape[0]\n        \n        # Calculate the grid size based on patch size\n        grid_size = int(np.sqrt(seq_len))\n        if grid_size * grid_size != seq_len:\n            # If not a perfect square, use approximate reshaping\n            grid_size = int(np.sqrt(seq_len))\n            spatial_features = image_features[:grid_size*grid_size].reshape(grid_size, grid_size, -1)\n        else:\n            spatial_features = image_features.reshape(grid_size, grid_size, -1)\n        \n        # Compute anomaly score for each patch\n        anomaly_scores = []\n        for i in range(grid_size):\n            row_scores = []\n            for j in range(grid_size):\n                patch_feature = spatial_features[i, j].reshape(1, -1)\n                score = self.compute_anomaly_score(patch_feature)\n                row_scores.append(score[0])\n            anomaly_scores.append(row_scores)\n        \n        # Create heatmap\n        heatmap = np.array(anomaly_scores)\n        heatmap = cv2.resize(heatmap, (w, h))\n        \n        return heatmap\n\n# Train anomaly detector\nprint(\"Training anomaly detector...\")\nanomaly_detector = AnomalyDetector(n_components=config.num_components)\nanomaly_detector.fit(train_features, train_labels)\n\n# SAM2 prompt engineering for anomaly segmentation\ndef generate_sam_prompts_from_heatmap(heatmap, threshold=0.7, num_points=10):\n    # Threshold heatmap to get binary mask\n    binary_mask = (heatmap > threshold).astype(np.uint8)\n    \n    # Apply morphological operations to clean up the mask\n    kernel = np.ones((5, 5), np.uint8)\n    binary_mask = cv2.morphologyEx(binary_mask, cv2.MORPH_CLOSE, kernel)\n    binary_mask = cv2.morphologyEx(binary_mask, cv2.MORPH_OPEN, kernel)\n    \n    # Find contours in the binary mask\n    contours, _ = cv2.findContours(binary_mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    \n    # Generate point prompts from contours\n    point_coords = []\n    point_labels = []  # 1 for foreground, 0 for background\n    \n    for contour in contours:\n        # Skip small contours\n        if cv2.contourArea(contour) < 100:\n            continue\n            \n        # Sample points along the contour\n        for i in range(0, len(contour), max(1, len(contour) // num_points)):\n            point = contour[i][0]\n            point_coords.append(point)\n            point_labels.append(1)  # Foreground point\n    \n    # Also add some negative (background) points\n    y, x = np.where(heatmap < threshold * 0.3)  # More conservative for background\n    if len(x) > 0 and len(y) > 0:\n        bg_indices = np.random.choice(len(x), min(5, len(x)), replace=False)\n        for idx in bg_indices:\n            point_coords.append([x[idx], y[idx]])\n            point_labels.append(0)  # Background point\n    \n    # Convert to numpy arrays\n    if point_coords:\n        point_coords = np.array(point_coords)\n        point_labels = np.array(point_labels)\n    else:\n        # If no points found, use center point as fallback\n        h, w = heatmap.shape\n        point_coords = np.array([[w//2, h//2]])\n        point_labels = np.array([1])\n    \n    # Also generate box prompts from contours\n    boxes = []\n    for contour in contours:\n        if cv2.contourArea(contour) < 100:\n            continue\n        x, y, w, h = cv2.boundingRect(contour)\n        boxes.append([x, y, x + w, y + h])\n    \n    boxes = np.array(boxes) if boxes else None\n    \n    return point_coords, point_labels, boxes\n\n# Combined inference function\ndef detect_anomalies(image_path, visualize=False):\n    # Load and preprocess DICOM image\n    try:\n        image = read_dicom_image(image_path)\n        original_size = image.size[::-1]  # (H, W)\n    except:\n        # Return dummy results if image loading fails\n        print(f\"Failed to load DICOM image: {image_path}\")\n        return np.zeros((256, 256)), np.zeros((256, 256))\n    \n    # Transform for DINOv3\n    input_image = test_transform(image).unsqueeze(0).to(device)\n    \n    # Extract features with DINOv3\n    with torch.no_grad():\n        if dinov3_processor is not None:\n            # Use processor if available\n            try:\n                inputs = dinov3_processor(images=input_image, return_tensors=\"pt\").to(device)\n                outputs = dinov3_model(**inputs)\n                features = outputs.last_hidden_state\n            except:\n                # Fallback if processor fails\n                features = dinov3_model(input_image).last_hidden_state\n        else:\n            # Fallback for dummy model\n            features = dinov3_model(input_image)\n    \n    # For DINOv3, we typically use the [CLS] token or average the patch tokens\n    if features.dim() == 3 and features.shape[1] > 1:\n        # Use average of patch tokens (excluding CLS token if present)\n        features = features.mean(dim=1)\n    \n    # Compute anomaly heatmap\n    heatmap = anomaly_detector.create_anomaly_heatmap(\n        features.squeeze().cpu().numpy(), original_size\n    )\n    \n    # Generate prompts for SAM2\n    point_coords, point_labels, boxes = generate_sam_prompts_from_heatmap(\n        heatmap, threshold=config.anomaly_threshold\n    )\n    \n    if sam2_predictor is not None:\n        try:\n            # Prepare image for SAM2\n            sam_image = np.array(image)\n            sam2_predictor.set_image(sam_image)\n            \n            # Segment with SAM2 using the prompts\n            masks, scores, logits = sam2_predictor.predict(\n                point_coords=point_coords,\n                point_labels=point_labels,\n                box=boxes,\n                multimask_output=True\n            )\n            \n            # Get the best mask\n            best_mask = masks[np.argmax(scores)]\n        except Exception as e:\n            print(f\"SAM2 prediction failed: {e}\")\n            # Fallback: use thresholded heatmap as mask\n            best_mask = (heatmap > config.anomaly_threshold).astype(np.uint8)\n    else:\n        # Fallback: use thresholded heatmap as mask\n        best_mask = (heatmap > config.anomaly_threshold).astype(np.uint8)\n    \n    if visualize:\n        # Visualize results\n        fig, axes = plt.subplots(1, 4, figsize=(20, 5))\n        \n        # Original image\n        axes[0].imshow(image)\n        axes[0].set_title('Original DICOM Image')\n        axes[0].axis('off')\n        \n        # Anomaly heatmap\n        axes[1].imshow(heatmap, cmap='hot')\n        axes[1].set_title('Anomaly Heatmap')\n        axes[1].axis('off')\n        \n        # SAM2 prompts\n        axes[2].imshow(image)\n        if len(point_coords) > 0:\n            # Plot foreground points in green\n            fg_coords = point_coords[point_labels == 1]\n            if len(fg_coords) > 0:\n                axes[2].scatter(fg_coords[:, 0], fg_coords[:, 1], c='g', s=50)\n            \n            # Plot background points in red\n            bg_coords = point_coords[point_labels == 0]\n            if len(bg_coords) > 0:\n                axes[2].scatter(bg_coords[:, 0], bg_coords[:, 1], c='r', s=50)\n        \n        # Plot boxes\n        if boxes is not None:\n            for box in boxes:\n                x1, y1, x2, y2 = box\n                axes[2].plot([x1, x2, x2, x1, x1], [y1, y1, y2, y2, y1], 'b-', linewidth=2)\n        \n        axes[2].set_title('SAM2 Prompts')\n        axes[2].axis('off')\n        \n        # Segmentation result\n        axes[3].imshow(image)\n        axes[3].imshow(best_mask, alpha=0.5, cmap='jet')\n        axes[3].set_title('Segmented Anomaly')\n        axes[3].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n    \n    return heatmap, best_mask\n\n# Test the pipeline on a sample image\nsample_images = [f for f in os.listdir(config.train_image_dir) if f.endswith('.dcm')]\nif sample_images:\n    sample_image_path = os.path.join(config.train_image_dir, sample_images[0])\n    print(f\"Testing on sample image: {sample_images[0]}\")\n    heatmap, segmentation_mask = detect_anomalies(sample_image_path, visualize=True)\nelse:\n    print(\"No DICOM sample images found\")\n    print(f\"Files in directory: {os.listdir(config.train_image_dir)[:10]}\")  # Show first 10 files\n\nprint(\"Pipeline completed successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-07T17:39:20.251132Z","iopub.execute_input":"2025-09-07T17:39:20.252134Z","iopub.status.idle":"2025-09-07T17:40:50.260231Z","shell.execute_reply.started":"2025-09-07T17:39:20.252096Z","shell.execute_reply":"2025-09-07T17:40:50.259319Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}