Source code for medical_image.datasets.cbis_ddsm

"""
CBIS-DDSM mammography dataset with full-image and patch-based loading.

Supports the TCIA CBIS-DDSM directory layout with automatic pairing of
full mammogram images and their corresponding ROI mask DICOMs.
"""

import random
from pathlib import Path
from typing import Optional, Callable, Dict, Any, Tuple, List, Literal

import numpy as np
import torch
import torch.nn.functional as F
from skimage.feature import match_template
from skimage.filters import sobel

from medical_image.utils.logging import logger
from medical_image.data.dicom_image import DicomImage
from medical_image.datasets.base_dataset import BaseDataset
from medical_image.utils.mask_utils import stack_dicom_masks
from medical_image.utils.pairing import pair_cbis_ddsm, CBISDDSMSample


[docs] class CBISDDSMDataset(BaseDataset): """ PyTorch Dataset for the CBIS-DDSM mammography database. Supports two loading modes: - ``"full_image"``: Load the entire mammogram (optionally resized). - ``"patch"``: Extract sliding-window patches from each mammogram. Each sample via ``__getitem__`` returns: .. code-block:: python { "image": Tensor[1, H, W], # mammogram (or patch) "mask": Tensor[1, H, W], # OR-merged ROI masks "metadata": { "case_id": str, "patient_id": str, "side": str, "view": str, "num_masks": int, "patch_idx": int, # only in patch mode "patch_position": (y, x), # only in patch mode } } Use :meth:`get_detailed_sample` for the rich format with per-ROI bounding boxes, crops, and masks. Args: root_dir: Path to the manifest directory containing ``CBIS-DDSM/``. mode: ``"full_image"`` or ``"patch"``. patch_size: Patch side length (used when ``mode="patch"``). stride: Stride between patches (used when ``mode="patch"``). transform: Optional image transform. target_transform: Optional mask transform. target_size: Optional (H, W) resize for full_image mode. percentage: Optional float (0-1] to use a random subset of cases. seed: Random seed for reproducible subset selection. Example: .. code-block:: python dataset = CBISDDSMDataset( root_dir="data/ddsm", mode="patch", patch_size=512, stride=256, percentage=0.2, ) loader = DataLoader( dataset, batch_size=4, shuffle=True, num_workers=4, collate_fn=dataset.collate_fn, ) """
[docs] def __init__( self, root_dir: str, mode: Literal["full_image", "patch"] = "full_image", patch_size: int = 512, stride: int = 256, transform: Optional[Callable] = None, target_transform: Optional[Callable] = None, target_size: Optional[Tuple[int, int]] = None, percentage: Optional[float] = None, seed: int = 42, ): self.mode = mode self.patch_size = patch_size self.stride = stride self.percentage = percentage self.seed = seed # Will be populated by _build_sample_list self._case_samples: List[CBISDDSMSample] = [] self._patch_index: List[Tuple[int, int, int]] = [] # (case_idx, y, x) super().__init__(root_dir, transform, target_transform, target_size)
# ------------------------------------------------------------------ # Build sample list # ------------------------------------------------------------------ def _build_sample_list(self) -> None: self._case_samples = pair_cbis_ddsm(self.root_dir) if not self._case_samples: logger.warning(f"No CBIS-DDSM cases found in {self.root_dir}") self._samples = [] return # Apply percentage filtering if self.percentage is not None: n_select = max(1, int(len(self._case_samples) * self.percentage)) rng = random.Random(self.seed) self._case_samples = sorted( rng.sample(self._case_samples, n_select), key=lambda s: s.case_id, ) logger.info( f"CBIS-DDSM: selected {n_select} cases " f"({self.percentage * 100:.0f}%)" ) if self.mode == "full_image": # One sample per case self._samples = list(range(len(self._case_samples))) elif self.mode == "patch": # Pre-compute patch positions for each case self._samples = [] self._patch_index = [] for case_idx, case in enumerate(self._case_samples): # Need to peek at image dimensions h, w = self._peek_image_size(case.mammogram_path) positions = self._compute_patch_positions(h, w) for y, x in positions: sample_idx = len(self._samples) self._samples.append(sample_idx) self._patch_index.append((case_idx, y, x)) logger.info( f"CBIS-DDSM patch mode: {len(self._samples)} patches " f"from {len(self._case_samples)} cases " f"(size={self.patch_size}, stride={self.stride})" ) else: raise ValueError(f"Unknown mode: {self.mode}. Use 'full_image' or 'patch'.") # ------------------------------------------------------------------ # Load single sample (BaseDataset contract) # ------------------------------------------------------------------ def _load_sample(self, idx: int) -> Dict[str, Any]: if self.mode == "full_image": return self._load_full_image(idx) else: return self._load_patch(idx) def _load_full_image(self, idx: int) -> Dict[str, Any]: case = self._case_samples[idx] # Load mammogram dcm = DicomImage(file_path=case.mammogram_path) dcm.load() image_tensor = self._to_chw(dcm.pixel_data.float()) # Load and merge masks h, w = dcm.height, dcm.width if case.mask_paths: try: mask_tensor = stack_dicom_masks(case.mask_paths) # Resize mask to match mammogram if needed if mask_tensor.shape[1:] != (h, w): mask_tensor = F.interpolate( mask_tensor.unsqueeze(0), size=(h, w), mode="nearest", ).squeeze(0) except Exception as e: logger.warning(f"Failed to load masks for {case.case_id}: {e}") mask_tensor = torch.zeros(1, h, w, dtype=torch.float32) else: mask_tensor = torch.zeros(1, h, w, dtype=torch.float32) metadata = { "case_id": case.case_id, "patient_id": case.patient_id, "side": case.side, "view": case.view, "num_masks": len(case.mask_paths), "file_name": Path(case.mammogram_path).name, } return {"image": image_tensor, "mask": mask_tensor, "metadata": metadata} def _load_patch(self, idx: int) -> Dict[str, Any]: case_idx, py, px = self._patch_index[idx] case = self._case_samples[case_idx] # Load mammogram dcm = DicomImage(file_path=case.mammogram_path) dcm.load() image_tensor = self._to_chw(dcm.pixel_data.float()) # Extract patch from image image_patch = image_tensor[ :, py : py + self.patch_size, px : px + self.patch_size ] # Load mask and extract corresponding patch h, w = dcm.height, dcm.width if case.mask_paths: try: mask_tensor = stack_dicom_masks(case.mask_paths) if mask_tensor.shape[1:] != (h, w): mask_tensor = F.interpolate( mask_tensor.unsqueeze(0), size=(h, w), mode="nearest", ).squeeze(0) except Exception as e: logger.warning(f"Failed to load masks for {case.case_id}: {e}") mask_tensor = torch.zeros(1, h, w, dtype=torch.float32) else: mask_tensor = torch.zeros(1, h, w, dtype=torch.float32) mask_patch = mask_tensor[ :, py : py + self.patch_size, px : px + self.patch_size ] # Pad if patch is at the edge _, ph, pw = image_patch.shape if ph < self.patch_size or pw < self.patch_size: image_patch = F.pad( image_patch, (0, self.patch_size - pw, 0, self.patch_size - ph) ) mask_patch = F.pad( mask_patch, (0, self.patch_size - pw, 0, self.patch_size - ph) ) metadata = { "case_id": case.case_id, "patient_id": case.patient_id, "side": case.side, "view": case.view, "num_masks": len(case.mask_paths), "patch_idx": idx, "patch_position": (py, px), } return {"image": image_patch, "mask": mask_patch, "metadata": metadata} # ------------------------------------------------------------------ # Detailed sample with per-ROI bboxes, crops, and masks # ------------------------------------------------------------------
[docs] def get_detailed_sample(self, idx: int) -> Dict[str, Any]: """ Load a full-image sample with per-ROI bounding boxes, crops, and masks. Only available in ``"full_image"`` mode. Returns: Dict with: - ``"image"``: ``Tensor[1, H, W]`` — full mammogram - ``"bboxes"``: ``Tensor[N, 4]`` — ``(x_min, y_min, x_max, y_max)`` in full mammogram coordinates - ``"rois"``: ``List[Tensor]`` — ROI crop tensors (from 1-1.dcm) - ``"masks"``: ``List[Tensor]`` — ROI mask tensors (from 1-2.dcm) - ``"label"``: task string (e.g. ``"Calc-Test"``) - ``"meta"``: dict with patient_id, view, side, task """ if self.mode != "full_image": raise ValueError( "get_detailed_sample is only available in 'full_image' mode" ) case = self._case_samples[idx] # Load full mammogram dcm = DicomImage(file_path=case.mammogram_path) dcm.load() full_image = self._to_chw(dcm.pixel_data.float()) full_array = dcm.pixel_data.numpy().astype(np.float64) rois: List[torch.Tensor] = [] masks: List[torch.Tensor] = [] bboxes: List[List[int]] = [] for entry in case.roi_entries: # Load ROI crop via DicomImage roi_tensor: Optional[torch.Tensor] = None roi_array: Optional[np.ndarray] = None if entry.roi_path: try: roi_dcm = DicomImage(file_path=entry.roi_path) roi_dcm.load() roi_array = roi_dcm.pixel_data.numpy().astype(np.float64) roi_tensor = self._to_chw(roi_dcm.pixel_data.float()) except Exception as e: logger.warning(f"Failed to load ROI crop {entry.roi_path}: {e}") # Load mask via DicomImage mask_tensor: Optional[torch.Tensor] = None if entry.mask_path: try: mask_dcm = DicomImage(file_path=entry.mask_path) mask_dcm.load() mask_arr = (mask_dcm.pixel_data > 0).float() mask_tensor = self._to_chw(mask_arr) except Exception as e: logger.warning(f"Failed to load ROI mask {entry.mask_path}: {e}") if roi_tensor is not None: rois.append(roi_tensor) if mask_tensor is not None: masks.append(mask_tensor) # Compute bounding box from ROI crop position in full mammogram if roi_array is not None: try: bbox = self._locate_roi_in_mammogram(full_array, roi_array) bboxes.append(bbox) except Exception as e: logger.warning(f"Failed to compute bbox for {entry.roi_path}: {e}") bbox_tensor = ( torch.tensor(bboxes, dtype=torch.int64) if bboxes else torch.zeros(0, 4, dtype=torch.int64) ) return { "image": full_image, "bboxes": bbox_tensor, "rois": rois, "masks": masks, "label": case.task, "meta": { "patient_id": case.patient_id, "view": case.view, "side": case.side, "task": case.task, "mammogram_path": case.mammogram_path, }, }
# ------------------------------------------------------------------ # Bounding box computation # ------------------------------------------------------------------
[docs] @staticmethod def get_bounding_boxes(mask: np.ndarray) -> List[Tuple[int, int, int, int]]: """ Compute bounding boxes from a binary mask array. Finds connected components of non-zero pixels and returns the bounding box of each component. Args: mask: 2D numpy array (H, W) with non-zero pixels marking ROIs. Returns: List of ``(x_min, y_min, x_max, y_max)`` tuples. """ from scipy.ndimage import label as ndimage_label if mask.ndim != 2: raise ValueError(f"Expected 2D mask, got shape {mask.shape}") binary = (mask > 0).astype(np.uint8) if binary.sum() == 0: return [] labeled, n_components = ndimage_label(binary) boxes = [] for comp_id in range(1, n_components + 1): ys, xs = np.where(labeled == comp_id) x_min, x_max = int(xs.min()), int(xs.max()) y_min, y_max = int(ys.min()), int(ys.max()) boxes.append((x_min, y_min, x_max, y_max)) return boxes
@staticmethod def _locate_roi_in_mammogram(full_image: np.ndarray, roi_crop: np.ndarray): th, tw = roi_crop.shape # --- Step 1: Normalize (z-score) --- full_norm = (full_image - np.mean(full_image)) / (np.std(full_image) + 1e-8) roi_norm = (roi_crop - np.mean(roi_crop)) / (np.std(roi_crop) + 1e-8) # --- Step 2: Edge enhancement (VERY important for mammograms) --- full_edges = sobel(full_norm) roi_edges = sobel(roi_norm) # --- Step 3: Normalized cross-correlation --- result = match_template(full_edges, roi_edges, pad_input=False) # --- Step 4: Locate best match --- y_off, x_off = np.unravel_index(np.argmax(result), result.shape) return [int(x_off), int(y_off), int(x_off + tw), int(y_off + th)] # ------------------------------------------------------------------ # Custom collate_fn # ------------------------------------------------------------------
[docs] @staticmethod def collate_fn(batch: List[Dict[str, Any]]) -> Dict[str, Any]: """ Custom collate function that handles variable-sized masks. Pads all images and masks to the maximum spatial dimensions in the batch. Returns: Dict with stacked ``"image"`` and ``"mask"`` tensors, and a list of ``"metadata"`` dicts. """ # Find max spatial dimensions max_h = max(s["image"].shape[1] for s in batch) max_w = max(s["image"].shape[2] for s in batch) images = [] masks = [] metadata = [] for s in batch: img = s["image"] msk = s["mask"] # Pad to max dimensions _, h, w = img.shape pad_h = max_h - h pad_w = max_w - w if pad_h > 0 or pad_w > 0: img = F.pad(img, (0, pad_w, 0, pad_h)) msk = F.pad(msk, (0, pad_w, 0, pad_h)) images.append(img) masks.append(msk) metadata.append(s["metadata"]) return { "image": torch.stack(images), "mask": torch.stack(masks), "metadata": metadata, }
# ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ @staticmethod def _peek_image_size(dcm_path: str) -> Tuple[int, int]: """ Read DICOM header to get image dimensions without loading pixel data. """ import pydicom ds = pydicom.dcmread(dcm_path, stop_before_pixels=True) return int(ds.Rows), int(ds.Columns) def _compute_patch_positions(self, h: int, w: int) -> List[Tuple[int, int]]: """ Compute top-left (y, x) positions for sliding-window patches. """ positions = [] for y in range(0, h, self.stride): for x in range(0, w, self.stride): positions.append((y, x)) return positions