{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport pydicom\nimport cv2\nimport numpy as np\nimport os\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom collections import Counter\nfrom sklearn.metrics import roc_auc_score\n\n# DATA ANALYSIS\nDATA_DIR = '/kaggle/input/competitions/rsna-pneumonia-detection-challenge'\nIMAGES_DIR = os.path.join(DATA_DIR, 'stage_2_train_images')\nLABELS_PATH = os.path.join(DATA_DIR, 'stage_2_train_labels.csv')\n\ndf = pd.read_csv(LABELS_PATH)\n\n# Basic statistics\nprint(f\"\\nTotal entries in CSV: {len(df)}\")\nprint(f\"Unique patients: {df['patientId'].nunique()}\")\nprint(f\"Columns: {df.columns.tolist()}\")\nprint(f\"\\nFirst 5 rows:\")\nprint(df.head())\n\n# Class distribution\ntarget_counts = df.groupby('patientId')['Target'].max().value_counts()\nprint(f\"\\nClass distribution:\")\nprint(f\"Normal (0):    {target_counts.get(0, 0)} patients\")\nprint(f\"Pneumonia (1): {target_counts.get(1, 0)} patients\")\nprint(f\"Ratio (pneumonia/normal): {target_counts.get(1, 0) / target_counts.get(0, 1):.3f}\")\n\n# Bounding box statistics (only for positive cases)\npositive_cases = df[df['Target'] == 1]\nif len(positive_cases) > 0:\n    print(f\"\\nBounding box statistics (pneumonia cases):\")\n    print(f\"Total boxes: {len(positive_cases)}\")\n    print(f\"Patients with boxes: {positive_cases['patientId'].nunique()}\")\n    print(f\"Boxes per patient (avg): {len(positive_cases) / positive_cases['patientId'].nunique():.2f}\")\n    \n    # Box sizes\n    box_areas = positive_cases['width'] * positive_cases['height']\n    print(f\"Box area - mean: {box_areas.mean():.0f}, std: {box_areas.std():.0f}\")\n    print(f\"Box area - min: {box_areas.min():.0f}, max: {box_areas.max():.0f}\")\n    \n    print(f\"Box width  - mean: {positive_cases['width'].mean():.0f}, std: {positive_cases['width'].std():.0f}\")\n    print(f\"Box height - mean: {positive_cases['height'].mean():.0f}, std: {positive_cases['height'].std():.0f}\")\n\nprint(f\"\\nSample DICOM analysis (first 5 files):\")\nimport glob\ndcm_files = glob.glob(os.path.join(IMAGES_DIR, '*.dcm'))\nprint(f\"Total DICOM files found: {len(dcm_files)}\")\n\nsex_counter = {'M': 0, 'F': 0}\nfor i, dcm_path in enumerate(dcm_files[:100]):  # Check first 100 for speed\n    try:\n        dcm = pydicom.dcmread(dcm_path, stop_before_pixels=True)\n        sex = dcm.get('PatientSex', 'Unknown')\n        if sex in sex_counter:\n            sex_counter[sex] += 1\n    except:\n        pass\n\nprint(f\"\\nSex distribution (from {sum(sex_counter.values())} sampled files):\")\nprint(f\"Male (M):   {sex_counter['M']} ({sex_counter['M']/sum(sex_counter.values())*100:.1f}%)\")\nprint(f\"Female (F): {sex_counter['F']} ({sex_counter['F']/sum(sex_counter.values())*100:.1f}%)\")\n\nsample_dcm = pydicom.dcmread(dcm_files[0])\nprint(f\"\\nSample DICOM metadata:\")\nprint(f\"  Shape: {sample_dcm.pixel_array.shape}\")\nprint(f\"  Min pixel value: {sample_dcm.pixel_array.min()}\")\nprint(f\"  Max pixel value: {sample_dcm.pixel_array.max()}\")\nprint(f\"  PatientSex: {sample_dcm.get('PatientSex', 'N/A')}\")\nprint(f\"  PatientAge: {sample_dcm.get('PatientAge', 'N/A')}\")\n\nclass PneumoniaDataset(Dataset):\n    \"\"\"\n    PyTorch Dataset for RSNA Pneumonia Detection Challenge.\n    Loads DICOM files, extracts pixels and sex, applies transforms.\n    \"\"\"\n\n    def __init__(self, df, images_dir, is_train=True, normalize_stats=None, image_size=320):\n        \"\"\"\n        Args:\n            df: DataFrame with columns [patientId, Target, x, y, width, height]\n            images_dir: Path to DICOM files\n            is_train: If True, apply augmentations\n            normalize_stats: tuple (mean, std)\n            image_size: Target image size\n        \"\"\"\n        self.df = df.groupby('patientId').agg({\n            'Target': 'max',\n            'x': 'first',\n            'y': 'first',\n            'width': 'first',\n            'height': 'first'\n        }).reset_index()\n\n        self.images_dir = images_dir\n        self.is_train = is_train\n        self.image_size = image_size\n\n        # Normalization\n        if normalize_stats is not None:\n            mean, std = normalize_stats\n            self.normalize = transforms.Normalize(\n                mean=[mean],\n                std=[std]\n            )\n        else:\n            self.normalize = None\n\n        # Base transform\n        self.base_transform = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((image_size, image_size)),\n            transforms.ToTensor(),\n        ])\n\n        # Training augmentations\n        self.train_transform = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.Resize((image_size, image_size)),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomRotation(degrees=10),\n            transforms.RandomAffine(\n                degrees=0,\n                translate=(0.05, 0.05)\n            ),\n            transforms.ColorJitter(\n                brightness=0.1,\n                contrast=0.1\n            ),\n            transforms.ToTensor(),\n        ])\n\n    def __len__(self):\n        return len(self.df)\n\n    @staticmethod\n    def preprocess_raw_image(image: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Normalize raw DICOM pixel values to [0, 1].\n        \"\"\"\n        image = image.astype(np.float32)\n        image = (image - image.min()) / (image.max() - image.min() + 1e-8)\n\n        return image\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        patient_id = row['patientId']\n        target = row['Target']\n\n        # Load DICOM\n        dcm_path = os.path.join(\n            self.images_dir,\n            f\"{patient_id}.dcm\"\n        )\n\n        dcm = pydicom.dcmread(dcm_path)\n        image = dcm.pixel_array\n\n        # Normalize DICOM intensities\n        image = self.preprocess_raw_image(image)\n\n        # Get sex for fairness\n        sex_str = str(dcm.get('PatientSex', 'Unknown'))\n        sex = 1 if sex_str == 'M' else 0\n\n        # Apply transforms\n        if self.is_train:\n            image = self.train_transform(image)\n        else:\n            image = self.base_transform(image)\n\n        # Normalize with dataset statistics\n        if self.normalize is not None:\n            image = self.normalize(image)\n\n        return {\n            'image': image,\n            'target': torch.tensor(\n                target,\n                dtype=torch.float32\n            ),\n            'patient_id': patient_id,\n            'sex': sex\n        }\n\n\nclass SimplePneumoniaClassifier(nn.Module):\n    \"\"\"\n    A pneumonia classification model that takes X-ray images and predicts whether pneumonia is present.\n    \n    Your model should follow this basic structure, but you are free to modify the internal architecture.\n    \"\"\"\n    def __init__(self, checkpoint_dir='checkpoints', normalize_stats=None, image_size=320):\n        \"\"\"\n        Initialize your model.\n        \n        Args:\n            checkpoint_dir (str): Directory where model checkpoints will be saved\n        \"\"\"\n        super(SimplePneumoniaClassifier, self).__init__()\n\n        self.checkpoint_dir = checkpoint_dir\n        os.makedirs(checkpoint_dir, exist_ok=True)\n\n        self.image_size = image_size\n        self.normalize_stats = normalize_stats\n\n        # Device\n        self.device = torch.device(\n            'cuda' if torch.cuda.is_available() else 'cpu'\n        )\n\n        # Load pretrained EfficientNet-B0\n        from torchvision.models import (efficientnet_b0, EfficientNet_B0_Weights)\n\n        self.backbone = efficientnet_b0(\n            weights=EfficientNet_B0_Weights.IMAGENET1K_V1\n        )\n\n        # Convert first conv layer from RGB -> grayscale\n        old_conv = self.backbone.features[0][0]\n\n        new_conv = nn.Conv2d(\n            in_channels=1,\n            out_channels=old_conv.out_channels,\n            kernel_size=old_conv.kernel_size,\n            stride=old_conv.stride,\n            padding=old_conv.padding,\n            bias=False\n        )\n\n        with torch.no_grad():\n            new_conv.weight.copy_(\n                old_conv.weight.mean(\n                    dim=1,\n                    keepdim=True\n                )\n            )\n\n        self.backbone.features[0][0] = new_conv\n\n        # Replace classifier\n        in_features = self.backbone.classifier[1].in_features\n\n        self.backbone.classifier = nn.Sequential(\n            nn.Dropout(p=0.2, inplace=True),\n            nn.Linear(in_features, 1)\n        )\n\n        self.to(self.device)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Forward pass of the model.\n        \n        Args:\n            x (torch.Tensor): Input tensor of shape [batch_size, 1, height, width] \n                             containing grayscale X-ray images\n        \n        Returns:\n            torch.Tensor: Output tensor with shape [batch_size, 1] containing probabilities \n                         of pneumonia (values between 0 and 1)\n        \"\"\"\n        return self.backbone(x)\n\n    def load_checkpoint(self, checkpoint_path: str) -> dict:\n        \"\"\"\n        Load model weights from a checkpoint file.\n        \n        Args:\n            checkpoint_path (str): Path to the checkpoint file\n            \n        Returns:\n            dict: Checkpoint data including 'epoch' and other training metadata\n        \"\"\"\n        checkpoint = torch.load(checkpoint_path, map_location=self.device, weights_only=False)\n\n        # Load the state dictionary into the model\n        self.load_state_dict(checkpoint['model_state_dict'])\n\n        if 'normalize_stats' in checkpoint:\n            self.normalize_stats = checkpoint['normalize_stats']\n        if 'fair_thresholds' in checkpoint:\n            self.fair_thresholds = checkpoint['fair_thresholds']\n        if 'image_size' in checkpoint:\n            self.image_size = checkpoint['image_size']\n\n        return checkpoint\n\n    def preprocess_for_inference(self, image):\n        \"\"\"\n        Preprocess raw image for inference.\n        \"\"\"\n        # Convert tensor -> numpy if needed\n        if isinstance(image, torch.Tensor):\n            image = image.cpu().numpy()\n\n        # Remove extra dimensions if present\n        image = np.squeeze(image)\n        # Convert to float32\n        image = image.astype(np.float32)\n        # Normalize to [0,1]\n        image = (image - image.min()) / (image.max() - image.min() + 1e-8)\n        # Resize\n        image = cv2.resize(image, (self.image_size, self.image_size))\n        # Convert to tensor\n        image = torch.from_numpy(image).float()\n        # Add channel dimension\n        image = image.unsqueeze(0)\n        \n        # Dataset normalization\n        if self.normalize_stats is not None:\n\n            mean, std = self.normalize_stats\n            normalize = transforms.Normalize(\n                mean=[mean],\n                std=[std]\n            )\n            image = normalize(image)\n\n        # Add batch dimension\n        image = image.unsqueeze(0)\n\n        return image.to(self.device)\n\n    def predict(self, image, device='cpu'):\n        \"\"\"\n        Make a prediction for a single image.\n        \n        Args:\n            image: Input image (can be numpy array or tensor)\n            device: Device to use for computation ('cpu', 'cuda', or 'mps')\n            \n        Returns:\n            dict: Dictionary containing:\n                - 'probability': Float value between 0 and 1\n                - 'class': Binary class (0 or 1)\n                - 'label': String label ('Normal' or 'Pneumonia')\n        \"\"\"\n        self.eval()\n        image = self.preprocess_for_inference(image)\n        with torch.no_grad():\n            logits = self.forward(image)\n            probability = torch.sigmoid(\n                logits\n            ).item()\n\n        pred_class = 1 if probability >= 0.5 else 0\n        label = 'Pneumonia' if pred_class == 1 else 'Normal'\n\n        return {\n            'probability': probability,\n            'class': pred_class,\n            'label': label\n        }\n\n\ndef compute_normalization_stats(df, images_dir, image_size=320, num_samples=2000):\n    \"\"\"\n    Compute mean and std across training images for proper normalization.\n    Samples a subset for speed.\n    \"\"\"\n    patient_ids = df['patientId'].unique()\n    sample_ids = np.random.choice(patient_ids, min(num_samples, len(patient_ids)), replace=False)\n\n    means = []\n    stds = []\n\n    for pid in sample_ids:\n        dcm_path = os.path.join(images_dir, f\"{pid}.dcm\")\n        dcm = pydicom.dcmread(dcm_path)\n        image = dcm.pixel_array.astype(np.float32)\n        image = (image - image.min()) / (image.max() - image.min() + 1e-8)\n\n        # Resize\n        image = cv2.resize(image,(image_size, image_size))\n\n        means.append(image.mean())\n        stds.append(image.std())\n\n    global_mean = np.mean(means)\n    global_std = np.mean(stds)\n\n    print(f\"Normalization stats — Mean: {global_mean:.2f}, Std: {global_std:.2f}\")\n\n    return global_mean, global_std\n\n# TRAINING\ndef train_model(train_dataset, val_dataset, normalize_stats): \n    \"\"\"\n    Train the pneumonia classifier.\n    \n    Args:\n        train_dataset: PneumoniaDataset for training\n        val_dataset: PneumoniaDataset for validation\n        normalize_stats: tuple (mean, std) to save in checkpoint\n    \"\"\"    \n    # Calculate pos_weight for class imbalance (normal/pneumonia ratio)\n    pos_count = (train_dataset.df['Target'] == 1).sum()\n    neg_count = (train_dataset.df['Target'] == 0).sum()\n    pos_weight = neg_count / pos_count\n    print(f\"\\npos_weight for BCE: {pos_weight:.2f}\")\n\n    # DataLoaders\n    train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=2)\n\n    # Model\n    model = SimplePneumoniaClassifier(normalize_stats=normalize_stats)\n\n    # Loss\n    criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([pos_weight]).to(model.device))\n\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', patience=2, factor=0.5)\n\n    num_epochs = 10\n    best_auc = 0.0\n\n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0.0\n        for batch in train_loader:\n            images = batch['image'].to(model.device)\n            targets = batch['target'].to(model.device).unsqueeze(1)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, targets)\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n            \n        # Validation\n        model.eval()\n        val_loss = 0.0\n        all_preds = []\n        all_targets = []\n\n        with torch.no_grad():\n            for batch in val_loader:\n                images = batch['image'].to(model.device)\n                targets = batch['target'].to(model.device).unsqueeze(1)\n                \n                outputs = model(images)\n                loss = criterion(outputs, targets)\n                val_loss += loss.item()\n                probs = torch.sigmoid(outputs)\n\n                all_preds.extend(probs.cpu().numpy().flatten())\n                all_targets.extend(targets.cpu().numpy().flatten())\n        \n        val_auc = roc_auc_score(all_targets, all_preds)\n\n        print(f\"Epoch {epoch+1}/{num_epochs}\")\n        print(f\"  Train Loss: {train_loss/len(train_loader):.4f}\")\n        print(f\"  Val Loss:   {val_loss/len(val_loader):.4f}\")\n        print(f\"  Val AUC:    {val_auc:.4f}\")\n        \n        scheduler.step(val_auc)\n\n        # Save best model\n        if val_auc > best_auc:\n            best_auc = val_auc\n            checkpoint = {\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_auc': val_auc,\n                'normalize_stats': normalize_stats,\n                'image_size': model.image_size,\n            }\n            torch.save(checkpoint, os.path.join(model.checkpoint_dir, 'best_model.pt'))\n            print(f\"  >>> Saved best model with AUC: {val_auc:.4f}\")\n    \n    print(f\"\\nTraining complete! Best AUC: {best_auc:.4f}\")\n    return model\n\ndef get_importance_heatmaps(model: SimplePneumoniaClassifier, \n                           images: list, \n                           window_size: int = 32, \n                           stride: int = 16) -> list:\n    \"\"\"\n    Generate occlusion sensitivity maps for a batch of images.\n    \n    This function should create heatmaps that highlight regions important for \n    the model's prediction. For pneumonia cases, the heatmap should focus on\n    the areas of the image that contain the pneumonia opacity.\n    \n    Args:\n        model: Trained PyTorch model (SimplePneumoniaClassifier)\n        images: List or tensor of input images\n        window_size: Size of the occlusion window (default: 32)\n        stride: Stride of the sliding window (default: 16)\n        \n    Returns:\n        heatmaps: List of numpy arrays, each representing a sensitivity map \n                 with the same height and width as the original image.\n                 Values should be normalized between 0 and 1.\n    \"\"\"\n    model.eval()\n\n    # Convert tensor batch -> list\n    if isinstance(images, torch.Tensor):\n        images = [img for img in images]\n\n    heatmaps = []\n\n    for image in images:\n        # Save original size\n        if isinstance(image, torch.Tensor):\n            original_image = image.squeeze().cpu().numpy()\n        else:\n            original_image = np.squeeze(image)\n        original_h, original_w = original_image.shape\n        \n        # Preprocess image for model\n        input_tensor = model.preprocess_for_inference(original_image)\n        \n        # Baseline prediction\n        with torch.no_grad():\n            baseline_logits = model(input_tensor)\n            baseline_prob = torch.sigmoid(baseline_logits).item()\n\n        # Working image in resized model space\n        working_image = input_tensor.squeeze().cpu().numpy()\n\n        # Remove normalization effect for occlusion\n        if model.normalize_stats is not None:\n            mean, std = model.normalize_stats\n            working_image = (working_image * std) + mean\n\n        h, w = working_image.shape\n\n        heatmap = np.zeros((h, w), dtype=np.float32)\n        count_map = np.zeros((h, w), dtype=np.float32)\n\n        occluded_images = []\n        positions = []\n\n        # Generate occluded images\n        for y in range(0, h - window_size + 1, stride):\n            for x in range(0, w - window_size + 1, stride):\n                occluded = working_image.copy()\n\n                blurred = cv2.GaussianBlur(\n                    working_image,\n                    (31, 31),\n                    0\n                )\n                \n                occluded[y:y + window_size, x:x + window_size] = \\\n                    blurred[y:y + window_size, x:x + window_size]\n                occluded_tensor = torch.from_numpy(occluded).float()\n\n                # Re-apply normalization\n                if model.normalize_stats is not None:\n                    normalize = transforms.Normalize(mean=[mean], std=[std])\n                    occluded_tensor = normalize(occluded_tensor.unsqueeze(0)).squeeze(0)\n\n                occluded_tensor = occluded_tensor.unsqueeze(0).unsqueeze(0)\n                occluded_images.append(occluded_tensor)\n                positions.append((x, y))\n\n        # Batch inference for speed\n        batch_size = 32\n        occlusion_scores = []\n\n        with torch.no_grad():\n            for i in range(0, len(occluded_images), batch_size):\n                batch = torch.cat(occluded_images[i:i + batch_size], dim=0).to(model.device)\n                logits = model(batch)\n                probs = torch.sigmoid(logits)\n                probs = probs.cpu().numpy().flatten()\n                occlusion_scores.extend(probs)\n\n        # Build heatmap\n        for (x, y), occluded_prob in zip(positions, occlusion_scores):\n            importance = baseline_prob - occluded_prob\n\n            # Keep only positive importance\n            importance = max(0.0, importance)\n\n            heatmap[y:y + window_size, x:x + window_size] += importance\n            count_map[y:y + window_size, x:x + window_size] += 1.0\n\n        # Average overlapping regions\n        heatmap = heatmap / (count_map + 1e-8)\n        # Smooth heatmap\n        heatmap = cv2.GaussianBlur(heatmap, (11, 11), 0)\n\n        # Sharpen important regions\n        heatmap = np.power(heatmap, 2)\n\n        # Normalize to [0,1]\n        heatmap = heatmap - heatmap.min()\n        heatmap = heatmap / (heatmap.max() + 1e-8)\n\n        # Resize back to original image size\n        heatmap = cv2.resize(heatmap, (original_w, original_h))\n        heatmaps.append(heatmap.astype(np.float32))\n\n    return heatmaps\n\ndef generate_gradcam_heatmap(model, image):\n    \"\"\"\n    Generate Grad-CAM heatmap for a single image.\n    \"\"\"\n    model.eval()\n    input_tensor = model.preprocess_for_inference(image)\n\n    # Target layer (last convolution block)\n    target_layer = model.backbone.features[-1]\n\n    activations = []\n    gradients = []\n\n    # Hooks\n    def forward_hook(module, input, output):\n        activations.append(output)\n\n    def backward_hook(module, grad_input, grad_output):\n        gradients.append(grad_output[0])\n\n    forward_handle = target_layer.register_forward_hook(forward_hook)\n\n    backward_handle = target_layer.register_full_backward_hook(backward_hook)\n\n    # Forward\n    logits = model(input_tensor)\n\n    # Backward\n    model.zero_grad()\n    logits[:, 0].backward()\n\n    # Get tensors\n    activation = activations[0]\n    gradient = gradients[0]\n\n    # Global average pooling on gradients\n    weights = gradient.mean(dim=(2, 3), keepdim=True)\n\n    # Weighted sum\n    cam = (weights * activation).sum(dim=1)\n    cam = torch.relu(cam)\n    cam = cam.squeeze().detach().cpu().numpy()\n\n    # Normalize\n    cam = cam - cam.min()\n    cam = cam / (cam.max() + 1e-8)\n\n    # Resize to original image size\n    original_image = np.squeeze(image)\n    h, w = original_image.shape\n    cam = cv2.resize(cam, (w, h))\n\n    # Remove hooks\n    forward_handle.remove()\n    backward_handle.remove()\n\n    return cam.astype(np.float32)\n\ndef compute_heatmap_iou(heatmap, bbox, threshold=0.5):\n    \"\"\"\n    Compute IoU between heatmap and ground-truth bounding box.\n\n    Args:\n        heatmap: numpy array in [0,1]\n        bbox: (x, y, w, h)\n        threshold: threshold for binary heatmap\n\n    Returns:\n        iou: Intersection over Union score\n    \"\"\"\n    x, y, w, h = bbox\n\n    # Binary prediction mask\n    pred_mask = (heatmap >= threshold).astype(np.uint8)\n\n    # Ground truth mask\n    gt_mask = np.zeros_like(pred_mask, dtype=np.uint8)\n\n    x2 = min(x + w, gt_mask.shape[1])\n    y2 = min(y + h, gt_mask.shape[0])\n\n    gt_mask[y:y2, x:x2] = 1\n\n    # Intersection / Union\n    intersection = np.logical_and(pred_mask, gt_mask).sum()\n\n    union = np.logical_or(pred_mask, gt_mask).sum()\n\n    if union == 0:\n        return 0.0\n\n    return intersection / union\n\ndef compare_explainability_methods(model,\n                                   image,\n                                   bbox=None,\n                                   window_size=32,\n                                   stride=16):\n    \"\"\"\n    Compare Occlusion Sensitivity vs Grad-CAM.\n    \"\"\"\n    original = np.squeeze(image)\n\n    occlusion_map = get_importance_heatmaps(\n        model,\n        [image],\n        window_size=window_size,\n        stride=stride\n    )[0]\n\n    gradcam_map = generate_gradcam_heatmap(model, image)\n\n    # Quantitative overlap with GT box\n    if bbox is not None:\n        occlusion_iou = compute_heatmap_iou(occlusion_map, bbox)\n        gradcam_iou = compute_heatmap_iou(gradcam_map, bbox)\n        print(f\"\\nQuantitative overlap with GT box:\")\n        print(f\"Occlusion IoU: {occlusion_iou:.4f}\")\n        print(f\"Grad-CAM IoU:  {gradcam_iou:.4f}\")\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n\n    # Original image\n    axes[0].imshow(original, cmap='gray')\n    if bbox is not None:\n        x, y, w, h = bbox\n        rect = plt.Rectangle((x, y), w, h, fill=False, color='lime', linewidth=2)\n        axes[0].add_patch(rect)\n    axes[0].set_title(\"Original Image\")\n    axes[0].axis('off')\n\n    # Occlusion\n    axes[1].imshow(original, cmap='gray')\n    axes[1].imshow(occlusion_map, cmap='jet', alpha=0.5)\n    if bbox is not None:\n        rect = plt.Rectangle((x, y), w, h, fill=False, color='lime', linewidth=2)\n        axes[1].add_patch(rect)\n    axes[1].set_title(\"Occlusion Sensitivity\")\n    axes[1].axis('off')\n\n    # Grad-CAM\n    axes[2].imshow(original, cmap='gray')\n    axes[2].imshow(gradcam_map, cmap='jet', alpha=0.5)\n    if bbox is not None:\n        rect = plt.Rectangle((x, y), w, h, fill=False, color='lime', linewidth=2)\n        axes[2].add_patch(rect)\n    axes[2].set_title(\"Grad-CAM\")\n    axes[2].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\ndef find_fair_thresholds(model, val_dataset):\n    \"\"\"\n    Find group-specific thresholds to achieve fairness.\n    Returns thresholds for M and F that balance TPR and prediction rate.\n    \"\"\"\n    from sklearn.metrics import roc_curve\n    \n    model.eval()\n    \n    # Collect predictions and labels by group\n    male_probs = []\n    male_targets = []\n    female_probs = []\n    female_targets = []\n    \n    val_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers=2)\n    \n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(model.device)\n            targets = batch['target'].numpy()\n            sexes = batch['sex'].numpy()  # 1 = M, 0 = F\n            \n            logits = model(images)\n            probs = torch.sigmoid(logits).cpu().numpy().flatten()\n            \n            for i in range(len(probs)):\n                if sexes[i] == 1:  # Male\n                    male_probs.append(probs[i])\n                    male_targets.append(targets[i])\n                else:  # Female\n                    female_probs.append(probs[i])\n                    female_targets.append(targets[i])\n    \n    male_probs = np.array(male_probs)\n    male_targets = np.array(male_targets)\n    female_probs = np.array(female_probs)\n    female_targets = np.array(female_targets)\n    \n    print(f\"Male samples:   {len(male_targets)}\")\n    print(f\"Female samples: {len(female_targets)}\")\n    \n    # Default threshold metrics\n    default_threshold = 0.5\n    \n    male_pred_default = (male_probs >= default_threshold).astype(int)\n    female_pred_default = (female_probs >= default_threshold).astype(int)\n    \n    male_tpr_default = (male_pred_default & male_targets.astype(bool)).sum() / male_targets.sum()\n    female_tpr_default = (female_pred_default & female_targets.astype(bool)).sum() / female_targets.sum()\n    \n    male_rate_default = male_pred_default.mean()\n    female_rate_default = female_pred_default.mean()\n    \n    print(f\"\\nDefault threshold (0.5):\")\n    print(f\"  Male TPR: {male_tpr_default:.4f}, Female TPR: {female_tpr_default:.4f}\")\n    print(f\"  TPR disparity: {abs(male_tpr_default - female_tpr_default):.4f}\")\n    print(f\"  Male pred rate: {male_rate_default:.4f}, Female pred rate: {female_rate_default:.4f}\")\n    print(f\"  Pred rate disparity: {abs(male_rate_default - female_rate_default):.4f}\")\n    \n    # Overall AUC check\n    all_probs = np.concatenate([male_probs, female_probs])\n    all_targets = np.concatenate([male_targets, female_targets])\n    overall_auc = roc_auc_score(all_targets, all_probs)\n    print(f\"  Overall AUC: {overall_auc:.4f}\")\n    \n    # Find thresholds that equalize TPR (prioritize TPR disparity)\n    best_male_threshold = 0.5\n    best_female_threshold = 0.5\n    best_tpr_disparity = float('inf')\n    best_rate_disparity = float('inf')\n    \n    # Two-pass search: first minimize TPR disparity, then rate disparity\n    candidates = []\n    \n    for male_thresh in np.arange(0.15, 0.85, 0.01):\n        male_pred = (male_probs >= male_thresh).astype(int)\n        if male_targets.sum() > 0:\n            male_tpr = (male_pred & male_targets.astype(bool)).sum() / male_targets.sum()\n        else:\n            continue\n        male_rate = male_pred.mean()\n        \n        for female_thresh in np.arange(0.15, 0.85, 0.01):\n            female_pred = (female_probs >= female_thresh).astype(int)\n            if female_targets.sum() > 0:\n                female_tpr = (female_pred & female_targets.astype(bool)).sum() / female_targets.sum()\n            else:\n                continue\n            female_rate = female_pred.mean()\n            \n            tpr_disparity = abs(male_tpr - female_tpr)\n            rate_disparity = abs(male_rate - female_rate)\n            \n            # Only consider candidates with low rate disparity\n            if rate_disparity < 0.01:\n                candidates.append({\n                    'male_thresh': male_thresh,\n                    'female_thresh': female_thresh,\n                    'tpr_disparity': tpr_disparity,\n                    'rate_disparity': rate_disparity,\n                    'male_tpr': male_tpr,\n                    'female_tpr': female_tpr,\n                    'male_rate': male_rate,\n                    'female_rate': female_rate\n                })\n    \n    # Find candidate with minimum TPR disparity\n    if candidates:\n        best = min(candidates, key=lambda x: x['tpr_disparity'])\n        best_male_threshold = best['male_thresh']\n        best_female_threshold = best['female_thresh']\n        best_tpr_disparity = best['tpr_disparity']\n        best_rate_disparity = best['rate_disparity']\n        male_tpr_best = best['male_tpr']\n        female_tpr_best = best['female_tpr']\n        male_rate_best = best['male_rate']\n        female_rate_best = best['female_rate']\n    else:\n        # Fallback: keep searching with relaxed rate constraint\n        for male_thresh in np.arange(0.2, 0.8, 0.02):\n            male_pred = (male_probs >= male_thresh).astype(int)\n            if male_targets.sum() > 0:\n                male_tpr = (male_pred & male_targets.astype(bool)).sum() / male_targets.sum()\n            else:\n                continue\n            male_rate = male_pred.mean()\n            \n            for female_thresh in np.arange(0.2, 0.8, 0.02):\n                female_pred = (female_probs >= female_thresh).astype(int)\n                if female_targets.sum() > 0:\n                    female_tpr = (female_pred & female_targets.astype(bool)).sum() / female_targets.sum()\n                else:\n                    continue\n                female_rate = female_pred.mean()\n                \n                tpr_disparity = abs(male_tpr - female_tpr)\n                rate_disparity = abs(male_rate - female_rate)\n                \n                if tpr_disparity < best_tpr_disparity and rate_disparity < 0.015:\n                    best_tpr_disparity = tpr_disparity\n                    best_male_threshold = male_thresh\n                    best_female_threshold = female_thresh\n                    male_tpr_best = male_tpr\n                    female_tpr_best = female_tpr\n                    male_rate_best = male_rate\n                    female_rate_best = female_rate\n                    best_rate_disparity = rate_disparity\n    \n    male_pred_best = (male_probs >= best_male_threshold).astype(int)\n    female_pred_best = (female_probs >= best_female_threshold).astype(int)\n    \n    male_tpr_best = (male_pred_best & male_targets.astype(bool)).sum() / male_targets.sum()\n    female_tpr_best = (female_pred_best & female_targets.astype(bool)).sum() / female_targets.sum()\n    \n    male_rate_best = male_pred_best.mean()\n    female_rate_best = female_pred_best.mean()\n    \n    print(f\"\\nFair thresholds:\")\n    print(f\"  Male threshold:   {best_male_threshold:.2f}\")\n    print(f\"  Female threshold: {best_female_threshold:.2f}\")\n    print(f\"  Male TPR: {male_tpr_best:.4f}, Female TPR: {female_tpr_best:.4f}\")\n    print(f\"  TPR disparity: {abs(male_tpr_best - female_tpr_best):.4f}\")\n    print(f\"  Male pred rate: {male_rate_best:.4f}, Female pred rate: {female_rate_best:.4f}\")\n    print(f\"  Pred rate disparity: {abs(male_rate_best - female_rate_best):.4f}\")\n    \n    # Check AUC still >= 0.8 with new thresholds\n    all_pred_binary = np.concatenate([male_pred_best, female_pred_best])\n    from sklearn.metrics import accuracy_score\n    fair_accuracy = accuracy_score(all_targets, all_pred_binary)\n    print(f\"  Overall accuracy with fair thresholds: {fair_accuracy:.4f}\")\n    print(f\"  Overall AUC (unchanged): {overall_auc:.4f}\")\n    \n    thresholds = {\n        'M': best_male_threshold,\n        'F': best_female_threshold,\n        'default': 0.5\n    }\n    \n    return thresholds\n\ndef fair_predict(model: SimplePneumoniaClassifier, \n                images: list, \n                sex_attribute: list = None) -> list:\n    \"\"\"\n    Make fair predictions on demographic attributes.\n    \n    Args:\n        model: Trained model (SimplePneumoniaClassifier)\n        images: List or tensor of input images\n        sex_attribute: List of sex attributes corresponding to images ('M' or 'F')\n                      Can be None if demographic information is not available\n        \n    Returns:\n        List of prediction dictionaries, each containing:\n            - 'probability': Raw probability from model (float between 0 and 1)\n            - 'threshold': Threshold used for this prediction\n            - 'class': Binary prediction (0 or 1) after applying threshold\n            - 'label': String label ('Normal' or 'Pneumonia')\n    \"\"\"\n    model.eval()    \n    results = []\n    \n    # Fair thresholds (computed from validation set)\n    fair_thresholds = getattr(model, 'fair_thresholds', {'M': 0.5, 'F': 0.5})\n    \n    for image, sex in zip(images, sex_attribute):\n        # Preprocess\n        input_tensor = model.preprocess_for_inference(image)\n        \n        with torch.no_grad():\n            logits = model(input_tensor)\n            prob = torch.sigmoid(logits).item()\n        \n        # Apply group-specific threshold\n        if sex == 'M':\n            threshold = fair_thresholds.get('M', 0.5)\n        elif sex == 'F':\n            threshold = fair_thresholds.get('F', 0.5)\n        else:\n            threshold = 0.5\n        \n        pred_class = 1 if prob >= threshold else 0\n        label = 'Pneumonia' if pred_class == 1 else 'Normal'\n        \n        results.append({\n            'probability': prob,\n            'threshold': threshold,\n            'class': pred_class,\n            'label': label\n        })\n    \n    return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T11:39:50.200043Z","iopub.execute_input":"2026-05-14T11:39:50.200412Z","iopub.status.idle":"2026-05-14T11:40:06.555196Z","shell.execute_reply.started":"2026-05-14T11:39:50.200382Z","shell.execute_reply":"2026-05-14T11:40:06.554460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\npatient_df = df.groupby('patientId')['Target'].max().reset_index()\n\ntrain_patients, val_patients = train_test_split(\n    patient_df,\n    test_size=0.2,\n    random_state=42,\n    stratify=patient_df['Target']\n)\n\ntrain_df = df[df['patientId'].isin(train_patients['patientId'])]\nval_df = df[df['patientId'].isin(val_patients['patientId'])]\n\nnormalize_stats = compute_normalization_stats(train_df, IMAGES_DIR)\n\ntrain_dataset = PneumoniaDataset(\n    train_df, IMAGES_DIR, is_train=True, normalize_stats=normalize_stats\n)\nval_dataset = PneumoniaDataset(\n    val_df, IMAGES_DIR, is_train=False, normalize_stats=normalize_stats\n)\n\n# Train model\nmodel = train_model(train_dataset, val_dataset, normalize_stats)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T11:40:06.556734Z","iopub.execute_input":"2026-05-14T11:40:06.557067Z","iopub.status.idle":"2026-05-14T12:26:49.179970Z","shell.execute_reply.started":"2026-05-14T11:40:06.557040Z","shell.execute_reply":"2026-05-14T12:26:49.178957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Test heatmap on one example\nsample_row = positive_cases.iloc[30]\npatient_id = sample_row['patientId']\ndcm_path = os.path.join(IMAGES_DIR, f\"{patient_id}.dcm\")\ndcm = pydicom.dcmread(dcm_path)\nimage = dcm.pixel_array.astype(np.float32)\nbbox = (\n    int(sample_row['x']), int(sample_row['y']),\n    int(sample_row['width']), int(sample_row['height'])\n)\ncompare_explainability_methods(model, image, bbox=bbox)\n\n# Fairness\nfair_thresholds = find_fair_thresholds(model, val_dataset)\nmodel.fair_thresholds = fair_thresholds\n\ncheckpoint_path = os.path.join(model.checkpoint_dir, 'best_model.pt')\nif os.path.exists(checkpoint_path):\n    checkpoint = torch.load(checkpoint_path, map_location=model.device, weights_only=False)\n    checkpoint['fair_thresholds'] = fair_thresholds\n    torch.save(checkpoint, checkpoint_path)\n    print(f\"\\nFair thresholds saved to checkpoint: M={fair_thresholds['M']:.2f}, F={fair_thresholds['F']:.2f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T12:26:49.181621Z","iopub.execute_input":"2026-05-14T12:26:49.181934Z","iopub.status.idle":"2026-05-14T12:27:33.577952Z","shell.execute_reply.started":"2026-05-14T12:26:49.181903Z","shell.execute_reply":"2026-05-14T12:27:33.577073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import shutil\n\nshutil.make_archive('checkpoints', 'zip', 'checkpoints')\n\nfrom IPython.display import FileLink\nFileLink('checkpoints.zip')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-14T12:27:33.579181Z","iopub.execute_input":"2026-05-14T12:27:33.579582Z","iopub.status.idle":"2026-05-14T12:27:35.988650Z","shell.execute_reply.started":"2026-05-14T12:27:33.579547Z","shell.execute_reply":"2026-05-14T12:27:35.987815Z"}},"outputs":[],"execution_count":null}]}