{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":5877069,"sourceType":"datasetVersion","datasetId":2998419},{"sourceId":12546895,"sourceType":"datasetVersion","datasetId":7921768},{"sourceId":133022801,"sourceType":"kernelVersion"}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# Breast Cancer Detection Pipeline Description\n\nThis notebook implements a deep learning pipeline for breast cancer detection from mammogram images using the RSNA Screening Mammography Breast Cancer Detection dataset (https://www.kaggle.com/competitions/rsna-breast-cancer-detection/data). inspired by:\n- `dangnh0611/kaggle_rsna_breast_cancer` (https://github.com/dangnh0611/kaggle_rsna_breast_cancer): Two-stage pipeline with YOLOX for ROI detection and ConvNeXt for classification.\n- `nyukat/breast_cancer_classifier` (https://github.com/nyukat/breast_cancer_classifier): Multi-view processing and Grad-CAM for interpretability.\n\n## Pipeline Overview\n1. **Data Preprocessing**:\n   - Converts high-resolution DICOM images to 512x512 PNGs to reduce memory usage and speed up training.\n   - Handles the RSNA dataset's structure: `train.csv` for metadata, `train_images/[patient_id]/[image_id].dcm` for images.\n2. **Dataset Setup**:\n   - Groups images by `patient_id` to process multiple views (CC and MLO) together, improving accuracy by leveraging complementary information\n   - Addresses class imbalance (~2% cancer prevalence) via oversampling positive cases (5x) and weighted loss.\n3. **Two-Stage Model**:\n   - **Stage 1: ROI Detection**: Uses pre-trained YOLOX-nano  to identify regions of interest, reducing noise from irrelevant areas.\n   - **Stage 2: Classification**: Uses EfficientNet-B3 (default) or ConvNeXt-small (if weights available) to classify ROIs as cancerous or non-cancerous.//Not decided yet\n4. **Training and Evaluation**:\n   - Trains with an 80/20 train-validation split, using weighted BCE loss to handle imbalance.\n   - Evaluates with AUC-ROC and probabilistic F1 score, prioritizing sensitivity for cancer detection.\n5. **Interpretability**:\n   - Generates Grad-CAM visualizations for each view to highlight regions influencing predictions, enhancing clinical trust \n\n## Why This Approach?\n- **RSNA Dataset**: Large-scale (~54,000 images), real-world screening data, multi-view support, and Kaggle accessibility make it ideal for deep learning.\n- **Class Imbalance**: Oversampling and weighted loss mitigate the 2% cancer prevalence, ensuring robust detection.\n- **Multi-View Processing**: Combines CC and MLO views for better accuracy\n- **Two-Stage Pipeline**: YOLOX reduces noise, and EfficientNet/ConvNeXt leverages modern architectures for high performance.\n- **Interpretability**: Grad-CAM ensures the model is clinically relevant.\n\n## Expected Output\n- Preprocessed PNG images in `/kaggle/working/png_images/`.\n- Training logs with epoch losses and validation AUC/F1 scores.\n- Grad-CAM visualizations for a sample patient’s views.\n- Saved model weights at `/kaggle/working/classifier_model.pth`.\n\"\"\"\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"# Install Dependencies\n# Install packages for DICOM processing, YOLOX, and visualization, including pylibjpeg for JPEG Lossless\n!pip install --no-cache-dir --verbose torchcam\n!pip install --no-cache-dir --verbose pydicom\n!pip install --no-cache-dir --verbose kagglehub\n!pip install --no-cache-dir --verbose pylibjpeg>=2.0 pylibjpeg-libjpeg>=2.1 pylibjpeg-openjpeg>=1.3  # For JPEG Lossless decompression\n!pip install --no-cache-dir --verbose git+https://github.com/MegEngine/YOLOX.git  # For YOLOX ROI detection\n\n# Verify pydicom with pylibjpeg\ntry:\n    import pylibjpeg\n    print(f\"pylibjpeg version: {pylibjpeg.__version__}\")\n    test_dicom = pydicom.dcmread(\"/kaggle/input/rsna-breast-cancer-detection/train_images/10006/1459541791.dcm\")\n    _ = test_dicom.pixel_array  # Test decompression\n    print(\"pydicom decompression test passed.\")\nexcept Exception as e:\n    print(f\"pydicom decompression test failed: {e}\")\n\nimport kagglehub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T17:58:21.886027Z","iopub.execute_input":"2025-07-22T17:58:21.886591Z","iopub.status.idle":"2025-07-22T17:58:49.726573Z","shell.execute_reply.started":"2025-07-22T17:58:21.886565Z","shell.execute_reply":"2025-07-22T17:58:49.72568Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Import Libraries\n","metadata":{}},{"cell_type":"code","source":"# Core libraries for data processing, deep learning, and visualization\nimport os\nimport pandas as pd\nimport numpy as np\nimport pydicom\nfrom PIL import Image\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torchvision import transforms, models\nimport cv2\nfrom sklearn.metrics import roc_auc_score, f1_score\nfrom sklearn.model_selection import train_test_split\nfrom torchcam.methods import GradCAM\nimport matplotlib.pyplot as plt\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T17:56:15.419465Z","iopub.execute_input":"2025-07-22T17:56:15.420018Z","iopub.status.idle":"2025-07-22T17:56:15.425566Z","shell.execute_reply.started":"2025-07-22T17:56:15.419985Z","shell.execute_reply":"2025-07-22T17:56:15.424795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pylibjpeg>=2.0 pylibjpeg-libjpeg>=2.1\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T17:55:43.794944Z","iopub.execute_input":"2025-07-22T17:55:43.79523Z","iopub.status.idle":"2025-07-22T17:55:46.840169Z","shell.execute_reply.started":"2025-07-22T17:55:43.795204Z","shell.execute_reply":"2025-07-22T17:55:46.839167Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Downloading Pre-trained Weights and  Converting DICOMs to PNGs\n","metadata":{}},{"cell_type":"code","source":"# Download Pre-trained Weights\n# Download YOLOX and ConvNeXt weights from Kaggle dataset for transfer learning\ntry:\n    weights_path = kagglehub.dataset_download(\"dangnh0611/rsna-breast-cancer-detection-best-ckpts\")\n    print(f\"Weights downloaded to: {weights_path}\")\nexcept Exception as e:\n    print(f\"Error downloading weights: {e}\")\n    weights_path = None\n\n# Dataset Paths\n# Define paths for RSNA dataset and output directory for preprocessed PNGs\nDATA_PATH = '/kaggle/input/rsna-breast-cancer-detection/'\nTRAIN_CSV = os.path.join(DATA_PATH, 'train.csv')\nTRAIN_DICOM_PATH = os.path.join(DATA_PATH, 'train_images')\nOUTPUT_PNG_PATH = '/kaggle/working/png_images/'\nos.makedirs(OUTPUT_PNG_PATH, exist_ok=True)\n\n# Preprocess DICOM to PNG\n# Convert high-resolution DICOM images to PNGs, handling JPEG Lossless compression\ndef dicom_to_png(dicom_path, output_path, resize_dim=512):\n    try:\n        ds = pydicom.dcmread(dicom_path)\n        img = ds.pixel_array  # Requires pylibjpeg for JPEG Lossless\n        img = (img - img.min()) / (img.max() - img.min() + 1e-6) * 255  # Normalize to 0-255\n        img = img.astype(np.uint8)\n        img = Image.fromarray(img).convert('RGB')\n        img = img.resize((resize_dim, resize_dim), Image.LANCZOS)  # Resize during preprocessing\n        img.save(output_path)\n        return True\n    except Exception as e:\n        print(f\"Error processing {dicom_path}: {e}\")\n        return False\n\n# Convert DICOMs to PNGs\n# Process all DICOM images in train.csv, save as PNGs, and log skipped files\ndef preprocess_dicom_to_png(csv_file, dicom_base_path, output_base_path, resize_dim=512):\n    df = pd.read_csv(csv_file)\n    skipped_files = []\n    for idx, row in df.iterrows():\n        patient_id = str(row['patient_id'])\n        image_id = str(row['image_id'])\n        dicom_path = os.path.join(dicom_base_path, patient_id, f\"{image_id}.dcm\")\n        output_path = os.path.join(output_base_path, f\"{patient_id}_{image_id}.png\")\n        if not os.path.exists(output_path):\n            success = dicom_to_png(dicom_path, output_path, resize_dim=resize_dim)\n            if not success:\n                skipped_files.append(dicom_path)\n    if skipped_files:\n        print(f\"Skipped {len(skipped_files)} files due to processing errors: {skipped_files[:5]}...\")\n    else:\n        print(\"All DICOM files processed successfully.\")\n    return skipped_files\n\nprint(\"Preprocessing train images...\")\nskipped_files = preprocess_dicom_to_png(TRAIN_CSV, TRAIN_DICOM_PATH, OUTPUT_PNG_PATH, resize_dim=512)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-22T17:56:19.285537Z","iopub.execute_input":"2025-07-22T17:56:19.285812Z","iopub.status.idle":"2025-07-22T17:58:18.972447Z","shell.execute_reply.started":"2025-07-22T17:56:19.285791Z","shell.execute_reply":"2025-07-22T17:58:18.971333Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Oversampling Positive Cases\ndef create_oversampled_indices(csv_file):\n    df = pd.read_csv(csv_file)\n    patient_groups = df.groupby('patient_id')\n    indices = []\n    positive_patients = []\n    negative_patients = []\n    \n    for idx, (patient_id, group) in enumerate(patient_groups):\n        if group['cancer'].max() == 1:\n            positive_patients.append(idx)\n        else:\n            negative_patients.append(idx)\n    \n    # Oversample positive patients (e.g., 5x)\n    indices.extend(positive_patients * 5)\n    indices.extend(negative_patients)\n    return indices\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# YOLOX Setup\n","metadata":{}},{"cell_type":"code","source":"\ntry:\n    from yolox.models import YOLOX, YOLOPAFPN, YOLOXHead\n    from yolox.utils import postprocess\nexcept ImportError:\n    print(\"YOLOX import failed. Falling back to full-image classification.\")\n    YOLOX = None\n\nclass YOLOXNano(nn.Module):\n    def __init__(self):\n        super().__init__()\n        backbone = YOLOPAFPN(depth=0.33, width=0.375)\n        head = YOLOXHead(num_classes=1, width=0.375)\n        self.model = YOLOX(backbone, head)\n    \n    def forward(self, x):\n        return self.model(x)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Custom Dataset Class with Multi-View Support\n","metadata":{}},{"cell_type":"code","source":"\nclass MammogramDataset(Dataset):\n    def __init__(self, csv_file, img_dir, transform=None):\n        self.data = pd.read_csv(csv_file)\n        self.img_dir = img_dir\n        self.transform = transform\n        self.patient_groups = self.data.groupby('patient_id')\n    \n    def __len__(self):\n        return len(self.patient_groups)\n    \n    def __getitem__(self, idx):\n        patient_id = list(self.patient_groups.groups.keys())[idx]\n        patient_data = self.patient_groups.get_group(patient_id)\n        images = []\n        labels = []\n        \n        for _, row in patient_data.iterrows():\n            image_id = str(row['image_id'])\n            img_path = os.path.join(self.img_dir, f\"{patient_id}_{image_id}.png\")\n            try:\n                image = Image.open(img_path).convert('RGB')\n                if self.transform:\n                    image = self.transform(image)\n                images.append(image)\n                labels.append(row['cancer'])\n            except FileNotFoundError:\n                print(f\"Image not found: {img_path}. Skipping.\")\n                continue\n        \n        if not images:\n            return None, None\n        \n        images = torch.stack(images)\n        label = max(labels)  # Positive if any image is positive\n        return images, label","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Transformations\n","metadata":{}},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.Resize((512, 512)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomRotation(10),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Load Data with Oversampling and Validation Split\n","metadata":{}},{"cell_type":"code","source":"\n\ntrain_dataset = MammogramDataset(TRAIN_CSV, OUTPUT_PNG_PATH, transform=transform)\noversampled_indices = create_oversampled_indices(TRAIN_CSV)\ntrain_idx, val_idx = train_test_split(oversampled_indices, test_size=0.2, random_state=42)\ntrain_subset = Subset(train_dataset, train_idx)\nval_subset = Subset(train_dataset, val_idx)\ntrain_loader = DataLoader(train_subset, batch_size=8, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_subset, batch_size=8, shuffle=False, num_workers=4)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Load YOLOX Model\nyolox_model = None\nif YOLOX is not None and weights_path is not None:\n    try:\n        yolox_model = YOLOXNano()\n        yolox_weights = os.path.join(weights_path, 'yolox_nano_416_roi_torch.pth')\n        if os.path.exists(yolox_weights):\n            yolox_model.load_state_dict(torch.load(yolox_weights, map_location='cpu'))\n            yolox_model = yolox_model.cuda() if torch.cuda.is_available() else yolox_model\n            yolox_model.eval()\n            print(\"YOLOX model loaded successfully.\")\n        else:\n            print(f\"YOLOX weights not found at {yolox_weights}. Using full images.\")\n            yolox_model = None\n    except Exception as e:\n        print(f\"Error loading YOLOX model: {e}. Using full images.\")\n        yolox_model = None\nelse:\n    print(\"YOLOX not available. Using full images.\")\n\n# Define Classification Model with Multi-View Support\nclass BreastCancerClassifier(nn.Module):\n    def __init__(self, use_convnext=False):\n        super().__init__()\n        if use_convnext:\n            self.backbone = models.convnext_small(pretrained=False)\n            self.backbone.classifier[2] = nn.Linear(self.backbone.classifier[2].in_features, 128)\n        else:\n            self.backbone = models.efficientnet_b3(pretrained=True)\n            self.backbone.classifier = nn.Sequential(\n                nn.Dropout(p=0.4),\n                nn.Linear(self.backbone.classifier[1].in_features, 128)\n            )\n        self.fusion = nn.Linear(128 * 4, 128)  # Assume up to 4 views\n        self.classifier = nn.Linear(128, 1)\n    \n    def forward(self, x):\n        batch_size, num_views, c, h, w = x.shape\n        features = []\n        for i in range(num_views):\n            view = x[:, i, :, :, :]\n            feat = self.backbone(view)\n            features.append(feat)\n        features = torch.cat(features, dim=1)\n        fused = self.fusion(features)\n        out = self.classifier(fused)\n        return out\n\n# Initialize Model\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nuse_convnext = False\nif weights_path is not None:\n    try:\n        convnext_weights = os.path.join(weights_path, 'convnext_small.pth')\n        if os.path.exists(convnext_weights):\n            print(\"Switching to ConvNeXt model and loading weights.\")\n            use_convnext = True\n    except Exception as e:\n        print(f\"Error checking ConvNeXt weights: {e}. Using EfficientNet-B3.\")\n\nclassifier = BreastCancerClassifier(use_convnext=use_convnext).to(device)\nif use_convnext and weights_path is not None:\n    try:\n        classifier.load_state_dict(torch.load(convnext_weights, map_location=device), strict=False)\n        print(\"ConvNeXt weights loaded successfully.\")\n    except Exception as e:\n        print(f\"Error loading ConvNeXt weights: {e}. Using EfficientNet-B3 without pre-trained weights.\")\n\n# Loss and Optimizer\npos_weight = torch.tensor([10.0]).to(device)\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\noptimizer = optim.AdamW(classifier.parameters(), lr=1e-4)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)\n\n# ROI Detection Function\ndef detect_roi(image, yolox_model):\n    if yolox_model is None:\n        return image\n    try:\n        with torch.no_grad():\n            outputs = yolox_model(image)\n            predictions = postprocess(outputs, conf_thre=0.3, nms_thre=0.45)[0]\n            if predictions is None:\n                return image\n            x1, y1, x2, y2 = predictions[0, :4].int()\n            roi = image[:, :, y1:y2, x1:x2]\n            if roi.shape[2] == 0 or roi.shape[3] == 0:\n                return image\n            return roi\n    except Exception as e:\n        print(f\"ROI detection error: {e}. Returning original image.\")\n        return image\n\n# Training Loop\ndef train_model(model, train_loader, val_loader, yolox_model, epochs=5):\n    model.train()\n    for epoch in range(epochs):\n        running_loss = 0.0\n        for images, labels in train_loader:\n            if images is None or labels is None:\n                continue\n            images, labels = images.to(device), labels.to(device).float()\n            \n            # Apply ROI detection to each view\n            roi_images = []\n            for batch_idx in range(images.shape[0]):\n                patient_views = images[batch_idx]\n                rois = [detect_roi(view.unsqueeze(0), yolox_model).squeeze(0) for view in patient_views]\n                try:\n                    rois = torch.stack(rois)\n                except RuntimeError:\n                    print(\"ROI stacking failed. Using original views.\")\n                    rois = patient_views\n                roi_images.append(rois)\n            roi_images = torch.stack(roi_images).to(device)\n            \n            optimizer.zero_grad()\n            outputs = model(roi_images).squeeze()\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n            running_loss += loss.item() * images.size(0)\n        \n        scheduler.step()\n        epoch_loss = running_loss / len(train_loader.dataset)\n        print(f'Epoch {epoch+1}/{epochs}, Loss: {epoch_loss:.4f}')\n        \n        # Evaluate on validation set\n        auc, f1 = evaluate_model(model, val_loader, yolox_model)\n\n# Evaluation Function\ndef evaluate_model(model, data_loader, yolox_model):\n    model.eval()\n    y_true, y_pred = [], []\n    with torch.no_grad():\n        for images, labels in data_loader:\n            if images is None or labels is None:\n                continue\n            images, labels = images.to(device), labels.to(device).float()\n            roi_images = []\n            for batch_idx in range(images.shape[0]):\n                patient_views = images[batch_idx]\n                rois = [detect_roi(view.unsqueeze(0), yolox_model).squeeze(0) for view in patient_views]\n                try:\n                    rois = torch.stack(rois)\n                except RuntimeError:\n                    print(\"ROI stacking failed. Using original views.\")\n                    rois = patient_views\n                roi_images.append(rois)\n            roi_images = torch.stack(roi_images).to(device)\n            outputs = model(roi_images).squeeze()\n            probs = torch.sigmoid(outputs)\n            y_true.extend(labels.cpu().numpy())\n            y_pred.extend(probs.cpu().numpy())\n    \n    auc = roc_auc_score(y_true, y_pred) if len(y_true) > 0 else 0.0\n    f1 = f1_score(y_true, (np.array(y_pred) > 0.5).astype(int)) if len(y_true) > 0 else 0.0\n    print(f'AUC: {auc:.4f}, Probabilistic F1: {f1:.4f}')\n    return auc, f1\n\n# Grad-CAM Visualization\ndef visualize_gradcam(model, images, target_layer='backbone.features'):\n    model.eval()\n    cam_extractor = GradCAM(model, target_layer)\n    images = images.to(device)\n    outputs = model(images.unsqueeze(0)).squeeze()\n    prob = torch.sigmoid(outputs).item()\n    \n    cams = []\n    for i in range(images.shape[0]):\n        image = images[i].unsqueeze(0)\n        cam = cam_extractor(0, model(image.unsqueeze(0))).cpu().numpy()[0]\n        cam = cv2.resize(cam, (512, 512))\n        cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-6)\n        cams.append(cam)\n    \n    plt.figure(figsize=(15, 5))\n    for i, (image, cam) in enumerate(zip(images, cams)):\n        image = image.cpu().numpy().transpose(1, 2, 0)\n        image = (image * np.array([0.229, 0.224, 0.225]) + np.array([0.485, 0.456, 0.406])).clip(0, 1)\n        heatmap = cv2.applyColorMap((cam * 255).astype(np.uint8), cv2.COLORMAP_JET)\n        overlay = (image * 255 * 0.5 + heatmap * 0.5).astype(np.uint8)\n        \n        plt.subplot(2, len(cams), i + 1)\n        plt.imshow(image)\n        plt.title(f'View {i + 1}')\n        plt.subplot(2, len(cams), len(cams) + i + 1)\n        plt.imshow(overlay)\n        plt.title(f'Grad-CAM (Prob: {prob:.4f})')\n    plt.tight_layout()\n    plt.show()\n\n\n```","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train the Model\ntrain_model(classifier, train_loader, val_loader, yolox_model, epochs=5)\n\n# Evaluate on Validation Set\nevaluate_model(classifier, val_loader, yolox_model)\n\n# Visualize Grad-CAM on a Sample Patient\nsample_images, sample_label = train_dataset[0]\nif sample_images is not None:\n    visualize_gradcam(classifier, sample_images)\n\n# Save Model\ntorch.save(classifier.state_dict(), '/kaggle/working/classifier_model.pth')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}