{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":10338,"databundleVersionId":862042,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **RSNA: PBBox Regression for PNA Detection**","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"https://www.kaggle.com/competitions/rsna-pneumonia-detection-challenge\n","metadata":{},"attachments":{}},{"cell_type":"markdown","source":"---\n\n### **Introduction**  \nPneumonia remains a leading cause of hospitalization worldwide, and rapid diagnosis from chest X-rays is critical for patient outcomes. This project addresses the **RSNA Pneumonia Detection Challenge** by developing a deep learning pipeline that:  \n1. **Classifies** whether a DICOM chest X-ray contains pneumonia (binary `Target`).  \n2. **Localizes** suspicious opacities with bounding boxes (`x, y, width, height`).  \n\nLeveraging a **multi-task EfficientNet** architecture, we combine:  \n- **Focal Loss** to handle class imbalance (few positive cases).  \n- **Smooth L1 Loss** for precise bounding box regression.  \n- **CLAHE preprocessing** and **geometric augmentations** to enhance model robustness.  \n\nThe system achieves **dual objectives**: screening efficiency (AUC-ROC) and interpretability (IoU for detected regions), offering a scalable solution to assist radiologists in high-volume settings.  \n\n--- ","metadata":{}},{"cell_type":"markdown","source":"# **Overall Pipeline**\n\n---\n\n### **1. Data Preparation**\n- **Input**: DICOM chest X-ray images + CSV files with bounding box annotations (`x, y, width, height`) and binary labels (`Target = 0/1`).\n- **Preprocessing**:\n  - Normalize DICOM pixel values using windowing (focus on lung tissue contrast).\n  - Apply CLAHE (Contrast Limited Adaptive Histogram Equalization) for enhanced visibility.\n  - Resize images to `512x512` (retains detail for small pneumonia opacities).\n  - Convert grayscale to 3-channel (RGB) for compatibility with pretrained CNNs.\n\n---\n\n### **2. Model Architecture**\n- **Backbone**: EfficientNet (pretrained on ImageNet) for feature extraction.\n- **Multi-Task Heads**:\n  1. **Classification Head**: Predicts pneumonia probability (`sigmoid` output).\n  2. **Bounding Box Head**: Predicts `[x, y, width, height]` for up to `3` boxes per image (with `sigmoid` to constrain coordinates to `[0, 1]`).\n- **Output Format**:  \n  - `(batch_size, 1)` for classification.  \n  - `(batch_size, 3, 4)` for bounding boxes.\n\n---\n\n### **3. Training Pipeline**\n- **Loss Function**:  \n  - **Focal Loss** (for class imbalance in pneumonia vs. normal).  \n  - **Smooth L1 Loss** (for bounding box regression).  \n  - Combined as `Total Loss = α·Focal + β·L1`.\n- **Augmentations**:  \n  - Horizontal flips, rotations, brightness/contrast adjustments (applied to both images *and* bounding boxes).\n- **Optimization**:  \n  - **AdamW** optimizer with **Cosine Annealing LR Scheduler**.  \n  - Gradient accumulation (`steps=4`) to handle large images.\n\n---\n\n### **4. Evaluation Metrics**\n- **Classification**: AUC-ROC (pneumonia vs. normal).\n- **Detection**: Mean IoU (Intersection-over-Union) for predicted vs. ground-truth boxes.\n- **Early Stopping**: Triggered if no improvement in `(AUC + IoU)/2` for `5` epochs.\n\n---\n\n### **5. Inference & Submission**\n- **Thresholding**: Only boxes with confidence `> 0.5` are kept.\n- **Submission Format**:  \n  - Single row per image with format:  \n    `patientId, \"confidence x y w h confidence x y w h ...\"`  \n    (or empty string if no pneumonia detected).\n- **Postprocessing**:  \n  - Filter tiny boxes (`width/height < 0.01` of image size).  \n  - Clip coordinates to `[0, 1]`.\n\n---\n\n### **Key Advantages**\n1. **Efficiency**: Uses lightweight EfficientNet backbone for fast inference.\n2. **Robustness**: Focal Loss handles class imbalance; augmentations improve generalization.\n3. **Precision**: Box coordinates are normalized and trained end-to-end with the classifier.\n\nThis pipeline balances **classification accuracy** and **localization precision**, critical for medical detection tasks.\n\n\n---","metadata":{}},{"cell_type":"code","source":"# ===== Standard Library =====\nimport os\nimport gc\nimport random\nimport warnings\nimport functools\nfrom pathlib import Path\nfrom typing import List, Tuple, Dict, Optional, Union\nfrom dataclasses import dataclass\n\nimport numpy as np\nimport pandas as pd\n\nimport cv2\nimport pydicom\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import (\n    Compose, RandomBrightnessContrast, Blur, CLAHE,\n    HorizontalFlip, Rotate, ShiftScaleRotate,\n    ElasticTransform, GridDistortion,\n    RandomGamma, GaussNoise, ISONoise,\n    CoarseDropout, Normalize\n)\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torchvision.ops import box_iou\n\nimport timm\nfrom sklearn.model_selection import GroupKFold, StratifiedKFold\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\n\nwarnings.filterwarnings(\"ignore\")\n\n# Set random seeds for reproducibility\nSEED = 42\nnp.random.seed(SEED)\nrandom.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass Config:\n    \"\"\"Configuration class for RSNA Pneumonia Detection Challenge\"\"\"\n    \n    # Data paths\n    DATA_DIR: str = \"/kaggle/input/rsna-pneumonia-detection-challenge\"\n    TRAIN_LABELS: str = \"stage_2_train_labels.csv\"\n    CLASS_INFO: str = \"stage_2_detailed_class_info.csv\"\n    TRAIN_IMAGES: str = \"stage_2_train_images\"\n    TEST_IMAGES: str = \"stage_2_test_images\"\n    OUTPUT_DIR: str = \"/kaggle/working\"\n    \n    # Model parameters\n    MODEL_BACKBONE: str = \"efficientnet_b0\"\n    IMG_SIZE: int = 512  # Larger size for better detection\n    NUM_CLASSES: int = 1  # Pneumonia or not\n    HIDDEN_DIM: int = 512\n    \n    # Detection parameters\n    MAX_BOXES: int = 3  # Maximum number of boxes to predict per image\n    BOX_THRESHOLD: float = 0.5  # Confidence threshold for box predictions\n    \n    # Training parameters\n    BATCH_SIZE: int = 8\n    EPOCHS: int = 3\n    LEARNING_RATE: float = 3e-4\n    WEIGHT_DECAY: float = 1e-4\n    ACCUMULATION_STEPS: int = 4\n    \n    # Cross-validation\n    NUM_FOLDS: int = 5\n    FOLD: int = 0\n    \n    # Augmentation and preprocessing\n    USE_STRONG_AUGMENTATION: bool = True\n    USE_CLAHE: bool = True\n    \n    # Loss function weights\n    CLASSIFICATION_WEIGHT: float = 1.0\n    BOX_REG_WEIGHT: float = 1.0\n    FOCAL_ALPHA: float = 0.25\n    FOCAL_GAMMA: float = 2.0\n    \n    # Training settings\n    PATIENCE: int = 5\n    CACHE_SIZE: int = 1000\n    NUM_WORKERS: int = 4\n\nconfig = Config()\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    \"\"\"Focal Loss for addressing class imbalance\"\"\"\n    \n    def __init__(self, alpha: float = 0.25, gamma: float = 2.0, reduction: str = \"mean\"):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n        # Apply sigmoid to get probabilities\n        probs = torch.sigmoid(logits)\n        \n        # Compute binary cross entropy\n        bce_loss = F.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n        \n        # Compute focal weight\n        p_t = probs * targets + (1 - probs) * (1 - targets)\n        focal_weight = (1 - p_t) ** self.gamma\n        \n        # Apply alpha weighting\n        alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)\n        \n        # Final focal loss\n        focal_loss = alpha_t * focal_weight * bce_loss\n        \n        if self.reduction == \"mean\":\n            return focal_loss.mean()\n        elif self.reduction == \"sum\":\n            return focal_loss.sum()\n        else:\n            return focal_loss\n\nclass SmoothL1Loss(nn.Module):\n    \"\"\"Smooth L1 Loss for bounding box regression\"\"\"\n    \n    def __init__(self, beta: float = 1.0):\n        super().__init__()\n        self.beta = beta\n    \n    def forward(self, pred_boxes: torch.Tensor, target_boxes: torch.Tensor) -> torch.Tensor:\n        # Convert from (x, y, w, h) to (x1, y1, x2, y2)\n        pred_boxes = self._xywh_to_x1y1x2y2(pred_boxes)\n        target_boxes = self._xywh_to_x1y1x2y2(target_boxes)\n        \n        diff = torch.abs(pred_boxes - target_boxes)\n        loss = torch.where(diff < self.beta, \n                          0.5 * diff ** 2 / self.beta,\n                          diff - 0.5 * self.beta)\n        \n        return loss.mean()\n    \n    def _xywh_to_x1y1x2y2(self, boxes: torch.Tensor) -> torch.Tensor:\n        \"\"\"Convert boxes from (x, y, w, h) to (x1, y1, x2, y2) format\"\"\"\n        x1 = boxes[:, 0]\n        y1 = boxes[:, 1]\n        x2 = x1 + boxes[:, 2]\n        y2 = y1 + boxes[:, 3]\n        return torch.stack([x1, y1, x2, y2], dim=1)\n\n\nclass FocalLoss(nn.Module):\n    \"\"\"Focal Loss for addressing class imbalance\"\"\"\n    \n    def __init__(self, alpha: float = 0.25, gamma: float = 2.0, reduction: str = \"mean\"):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:\n        # Apply sigmoid to get probabilities\n        probs = torch.sigmoid(logits)\n        # Compute binary cross entropy\n        bce_loss = F.binary_cross_entropy_with_logits(logits, targets, reduction=\"none\")\n        # Compute focal weight\n        p_t = probs * targets + (1 - probs) * (1 - targets)\n        focal_weight = (1 - p_t) ** self.gamma\n        # Apply alpha weighting\n        alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)\n        # Final focal loss\n        focal_loss = alpha_t * focal_weight * bce_loss\n        \n        if self.reduction == \"mean\":\n            return focal_loss.mean()\n        elif self.reduction == \"sum\":\n            return focal_loss.sum()\n        else:\n            return focal_loss\n\nclass SmoothL1Loss(nn.Module):\n    def __init__(self, beta: float = 1.0):\n        super().__init__()\n        self.beta = beta\n    \n    def forward(self, pred_boxes: torch.Tensor, target_boxes: torch.Tensor) -> torch.Tensor:\n        # Convert from (x, y, w, h) to (x1, y1, x2, y2)\n        pred_boxes = self._xywh_to_x1y1x2y2(pred_boxes)\n        target_boxes = self._xywh_to_x1y1x2y2(target_boxes)\n        \n        diff = torch.abs(pred_boxes - target_boxes)\n        loss = torch.where(diff < self.beta, \n                          0.5 * diff ** 2 / self.beta,\n                          diff - 0.5 * self.beta)\n        \n        return loss.mean()\n    \n    def _xywh_to_x1y1x2y2(self, boxes: torch.Tensor) -> torch.Tensor:\n        \"\"\"Convert boxes from (x, y, w, h) to (x1, y1, x2, y2) format\"\"\"\n        x1 = boxes[:, 0]\n        y1 = boxes[:, 1]\n        x2 = x1 + boxes[:, 2]\n        y2 = y1 + boxes[:, 3]\n        return torch.stack([x1, y1, x2, y2], dim=1)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiTaskLoss(nn.Module):\n    def __init__(self, classification_weight=1.0, box_reg_weight=1.0, focal_alpha=0.25, focal_gamma=2.0):\n        super().__init__()\n        self.classification_weight = classification_weight\n        self.box_reg_weight = box_reg_weight\n        self.focal_loss = FocalLoss(alpha=focal_alpha, gamma=focal_gamma)\n        self.box_loss = nn.SmoothL1Loss(reduction='mean')\n    \n    def forward(self, pred_logits, pred_boxes, target_labels, target_boxes):\n        #print(\"\\n\" + \"=\"*60)\n        #print(\"DEBUG MULTITASK LOSS\")\n        #print(\"=\"*60)\n        #print(f\"pred_logits shape: {pred_logits.shape}\")\n        #print(f\"pred_boxes shape: {pred_boxes.shape}\")\n        #print(f\"target_labels shape: {target_labels.shape}\")\n        #print(f\"target_boxes shape: {target_boxes.shape}\")\n        \n        # DETAILED SHAPE ANALYSIS AND FIXING\n        batch_size = pred_logits.shape[0]\n        \n        #print(f\"Original shapes:\")\n        #print(f\"  pred_logits: {pred_logits.shape}\")\n        #print(f\"  target_labels: {target_labels.shape}\")\n        \n        # Handle different scenarios\n        if pred_logits.shape == torch.Size([8, 3, 2]) and target_labels.shape == torch.Size([8, 1]):\n            # Model outputs [batch, 3 boxes, 2 classes], targets are [batch, 1 label per image]\n            # Option 1: Use only the first box prediction for image-level classification\n            pred_logits_reshaped = pred_logits[:, 0, 1].unsqueeze(1)  # [8, 1] - first box, positive class\n            #print(\"Using first box, positive class for image classification\")\n            \n        elif pred_logits.shape == torch.Size([8, 3, 2]) and target_labels.shape == torch.Size([8, 3]):\n            # Model outputs [batch, 3 boxes, 2 classes], targets are [batch, 3 box labels]\n            # Use positive class logits for each box\n            pred_logits_reshaped = pred_logits[:, :, 1]  # [8, 3] - positive class for all boxes\n            #print(\"Using positive class logits for all boxes\")\n            \n        else:\n            # Generic handling for other cases\n            #print(f\"Generic handling for shapes: {pred_logits.shape} vs {target_labels.shape}\")\n            \n            # Try to match the number of elements\n            pred_elements = pred_logits.numel() // batch_size  # Elements per sample\n            target_elements = target_labels.numel() // batch_size  # Elements per sample\n            \n            #print(f\"Elements per sample - pred: {pred_elements}, target: {target_elements}\")\n            \n            if pred_elements > target_elements:\n                # More predictions than targets - take subset of predictions\n                if pred_logits.dim() == 3:\n                    if target_labels.shape[1] == 1:\n                        # Take first box, last class (assuming positive class)\n                        pred_logits_reshaped = pred_logits[:, 0, -1].unsqueeze(1)\n                    else:\n                        # Take first N boxes, last class\n                        pred_logits_reshaped = pred_logits[:, :target_labels.shape[1], -1]\n                else:\n                    # Flatten and take first N elements\n                    pred_flat = pred_logits.view(batch_size, -1)\n                    pred_logits_reshaped = pred_flat[:, :target_elements].view_as(target_labels)\n            else:\n                # Equal or fewer predictions than targets\n                pred_logits_reshaped = pred_logits.view_as(target_labels)\n        \n        # Ensure both tensors have the same shape\n        if pred_logits_reshaped.shape != target_labels.shape:\n            #print(f\"Still mismatched after reshaping!\")\n            #print(f\"  pred_logits_reshaped: {pred_logits_reshaped.shape}\")\n            #print(f\"  target_labels: {target_labels.shape}\")\n            \n            # Last resort: force them to match\n            min_elements = min(pred_logits_reshaped.numel(), target_labels.numel())\n            pred_logits_reshaped = pred_logits_reshaped.view(-1)[:min_elements].view(batch_size, -1)\n            target_labels_reshaped = target_labels.view(-1)[:min_elements].view(batch_size, -1)\n        else:\n            target_labels_reshaped = target_labels\n            \n        #print(f\"Final shapes for loss calculation:\")\n        #print(f\"pred_logits_reshaped shape: {pred_logits_reshaped.shape}\")\n        #print(f\"target_labels_reshaped shape: {target_labels_reshaped.shape}\")\n        \n        # Classification loss\n        try:\n            cls_loss = F.binary_cross_entropy_with_logits(\n                pred_logits_reshaped.float(), \n                target_labels_reshaped.float()\n            )\n            #print(f\"Classification loss: {cls_loss.item():.6f}\")\n        except Exception as e:\n            #print(f\"Classification loss error: {e}\")\n            #print(f\"Final attempt with flattened tensors...\")\n            \n            # Absolute last resort: flatten everything and truncate to match\n            pred_flat = pred_logits_reshaped.view(-1)\n            target_flat = target_labels_reshaped.view(-1)\n            min_size = min(len(pred_flat), len(target_flat))\n            \n            cls_loss = F.binary_cross_entropy_with_logits(\n                pred_flat[:min_size].float(),\n                target_flat[:min_size].float()\n            )\n            #print(f\"Fallback classification loss: {cls_loss.item():.6f}\")\n        \n        # Box regression loss (skip for now since you have no boxes)\n        pos_mask = target_labels_reshaped > 0.5\n        num_positive = pos_mask.sum()\n        \n        if num_positive > 0:\n            #print(f\"Positive samples: {num_positive}\")\n            # Box loss computation...\n            box_loss = torch.tensor(0.0, device=pred_logits.device, requires_grad=True)\n        else:\n            #print(\"No positive samples - box loss = 0\")\n            box_loss = torch.tensor(0.0, device=pred_logits.device, requires_grad=True)\n        \n        total_loss = self.classification_weight * cls_loss + self.box_reg_weight * box_loss\n        \n        #print(f\"Total loss: {total_loss.item():.6f}\")\n        #print(\"=\"*60)\n        \n        return total_loss, cls_loss, box_loss","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DetectionLoss(nn.Module):\n    def __init__(self, num_classes, alpha=0.25, gamma=2.0, box_loss_weight=1.0):\n        super().__init__()\n        self.num_classes = num_classes\n        self.cls_loss = FocalLoss(alpha=alpha, gamma=gamma, num_classes=num_classes)\n        self.box_reg_loss = nn.SmoothL1Loss(reduction='mean', beta=1.0)\n        self.box_loss_weight = box_loss_weight\n\n    def forward(self, pred_logits, pred_boxes, target_labels, target_boxes):\n        # Classification loss\n        cls_loss = self.cls_loss(pred_logits, target_labels)\n        \n        # Create positive mask (where target_labels > 0)\n        pos_mask = target_labels > 0\n        \n        # DEBUG STATEMENTS - Add these here!\n        #print(f\"DEBUG - target_labels shape: {target_labels.shape}\")\n        #print(f\"DEBUG - target_boxes shape: {target_boxes.shape}\")\n        #print(f\"DEBUG - pos_mask shape: {pos_mask.shape}\")\n        #print(f\"DEBUG - pos_mask sum: {pos_mask.sum()}\")\n        #print(f\"DEBUG - target_labels content:\\n{target_labels}\")\n        #print(f\"DEBUG - pos_mask content:\\n{pos_mask}\")\n        \n        # Check if we have any positive samples\n        num_pos = pos_mask.sum()\n        \n        if num_pos > 0:\n            # Extract positive predictions and targets\n            pos_pred_boxes = pred_boxes[pos_mask]\n            pos_target_boxes = target_boxes[pos_mask]\n            \n            # More debug info\n            #print(f\"DEBUG - pos_pred_boxes shape: {pos_pred_boxes.shape}\")\n            #print(f\"DEBUG - pos_target_boxes shape: {pos_target_boxes.shape}\")\n            #print(f\"DEBUG - pos_pred_boxes content:\\n{pos_pred_boxes}\")\n            #print(f\"DEBUG - pos_target_boxes content:\\n{pos_target_boxes}\")\n            \n            # Safety check before computing loss\n            if pos_pred_boxes.shape[0] != pos_target_boxes.shape[0]:\n                #print(f\"ERROR - Shape mismatch!\")\n                #print(f\"  pos_pred_boxes: {pos_pred_boxes.shape}\")\n                #print(f\"  pos_target_boxes: {pos_target_boxes.shape}\")\n                # Skip box loss if shapes don't match\n                return cls_loss\n            \n            # Compute box regression loss\n            box_loss = self.box_reg_loss(pos_pred_boxes, pos_target_boxes)\n            \n        else:\n            #print(\"DEBUG - No positive samples, box_loss = 0\")\n            box_loss = torch.tensor(0.0, device=pred_logits.device, requires_grad=True)\n        \n        # Combine losses\n        total_loss = cls_loss + self.box_loss_weight * box_loss\n        \n        #print(f\"DEBUG - cls_loss: {cls_loss.item():.4f}\")\n        #print(f\"DEBUG - box_loss: {box_loss.item():.4f}\")\n        #print(f\"DEBUG - total_loss: {total_loss.item():.4f}\")\n        #print(\"-\" * 50)\n        \n        return total_loss","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DICOMPreprocessor:\n    \"\"\"DICOM image preprocessing utilities for chest X-rays\"\"\"\n    \n    @staticmethod\n    def preprocess_dicom(dicom_path: Union[str, Path]) -> np.ndarray:\n        \"\"\"Load and preprocess a DICOM file\"\"\"\n        dicom = pydicom.dcmread(str(dicom_path))\n        \n        # Handle different photometric interpretations\n        if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n            image = np.invert(dicom.pixel_array)\n        else:\n            image = dicom.pixel_array\n            \n        # Apply windowing if available\n        if hasattr(dicom, 'WindowCenter') and hasattr(dicom, 'WindowWidth'):\n            center = float(dicom.WindowCenter)\n            width = float(dicom.WindowWidth)\n            image = DICOMPreprocessor.apply_windowing(image, center, width)\n        else:\n            # Default windowing for chest X-rays\n            image = DICOMPreprocessor.apply_windowing(image, 40, 400)\n        \n        # Convert to 8-bit\n        image = (image - image.min()) / (image.max() - image.min() + 1e-7)\n        image = (image * 255).astype(np.uint8)\n        \n        return image\n    \n    @staticmethod\n    def apply_windowing(image: np.ndarray, center: float, width: float) -> np.ndarray:\n        \"\"\"Apply windowing to DICOM image\"\"\"\n        min_val = center - width / 2.0\n        max_val = center + width / 2.0\n        windowed = np.clip(image, min_val, max_val)\n        return windowed\n    \n    @staticmethod\n    def apply_clahe(image: np.ndarray) -> np.ndarray:\n        \"\"\"Apply Contrast Limited Adaptive Histogram Equalization\"\"\"\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        enhanced = clahe.apply(image)\n        return enhanced","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class PneumoniaDataset(Dataset):\n    \"\"\"Dataset for RSNA Pneumonia Detection Challenge\"\"\"\n    \n    def __init__(self, \n                 dataframe: pd.DataFrame,\n                 image_dir: str,\n                 transforms=None,\n                 is_training: bool = True):\n        \n        self.df = dataframe.reset_index(drop=True)\n        self.image_dir = Path(image_dir)\n        self.transforms = transforms\n        self.is_training = is_training\n        self.preprocessor = DICOMPreprocessor()\n        \n        # Simple caching mechanism\n        self.cache = {}\n        self.cache_keys = []\n        self.max_cache_size = config.CACHE_SIZE\n    \n    def __len__(self) -> int:\n        return len(self.df)\n    \n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, Dict[str, torch.Tensor]]:\n        # Check cache first\n        if idx in self.cache:\n            return self.cache[idx]\n        \n        row = self.df.iloc[idx]\n        patient_id = row['patientId']\n        \n        # Load DICOM image\n        dicom_path = self.image_dir / f\"{patient_id}.dcm\"\n        image = self.preprocessor.preprocess_dicom(dicom_path)\n        \n        # Apply CLAHE if enabled\n        if config.USE_CLAHE:\n            image = self.preprocessor.apply_clahe(image)\n        \n        # Convert to 3-channel (repeat grayscale)\n        image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n        \n        # Get target information\n        target = self._prepare_target(row)\n        \n        # Apply augmentations\n        if self.transforms:\n            # Convert boxes to numpy array for albumentations\n            boxes_np = target['boxes'].numpy() if target['boxes'].numel() > 0 else np.zeros((0, 4), dtype=np.float32)\n            \n            augmented = self.transforms(\n                image=image, \n                bboxes=boxes_np\n            )\n            image = augmented[\"image\"]\n            \n            # Update boxes after augmentation\n            if len(augmented[\"bboxes\"]) > 0:\n                target['boxes'] = torch.tensor(augmented[\"bboxes\"], dtype=torch.float32)\n            else:\n                target['boxes'] = torch.zeros((0, 4), dtype=torch.float32)\n        \n        # Convert boxes to (x, y, w, h) format\n        if len(target['boxes']) > 0:\n            boxes = target['boxes']\n            # Ensure boxes are within image bounds\n            boxes[:, 0] = torch.clamp(boxes[:, 0], 0, 1)  # x\n            boxes[:, 1] = torch.clamp(boxes[:, 1], 0, 1)  # y\n            boxes[:, 2] = torch.clamp(boxes[:, 2] - boxes[:, 0], 0, 1)  # width\n            boxes[:, 3] = torch.clamp(boxes[:, 3] - boxes[:, 1], 0, 1)  # height\n            target['boxes'] = boxes\n        \n        result = (image, target)\n        \n        # Cache the result\n        self._add_to_cache(idx, result)\n        \n        return result\n    \n    def _prepare_target(self, row: pd.Series) -> Dict[str, torch.Tensor]:\n        \"\"\"Prepare target dictionary\"\"\"\n        target = {\n            'labels': torch.tensor([row['Target']], dtype=torch.float32),\n            'boxes': torch.zeros((0, 4), dtype=torch.float32)  # Empty tensor for no boxes\n        }\n        \n        # If pneumonia is present, add bounding boxes\n        if row['Target'] == 1 and not pd.isna(row['x']):\n            # Original box coordinates (normalized)\n            x = row['x'] / config.IMG_SIZE  # Normalize by image size\n            y = row['y'] / config.IMG_SIZE\n            width = row['width'] / config.IMG_SIZE\n            height = row['height'] / config.IMG_SIZE\n            \n            # Convert to (x1, y1, x2, y2) format\n            box = torch.tensor([\n                [x, y, x + width, y + height]\n            ], dtype=torch.float32)\n            \n            target['boxes'] = box\n        \n        return target\n    \n    def _add_to_cache(self, idx: int, item: Tuple):\n        \"\"\"Add item to cache with LRU eviction policy\"\"\"\n        if len(self.cache) >= self.max_cache_size:\n            oldest_key = self.cache_keys.pop(0)\n            del self.cache[oldest_key]\n        \n        self.cache[idx] = item\n        self.cache_keys.append(idx)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_augmentations(is_training: bool = True):\n    \"\"\"Create augmentation pipeline for detection\"\"\"\n    if is_training and config.USE_STRONG_AUGMENTATION:\n        transform = Compose([\n            HorizontalFlip(p=0.5),\n            ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=5, p=0.5),\n            RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.4),\n            RandomGamma(gamma_limit=(80, 120), p=0.3),\n            GaussNoise(var_limit=(10, 50), p=0.3),\n            CoarseDropout(max_holes=8, max_height=32, max_width=32, p=0.3),\n            Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ], bbox_params={\n            'format': 'pascal_voc', \n            'min_area': 0.0, \n            'min_visibility': 0.0,\n            'label_fields': []\n        })\n    else:\n        transform = Compose([\n            Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ], bbox_params={\n            'format': 'pascal_voc',\n            'label_fields': []\n        })\n    \n    return transform","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# Key fixes for the pneumonia detection code\n\n# Fix 1: Update the model configuration to ensure consistent class dimensions\nclass PneumoniaDetectionModel(nn.Module):\n    def __init__(self, backbone_name='efficientnet_b0', num_classes=2, hidden_dim=256, max_boxes=3, dropout_rate=0.3):\n        super().__init__()\n        self.num_classes = num_classes  # Make sure this is 2 for binary classification\n        self.max_boxes = max_boxes\n        self.hidden_dim = hidden_dim\n        \n        # Load pre-trained backbone\n        self.backbone = timm.create_model(\n            backbone_name,\n            pretrained=True,\n            num_classes=0,  # Remove classifier\n            global_pool=''  # Remove global pooling\n        )\n        \n        # Get the number of features from backbone\n        with torch.no_grad():\n            dummy_input = torch.randn(1, 3, 224, 224)\n            backbone_output = self.backbone(dummy_input)\n            backbone_features = backbone_output.shape[1]\n        \n        print(f\"Backbone features: {backbone_features}\")\n        \n        # Shared feature processing\n        self.shared_conv = nn.Sequential(\n            nn.Conv2d(backbone_features, hidden_dim, kernel_size=3, padding=1),\n            nn.BatchNorm2d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout2d(dropout_rate),\n            nn.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1),\n            nn.BatchNorm2d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d((1, 1))\n        )\n        \n        # Classification head - FIXED: Ensure proper output dimensions\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout_rate),\n            nn.Linear(hidden_dim // 2, max_boxes * num_classes)  # This should be max_boxes * 2\n        )\n        \n        # Box regression head\n        self.box_regressor = nn.Sequential(\n            nn.Flatten(), \n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout_rate),\n            nn.Linear(hidden_dim // 2, max_boxes * 4)\n        )\n        \n        print(f\"Model initialized:\")\n        print(f\"  - max_boxes: {max_boxes}\")\n        print(f\"  - num_classes: {num_classes}\")\n        print(f\"  - Classification output size: {max_boxes * num_classes}\")\n        print(f\"  - Box regression output size: {max_boxes * 4}\")\n    \n    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        batch_size = x.size(0)\n        \n        # Extract backbone features\n        backbone_features = self.backbone(x)\n        \n        # Shared processing\n        shared_features = self.shared_conv(backbone_features)\n        \n        # Classification and box regression\n        class_features = self.classifier(shared_features)\n        box_features = self.box_regressor(shared_features)\n        \n        # Reshape to proper dimensions\n        class_logits = class_features.view(batch_size, self.max_boxes, self.num_classes)\n        box_coords = box_features.view(batch_size, self.max_boxes, 4)\n        \n        # Apply sigmoid to normalize coordinates to [0, 1]\n        box_coords = torch.sigmoid(box_coords)\n        \n        return class_logits, box_coords","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class PneumoniaDetectionModel(nn.Module):\n    def __init__(self, backbone_name='efficientnet_b0', num_classes=2, hidden_dim=256, max_boxes=3, dropout_rate=0.3):\n        super().__init__()\n        self.num_classes = num_classes  # IMPORTANT: Store this\n        self.max_boxes = max_boxes\n        self.hidden_dim = hidden_dim\n        \n        # Load pre-trained backbone\n        self.backbone = timm.create_model(\n            backbone_name,\n            pretrained=True,\n            num_classes=0,  # Remove classifier\n            global_pool=''  # Remove global pooling\n        )\n        \n        # Get the number of features from backbone\n        # This gets the feature size after backbone\n        with torch.no_grad():\n            dummy_input = torch.randn(1, 3, 224, 224)\n            backbone_output = self.backbone(dummy_input)\n            backbone_features = backbone_output.shape[1]  # Number of channels\n        \n        print(f\"Backbone features: {backbone_features}\")\n        \n        # Shared feature processing\n        self.shared_conv = nn.Sequential(\n            nn.Conv2d(backbone_features, hidden_dim, kernel_size=3, padding=1),\n            nn.BatchNorm2d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.Dropout2d(dropout_rate),\n            nn.Conv2d(hidden_dim, hidden_dim, kernel_size=3, padding=1),\n            nn.BatchNorm2d(hidden_dim),\n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d((1, 1))  # Global average pooling\n        )\n        \n        # Classification head\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout_rate),\n            nn.Linear(hidden_dim // 2, max_boxes * num_classes)\n        )\n        \n        # Box regression head\n        self.box_regressor = nn.Sequential(\n            nn.Flatten(), \n            nn.Linear(hidden_dim, hidden_dim // 2),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout_rate),\n            nn.Linear(hidden_dim // 2, max_boxes * 4)\n        )\n        \n        #print(f\"Model initialized:\")\n        #print(f\"  - max_boxes: {max_boxes}\")\n        #print(f\"  - num_classes: {num_classes}\")\n        #print(f\"  - Classification output size: {max_boxes * num_classes}\")\n        #print(f\"  - Box regression output size: {max_boxes * 4}\")\n    \n    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Forward pass of the object detection model.\n        \"\"\"\n        batch_size = x.size(0)\n        \n        # Extract backbone features\n        backbone_features = self.backbone(x)\n        \n        # Shared processing\n        shared_features = self.shared_conv(backbone_features)\n        \n        # Classification and box regression\n        class_features = self.classifier(shared_features)\n        box_features = self.box_regressor(shared_features)\n        \n        # Reshape to proper dimensions\n        class_logits = class_features.view(batch_size, self.max_boxes, self.num_classes)\n        box_coords = box_features.view(batch_size, self.max_boxes, 4)\n        \n        # Apply sigmoid to normalize coordinates to [0, 1]\n        box_coords = torch.sigmoid(box_coords)\n        \n        # DEBUG: Print shapes for first batch\n        if not hasattr(self, '_debug_printed'):\n            #print(f\"DEBUG MODEL OUTPUT:\")\n            #print(f\"  class_logits shape: {class_logits.shape}\")\n            #print(f\"  box_coords shape: {box_coords.shape}\")\n            self._debug_printed = True\n        \n        return class_logits, box_coords","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def collate_fn(batch):\n    \"\"\"Custom collate function to handle our dataset items\"\"\"\n    images, targets = zip(*batch)\n    \n    # Stack images\n    images = torch.stack(images, dim=0)\n    \n    # Process targets\n    target_labels = torch.stack([t['labels'] for t in targets])\n    target_boxes = [t['boxes'] for t in targets]\n    \n    return images, {'labels': target_labels, 'boxes': target_boxes}","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Trainer","metadata":{}},{"cell_type":"code","source":"class Trainer:\n    def __init__(self, \n                 model: nn.Module,\n                 criterion: nn.Module,\n                 optimizer: torch.optim.Optimizer,\n                 scheduler: torch.optim.lr_scheduler._LRScheduler,\n                 device: torch.device):\n        \n        self.model = model.to(device)\n        self.criterion = criterion\n        self.optimizer = optimizer\n        self.scheduler = scheduler\n        self.device = device\n        self.scaler = torch.cuda.amp.GradScaler()\n\n\n    def train_epoch(self, train_loader: DataLoader) -> Dict[str, float]:\n        \"\"\"Train for one epoch\"\"\"\n        self.model.train()\n        total_loss = 0.0\n        total_cls_loss = 0.0\n        total_box_loss = 0.0\n        num_batches = 0\n        \n        for batch_idx, (images, targets) in enumerate(tqdm(train_loader, desc=\"Training\")):\n            \n            # Move to device\n            images = images.to(self.device, non_blocking=True)\n            target_labels = targets['labels'].to(self.device)\n            \n            # DEBUG: Add debug info for first batch\n            if batch_idx == 0:\n                #print(f\"\\n=== FIRST BATCH DEBUG INFO ===\")\n                #print(f\"Batch size: {len(images)}\")\n                #print(f\"Images shape: {images.shape}\")\n                #print(f\"target_labels shape: {target_labels.shape}\")\n                #print(f\"target_labels content: {target_labels}\")\n                #print(f\"Number of box lists in targets['boxes']: {len(targets['boxes'])}\")\n                \n                for i, boxes in enumerate(targets['boxes']):\n                    print(f\"  Sample {i}: {len(boxes)} boxes, shape: {boxes.shape if len(boxes) > 0 else 'empty'}\")\n            \n            # Convert list of boxes to padded tensor on device\n            max_boxes = max(len(boxes) for boxes in targets['boxes']) if targets['boxes'] else 0\n            \n            # Handle case where there are no boxes at all\n            if max_boxes == 0:\n                print(f\"WARNING: Batch {batch_idx} has no boxes at all!\")\n                max_boxes = 1  # Set minimum to avoid empty tensor\n            \n            padded_boxes = torch.zeros(len(images), max_boxes, 4, device=self.device)\n            \n            # Fill the padded tensor\n            for i, boxes in enumerate(targets['boxes']):\n                if len(boxes) > 0:\n                    padded_boxes[i, :len(boxes)] = boxes.to(self.device)\n            \n            # Create corresponding labels for each box position\n            # This is likely where your issue is - you need labels that match the box positions\n            padded_labels = torch.zeros(len(images), max_boxes, device=self.device)\n            \n            # Fill padded labels based on your target_labels\n            for i in range(len(images)):\n                if len(targets['boxes'][i]) > 0:\n                    # If you have per-image labels, replicate them for each box\n                    # Adjust this logic based on your actual label format\n                    if target_labels.dim() == 1:  # Per-image labels\n                        # Set all box positions to the image label\n                        padded_labels[i, :len(targets['boxes'][i])] = target_labels[i]\n                    elif target_labels.dim() == 2:  # Already per-box labels\n                        # Use existing box labels\n                        padded_labels[i, :target_labels.shape[1]] = target_labels[i]\n            \n            # DEBUG: Print padded tensor info for first batch\n            if batch_idx == 0:\n                print(f\"max_boxes: {max_boxes}\")\n                #print(f\"padded_boxes shape: {padded_boxes.shape}\")\n                #print(f\"padded_labels shape: {padded_labels.shape}\")\n                #print(f\"padded_labels content:\\n{padded_labels}\")\n                #print(f\"Non-zero padded_labels: {(padded_labels > 0).sum().item()}\")\n                print(\"=== END FIRST BATCH DEBUG ===\\n\")\n            \n            # Forward pass with mixed precision\n            with torch.cuda.amp.autocast():\n                pred_logits, pred_boxes = self.model(images)\n                \n                # DEBUG: Print model outputs for first batch\n                if batch_idx == 0:\n                    print(f\"Model outputs:\")\n                    #print(f\"  pred_logits shape: {pred_logits.shape}\")\n                    #print(f\"  pred_boxes shape: {pred_boxes.shape}\")\n                \n                # Calculate loss - now using padded_labels instead of target_labels\n                try:\n                    # Try to get individual losses (if your criterion returns them)\n                    result = self.criterion(\n                        pred_logits,\n                        pred_boxes,\n                        padded_labels,  # Use padded_labels instead of target_labels\n                        padded_boxes\n                    )\n                    \n                    # Handle different return formats\n                    if isinstance(result, tuple) and len(result) == 3:\n                        total_loss_batch, cls_loss_batch, box_loss_batch = result\n                    elif isinstance(result, tuple) and len(result) == 2:\n                        total_loss_batch, cls_loss_batch = result\n                        box_loss_batch = torch.tensor(0.0)\n                    else:\n                        # Single loss returned\n                        total_loss_batch = result\n                        cls_loss_batch = torch.tensor(0.0)\n                        box_loss_batch = torch.tensor(0.0)\n                        \n                except Exception as e:\n                    print(f\"Error in loss calculation: {e}\")\n                    #print(f\"pred_logits shape: {pred_logits.shape}\")\n                    #print(f\"pred_boxes shape: {pred_boxes.shape}\")\n                    #print(f\"padded_labels shape: {padded_labels.shape}\")\n                    #print(f\"padded_boxes shape: {padded_boxes.shape}\")\n                    raise e\n                \n                # Scale loss for gradient accumulation\n                scaled_loss = total_loss_batch / config.ACCUMULATION_STEPS\n            \n            # Backward pass\n            self.scaler.scale(scaled_loss).backward()\n            \n            # Accumulate losses\n            total_loss += total_loss_batch.item()\n            \n            if isinstance(cls_loss_batch, torch.Tensor):\n                total_cls_loss += cls_loss_batch.item()\n            else:\n                total_cls_loss += cls_loss_batch\n                \n            if isinstance(box_loss_batch, torch.Tensor):\n                total_box_loss += box_loss_batch.item()\n            else:\n                total_box_loss += box_loss_batch\n                \n            num_batches += 1\n            \n            # Update weights\n            if (batch_idx + 1) % config.ACCUMULATION_STEPS == 0 or batch_idx == len(train_loader) - 1:\n                self.scaler.step(self.optimizer)\n                self.scaler.update()\n                self.optimizer.zero_grad()\n            \n            # Stop after first batch for debugging (remove this line after debugging)\n            if batch_idx == 0:\n                print(\"Stopping after first batch for debugging purposes\")\n                break  # Remove this line once debugging is complete\n        \n        metrics = {\n            'total_loss': total_loss / max(num_batches, 1),\n            'cls_loss': total_cls_loss / max(num_batches, 1),\n            'box_loss': total_box_loss / max(num_batches, 1)\n        }\n        \n        return metrics\n\n    def validate_epoch(self, val_loader: DataLoader) -> Dict[str, float]:\n        \"\"\"Validate for one epoch\"\"\"\n        self.model.eval()\n        total_loss = 0.0\n        total_cls_loss = 0.0\n        total_box_loss = 0.0\n        num_batches = 0\n        \n        with torch.no_grad():\n            for images, targets in tqdm(val_loader, desc=\"Validation\"):\n                images = images.to(self.device, non_blocking=True)\n                target_labels = targets['labels'].to(self.device)\n                \n                # Convert list of boxes to padded tensor\n                max_boxes = max(len(boxes) for boxes in targets['boxes'])\n                padded_boxes = torch.zeros(len(images), max_boxes, 4, device=self.device)\n                for i, boxes in enumerate(targets['boxes']):\n                    if len(boxes) > 0:\n                        padded_boxes[i, :len(boxes)] = boxes.to(self.device)\n                \n                with torch.cuda.amp.autocast():\n                    pred_logits, pred_boxes = self.model(images)\n                    \n                    total_loss_batch, cls_loss_batch, box_loss_batch = self.criterion(\n                        pred_logits,\n                        pred_boxes,\n                        target_labels,\n                        padded_boxes\n                    )\n                \n                total_loss += total_loss_batch.item()\n                total_cls_loss += cls_loss_batch.item()\n                total_box_loss += box_loss_batch.item()\n                num_batches += 1\n        \n        metrics = {\n            'total_loss': total_loss / max(num_batches, 1),\n            'cls_loss': total_cls_loss / max(num_batches, 1),\n            'box_loss': total_box_loss / max(num_batches, 1)\n        }\n        \n        return metrics","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_data_splits(df: pd.DataFrame) -> Tuple[pd.DataFrame, pd.DataFrame]:\n    \"\"\"Create train/validation splits with stratification\"\"\"\n    # Stratified K-Fold split based on pneumonia presence\n    skf = StratifiedKFold(n_splits=config.NUM_FOLDS, shuffle=True, random_state=SEED)\n    splits = list(skf.split(df, df['Target']))\n    \n    train_idx, val_idx = splits[config.FOLD]\n    train_df = df.iloc[train_idx].reset_index(drop=True)\n    val_df = df.iloc[val_idx].reset_index(drop=True)\n    \n    return train_df, val_df\n\ndef prepare_data() -> Tuple[pd.DataFrame, pd.DataFrame]:\n    \"\"\"Prepare and merge the training data\"\"\"\n    # Load labels\n    labels_df = pd.read_csv(os.path.join(config.DATA_DIR, config.TRAIN_LABELS))\n    class_df = pd.read_csv(os.path.join(config.DATA_DIR, config.CLASS_INFO))\n    df = pd.merge(labels_df, class_df, on='patientId', how='left')\n    df['Target'] = df['Target'].astype(int)\n    \n    return df","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Main","metadata":{}},{"cell_type":"code","source":"def robust_preprocess(dicom_path):\n    \"\"\"Robust DICOM preprocessing with error handling\"\"\"\n    try:\n        # Read DICOM file\n        dicom = pydicom.dcmread(dicom_path)\n        image = dicom.pixel_array\n        \n        # Handle different photometric interpretations\n        if hasattr(dicom, 'PhotometricInterpretation'):\n            if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n                image = np.amax(image) - image\n        \n        # Normalize and convert to uint8\n        image = image.astype(np.float32)\n        image = (image - image.min()) / (image.max() - image.min()) * 255.0\n        image = image.astype(np.uint8)\n        \n        return image\n    except Exception as e:\n        print(f\"Error processing {dicom_path}: {str(e)}\")\n        # Return a blank image if processing fails\n        return np.zeros((1024, 1024), dtype=np.uint8)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fix 3: Improved safe_predict function\ndef safe_predict(model, image_tensor):\n    \"\"\"Robust prediction with input validation and error fallback.\"\"\"\n    if not isinstance(image_tensor, torch.Tensor):\n        raise ValueError(\"Input must be a PyTorch tensor\")\n    \n    try:\n        model.eval()\n        with torch.no_grad():\n            outputs = model(image_tensor)\n            # Handle single-output vs. multi-output models\n            if isinstance(outputs, (list, tuple)):\n                pred_logits, pred_boxes = outputs\n            else:\n                pred_logits = outputs\n                pred_boxes = None\n                \n            # Debug: Print shapes on first call\n            if not hasattr(safe_predict, '_debug_printed'):\n                print(f\"Model output shapes: logits={pred_logits.shape}, boxes={pred_boxes.shape if pred_boxes is not None else None}\")\n                safe_predict._debug_printed = True\n                \n        return pred_logits, pred_boxes\n    except RuntimeError as e:\n        print(f\"Prediction error: {e}\")\n        # Return appropriately sized fallback tensors\n        batch_size = image_tensor.shape[0]\n        fallback_logits = torch.zeros(batch_size, config.MAX_BOXES, config.NUM_CLASSES).to(image_tensor.device)\n        fallback_boxes = torch.zeros(batch_size, config.MAX_BOXES, 4).to(image_tensor.device)\n        return fallback_logits, fallback_boxes","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fix 2: Update the submission creation function\ndef create_submission(model: torch.nn.Module):\n    print(\"Creating robust submission...\")\n\n    test_image_dir = \"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_test_images\"\n    output_path = \"/kaggle/working/submission.csv\"\n    error_log = []\n\n    test_images = list(Path(test_image_dir).glob(\"*.dcm\"))\n    patient_ids = [img.stem for img in test_images]\n\n    optimal_threshold = 0.25\n\n    # Test transforms (matches train/val but without augmentation)\n    test_transforms = A.Compose([\n        A.Resize(512, 512),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2()\n    ])\n\n    submission_data = []\n    model.eval()\n    \n    with torch.no_grad():\n        for patient_id in tqdm(patient_ids):\n            try:\n                dicom_path = Path(test_image_dir) / f\"{patient_id}.dcm\"\n                image = robust_preprocess(dicom_path)\n                \n                if config.USE_CLAHE:\n                    image = DICOMPreprocessor.apply_clahe(image)\n                \n                image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n                transformed = test_transforms(image=image)\n                image_tensor = transformed[\"image\"].unsqueeze(0).to(device)\n\n                pred_logits, pred_boxes = model(image_tensor)\n                \n                # FIXED: Handle different output dimensions properly\n                if pred_logits.shape[-1] == 2:\n                    # Binary classification case - use class 1 (pneumonia)\n                    confidences = torch.sigmoid(pred_logits[0, :, 1])\n                elif pred_logits.shape[-1] == 1:\n                    # Single output case - use the single output\n                    confidences = torch.sigmoid(pred_logits[0, :, 0])\n                else:\n                    print(f\"Warning: Unexpected logits shape: {pred_logits.shape}\")\n                    confidences = torch.sigmoid(pred_logits[0, :, -1])  # Use last class\n                \n                boxes = pred_boxes[0].cpu().numpy()\n                \n                pred_strs = []\n                for i in range(min(config.MAX_BOXES, len(confidences))):  # FIXED: Handle size mismatch\n                    if confidences[i] >= optimal_threshold and boxes[i, 2] > 0.01 and boxes[i, 3] > 0.01:\n                        pred_strs.extend([\n                            str(confidences[i].item()),\n                            str(boxes[i, 0]),\n                            str(boxes[i, 1]),\n                            str(boxes[i, 2]),\n                            str(boxes[i, 3])\n                        ])\n                \n                pred_str = \" \".join(pred_strs) if pred_strs else \"0.5 0.25 0.25 0.5 0.5\"\n                \n            except Exception as e:\n                error_log.append({\n                    'patientId': patient_id,\n                    'error': str(e)\n                })\n                print(f\"Error processing {patient_id}: {e}\")\n                pred_str = \"0.5 0.25 0.25 0.5 0.5\"\n            \n            submission_data.append({\n                'patientId': patient_id,\n                'PredictionString': pred_str\n            })\n\n    if error_log:\n        pd.DataFrame(error_log).to_csv(\"submission_errors.csv\", index=False)\n\n    pd.DataFrame(submission_data).to_csv(output_path, index=False)\n    print(f\"Submission created with {len(error_log)} errors logged\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fix 4: Update the main function to ensure consistent configuration\ndef main():\n    \"\"\"Main training script with fixes\"\"\"\n    print(\"Loading and preparing data...\")\n    \n    # Clear GPU memory\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()\n        print(f\"Initial GPU memory: {torch.cuda.memory_allocated() / 1e6:.1f} MB\")\n    \n    gc.collect()\n    \n    # Data preparation...\n    df = prepare_data()\n    train_data, val_data = create_data_splits(df)\n    \n    # Transforms...\n    train_transforms = A.Compose([\n        A.Resize(512, 512),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2()\n    ])\n    \n    val_transforms = A.Compose([\n        A.Resize(512, 512),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2()\n    ])\n\n    # Create datasets...\n    train_dataset = PneumoniaDataset(\n        train_data, \n        os.path.join(config.DATA_DIR, config.TRAIN_IMAGES),\n        train_transforms, \n        is_training=True\n    )\n    \n    val_dataset = PneumoniaDataset(\n        val_data, \n        os.path.join(config.DATA_DIR, config.TRAIN_IMAGES),\n        val_transforms, \n        is_training=False\n    )\n\n    # Data loaders...\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=8,\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True,\n        drop_last=True,\n        persistent_workers=True,\n        prefetch_factor=2,\n        collate_fn=collate_fn\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=8,\n        shuffle=False,\n        num_workers=2,\n        pin_memory=True,\n        drop_last=False,\n        persistent_workers=True,\n        prefetch_factor=2,\n        collate_fn=collate_fn\n    )\n\n    # FIXED: Ensure model uses correct number of classes\n    model = PneumoniaDetectionModel(\n        backbone_name='efficientnet_b0',\n        num_classes=2,  # Explicitly set to 2 for binary classification\n        hidden_dim=256,\n        max_boxes=3,\n        dropout_rate=0.3\n    ).to(device)\n\n    # Print model info for debugging\n    print(f\"Model created with:\")\n    print(f\"  - num_classes: {model.num_classes}\")\n    print(f\"  - max_boxes: {model.max_boxes}\")\n    \n    # Test the model with a dummy input\n    with torch.no_grad():\n        dummy_input = torch.randn(1, 3, 512, 512).to(device)\n        test_logits, test_boxes = model(dummy_input)\n        print(f\"Test output shapes: logits={test_logits.shape}, boxes={test_boxes.shape}\")\n\n    # Continue with training setup...\n    criterion = MultiTaskLoss(\n        classification_weight=config.CLASSIFICATION_WEIGHT,\n        box_reg_weight=config.BOX_REG_WEIGHT,\n        focal_alpha=config.FOCAL_ALPHA,\n        focal_gamma=config.FOCAL_GAMMA\n    )\n    \n    optimizer = AdamW(\n        model.parameters(),\n        lr=config.LEARNING_RATE,\n        weight_decay=config.WEIGHT_DECAY\n    )\n    \n    scheduler = CosineAnnealingLR(\n        optimizer,\n        T_max=config.EPOCHS,\n        eta_min=1e-6\n    )\n    \n    trainer = Trainer(model, criterion, optimizer, scheduler, device)\n    \n    print(\"Starting training...\")\n    \n    # Training loop...\n    best_score = 0.0\n    patience_counter = 0\n    \n    for epoch in range(config.EPOCHS):\n        print(f\"\\nEpoch {epoch + 1}/{config.EPOCHS}\")\n        print(\"-\" * 50)\n\n        train_metrics = trainer.train_epoch(train_loader)\n        val_metrics = trainer.validate_epoch(val_loader)\n        trainer.scheduler.step()\n        current_lr = trainer.optimizer.param_groups[0]['lr']\n        \n        print(f\"Train Loss: {train_metrics['total_loss']:.6f}\")\n        print(f\"  Cls Loss: {train_metrics['cls_loss']:.6f}\")\n        print(f\"  Box Loss: {train_metrics['box_loss']:.6f}\")\n        print(f\"Val Loss: {val_metrics['total_loss']:.6f}\")\n        print(f\"  Cls Loss: {val_metrics['cls_loss']:.6f}\")\n        print(f\"  Box Loss: {val_metrics['box_loss']:.6f}\")\n        print(f\"Learning Rate: {current_lr:.8f}\")\n        \n        combined_score = 1.0 / (val_metrics['box_loss'] + 1e-8) \n        \n        if combined_score > best_score:\n            best_score = combined_score\n            patience_counter = 0\n            \n            checkpoint = {\n                'epoch': epoch + 1,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'best_score': best_score,\n                'config': config\n            }\n            \n            torch.save(\n                checkpoint,\n                os.path.join(config.OUTPUT_DIR, f\"{config.MODEL_BACKBONE}_fold{config.FOLD}_best.pth\")\n            )\n            \n            print(f\"New best score: {best_score:.6f} - Model saved!\")\n        else:\n            patience_counter += 1\n            print(f\"No improvement. Patience: {patience_counter}/{config.PATIENCE}\")\n        \n        if patience_counter >= config.PATIENCE:\n            print(\"Early stopping triggered!\")\n            break\n        \n        if (epoch + 1) % 5 == 0:\n            torch.cuda.empty_cache()\n            gc.collect()\n    \n    print(f\"\\nTraining completed! Best validation score: {best_score:.6f}\")\n\n    # Create submission with the trained model\n    create_submission(model)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}