diff --git a/datamint/dataset/__init__.py b/datamint/dataset/__init__.py index 6ea11f59..1cfe80ea 100644 --- a/datamint/dataset/__init__.py +++ b/datamint/dataset/__init__.py @@ -17,7 +17,7 @@ from .video_dataset import VideoDataset from .sliced_dataset import SlicedVolumeDataset from .sliced_video_dataset import SlicedVideoDataset -from .detection_dataset import DetectionDataset, detection_collate_fn +from .image_dataset import detection_collate_fn from .factory import build_dataset from .split_result import SplitResult @@ -32,7 +32,6 @@ 'VideoDataset', 'SlicedVolumeDataset', 'SlicedVideoDataset', - 'DetectionDataset', 'detection_collate_fn', # Factory 'build_dataset', diff --git a/datamint/dataset/base.py b/datamint/dataset/base.py index 08e1d2b0..2b0c1feb 100644 --- a/datamint/dataset/base.py +++ b/datamint/dataset/base.py @@ -86,6 +86,7 @@ def __init__( # all_annotations: bool = False, return_metainfo: bool = True, return_segmentations: bool = True, + return_boxes: bool = False, return_as_semantic_segmentation: bool = False, semantic_seg_merge_strategy: MergeStrategy | None = None, alb_transform: 'Callable | BaseCompose | None' = None, @@ -152,6 +153,7 @@ def __init__( # Store configuration self.return_metainfo = return_metainfo self.return_segmentations = return_segmentations + self.return_boxes = return_boxes self.return_as_semantic_segmentation = return_as_semantic_segmentation self.semantic_seg_merge_strategy: MergeStrategy | None = semantic_seg_merge_strategy self.include_unannotated = include_unannotated @@ -629,6 +631,8 @@ def _setup_labels(self) -> None: self.image_lsets, self.image_lcodes = self._infer_labels_set(framed=False) self.seglabel_list, self.seglabel2code = self._infer_segmentation_group() + self.box_class_map: dict[str, int] = self._build_box_class_map() + def _augment_labels_from_annotations(self) -> None: """Augment project-defined label sets with identifiers found in actual annotations. @@ -855,6 +859,11 @@ def segmentation_labels_set(self) -> list[str]: """Segmentation label names.""" return self.seglabel_list + @property + def box_labels_set(self) -> list[str]: + """Box annotation class names, alphabetically ordered.""" + return sorted(self.box_class_map, key=self.box_class_map.__getitem__) + def _infer_segmentation_group(self) -> tuple[list[str], dict[str, int]]: """Infer segmentation labels from annotations when no project is provided.""" seglabel_set: set[str] = set() @@ -872,6 +881,55 @@ def _infer_segmentation_group(self) -> tuple[list[str], dict[str, int]]: return seglabel_list, seglabel2code + def _build_box_class_map(self) -> dict[str, int]: + """Build {class_name: index} for box annotations, alphabetically sorted.""" + from datamint.entities.annotations import AnnotationType + class_names: set[str] = set() + for anns in self.resource_annotations: + for ann in anns: + if getattr(ann, 'annotation_type', None) == AnnotationType.SQUARE and ann.identifier: + class_names.add(ann.identifier) + return {name: idx for idx, name in enumerate(sorted(class_names))} + + def _load_boxes( + self, + annotations: 'Sequence[Annotation]', + ) -> tuple['Tensor', 'Tensor']: + """Extract box tensors from square annotations. + + Returns: + Tuple of (boxes, box_labels) where boxes is (N, 4) float32 in + pascal_voc pixel coords and box_labels is (N,) int64 class indices. + """ + valid: list[tuple[float, float, float, float, str]] = [] + for ann in annotations: + if not ann.identifier: + _LOGGER.warning("Skipping box annotation with no identifier.") + continue + geometry = getattr(ann, 'geometry', None) + if geometry is None: + continue + x1, y1, _ = geometry.point1 + x2, y2, _ = geometry.point2 + x1, y1, x2, y2 = float(x1), float(y1), float(x2), float(y2) + if x2 <= x1 or y2 <= y1: + _LOGGER.warning( + "Skipping degenerate box (x2<=x1 or y2<=y1): (%s, %s, %s, %s)", + x1, y1, x2, y2, + ) + continue + valid.append((x1, y1, x2, y2, ann.identifier)) + + if not valid: + return torch.zeros((0, 4), dtype=torch.float32), torch.zeros((0,), dtype=torch.int64) + + boxes = torch.tensor([(x1, y1, x2, y2) for x1, y1, x2, y2, _ in valid], dtype=torch.float32) + labels = torch.tensor( + [self.box_class_map.get(name, 0) for _, _, _, _, name in valid], + dtype=torch.int64, + ) + return boxes, labels + def _process_segmentation_group(self, groups: dict) -> tuple[list[str], dict[str, int]]: """Get segmentation labels from the server.""" try: @@ -979,28 +1037,41 @@ def __getitem__(self, index: int) -> dict[str, Any]: if isinstance(img, np.ndarray): img = self._preprocess_image_array(img) annotations = result['annotations'] - # resource = result['resource'] - # _LOGGER.debug(f"Loaded image {resource.filename} with shape {img.shape}") - # Process segmentations + # Load all requested annotation targets + targets: dict[str, Any] = {} + seg_labels = None + if self.return_segmentations: - seg_anns = AnnotationProcessor.filter_annotations(annotations, - type='segmentation', - scope='all') + seg_anns = AnnotationProcessor.filter_annotations(annotations, type='segmentation', scope='all') segmentations, seg_labels, _ = self.annotation_processor.load_segmentations(seg_anns) - # Apply albumentations if present - if self.alb_transform: - aug_result = self.apply_alb_transform(img, segmentations) - img = aug_result['image'] - result['image'] = img - segmentations = aug_result['segmentations'] + targets['masks'] = segmentations + + if self.return_boxes: + box_anns = [ann for ann in annotations if getattr(ann, 'annotation_type', None) == 'square'] + boxes, box_labels_tensor = self._load_boxes(box_anns) + targets['boxes'] = boxes + targets['box_labels'] = box_labels_tensor + + # Apply albumentations to all targets at once + if self.alb_transform: + aug = self.apply_alb_transform(img, targets) + img = aug.pop('image') + targets.update(aug) - segmentations, seg_labels = self._process_segmentations(segmentations, seg_labels, - output_shape=img.shape[1:]) + result['image'] = img - result['segmentations'] = segmentations + # Post-process and write to result + if self.return_segmentations: + masks = targets.get('masks', {}) + masks, seg_labels = self._process_segmentations(masks, seg_labels, output_shape=img.shape[1:]) + result['masks'] = masks if seg_labels: - result['seg_labels'] = seg_labels + result['mask_labels'] = seg_labels + + if self.return_boxes: + result['boxes'] = targets.get('boxes', torch.zeros((0, 4), dtype=torch.float32)) + result['box_labels'] = targets.get('box_labels', torch.zeros((0,), dtype=torch.int64)) # Process image-level labels result['image_labels'] = self._extract_image_labels(annotations) @@ -1012,17 +1083,19 @@ def __getitem__(self, index: int) -> dict[str, Any]: def apply_alb_transform( self, img: np.ndarray, - segmentations: dict[str, np.ndarray] + targets: dict[str, Any], ) -> dict[str, Any]: - """Apply albumentations transform to image and masks. + """Apply albumentations transform to image and annotation targets. - Returns: - Dict with transformed 'image' and 'segmentations' (dict). - It is recommended that 'image' has shape (C, depth, H, W) - and each segmentation of 'segmentations' has shape (num_instances, depth, H, W), so that - common downstream processing can be applied. - If not, please override :py:meth:`_process_segmentations` accordingly. + Args: + img: Image array. + targets: Dict of annotation targets to transform. May contain: + - ``'masks'``: per-annotator segmentation masks + - ``'boxes'``: (N, 4) float32 tensor in pascal_voc pixel coords + - ``'box_labels'``: (N,) int64 tensor of class indices + Returns: + Dict with ``'image'`` key plus the same target keys, all transformed. """ pass @@ -1052,6 +1125,8 @@ def build_mlflow_dataset(self) -> 'DatamintMLflowDataset': project_id = getattr(project, 'id', 'unknown') if project is not None else 'unknown' extra_params = { + 'return_segmentations': self.return_segmentations, + 'return_boxes': self.return_boxes, 'return_as_semantic_segmentation': self.return_as_semantic_segmentation, 'semantic_seg_merge_strategy': str(self.semantic_seg_merge_strategy), 'include_unannotated': self.include_unannotated, diff --git a/datamint/dataset/detection_dataset.py b/datamint/dataset/detection_dataset.py deleted file mode 100644 index f68dfe70..00000000 --- a/datamint/dataset/detection_dataset.py +++ /dev/null @@ -1,227 +0,0 @@ -""" -DetectionDataset - Dataset for object detection tasks. - -Loads images and bounding box annotations, returning detection-ready tensors -in pascal_voc (x1, y1, x2, y2) pixel coordinate format. -""" -import logging -from typing import Any -from typing_extensions import override - -import numpy as np -import torch - -from .image_dataset import ImageDataset - -_LOGGER = logging.getLogger(__name__) - - -class DetectionDataset(ImageDataset): - """Dataset for 2D object detection. - - Returns items as dicts with keys ``'image'`` (C×H×W tensor), - ``'boxes'`` (N×4 float32 tensor, pascal_voc pixel coords), - ``'labels'`` (N int64 tensor), ``'resource_id'`` (str), and - ``'identifiers'`` (list[str], for v2 instance tracking). - - The class map is built alphabetically from all ``identifier`` values - found on box annotations across the project, ensuring a stable - name→index mapping between runs. - - Use :func:`detection_collate_fn` as the DataLoader ``collate_fn``. - """ - - def __init__(self, *args: Any, **kwargs: Any) -> None: - kwargs['return_segmentations'] = False - super().__init__(*args, **kwargs) - - @override - def _setup_dataset(self) -> None: - super()._setup_dataset() - self._class_map: dict[str, int] = self._build_class_map() - - def _build_class_map(self) -> dict[str, int]: - """Return alphabetically-sorted ``{class_name: index}`` from all box annotations.""" - class_names: set[str] = set() - for anns in self.resource_annotations: - for ann in anns: - if getattr(ann, 'annotation_type', None) == 'square': - name = ann.identifier - if name: - class_names.add(name) - return {name: idx for idx, name in enumerate(sorted(class_names))} - - def _load_image(self, index: int) -> np.ndarray: - """Load the image at *index* as a (H, W, C) float32 numpy array. - - Delegates to the parent's ``_get_raw_item`` so that format-specific - handling (DICOM, NIfTI, WebP, uint16 normalisation, …) is not - duplicated here. - """ - - raw = self._get_raw_item(index) - img = raw['image'] # (C, N, H, W) - if img.ndim == 4 and img.shape[1] == 1: - img = img.squeeze(1) # (C, H, W) - elif img.ndim == 4: - img = img[:, 0] # first frame (DetectionDataset is 2D-only) - img = self._preprocess_image_array(img) - return np.ascontiguousarray(img.transpose(1, 2, 0)) # (H, W, C) - - def _fetch_boxes( - self, - resource, - frame_index: int | None = None, - ) -> list[tuple[float, float, float, float]]: - """Return ``(x1, y1, x2, y2)`` pixel-coordinate tuples for box annotations. - - Args: - resource: Resource with a ``fetch_annotations`` method. - frame_index: When given, only boxes whose ``frame_index`` matches - are returned. - """ - anns = resource.fetch_annotations(annotation_type='square') - boxes: list[tuple[float, float, float, float]] = [] - for ann in anns: - if frame_index is not None and ann.frame_index != frame_index: - continue - geometry = getattr(ann, 'geometry', None) - if geometry is None: - continue - x1, y1, _ = geometry.point1 - x2, y2, _ = geometry.point2 - boxes.append((float(x1), float(y1), float(x2), float(y2))) - return boxes - - @override - def apply_alb_transform( - self, - img: np.ndarray, - boxes: list[tuple[float, float, float, float]], - labels: list[int], - identifiers: list[str] | None = None, - ) -> dict[str, Any]: - """Apply albumentations transform with bounding box support. - - The transform must be built with - ``A.BboxParams(format='pascal_voc', label_fields=['labels', 'identifiers'])`` - so that albumentations keeps identifiers aligned with their boxes when - spatial transforms remove or reorder boxes. - """ - if self.alb_transform is None: - raise ValueError("alb_transform is not set") - - _ids = identifiers if identifiers is not None else ['' for _ in boxes] - aug = self.alb_transform(image=img, bboxes=boxes, labels=labels, identifiers=_ids) - aug_img = aug['image'] - aug_boxes: list = list(aug['bboxes']) - aug_labels: list = list(aug['labels']) - aug_identifiers: list = list(aug.get('identifiers', _ids[:len(aug_boxes)])) - - if isinstance(aug_img, np.ndarray): - aug_img = torch.from_numpy( - np.ascontiguousarray(aug_img.transpose(2, 0, 1)).astype(np.float32) - ) - - if aug_boxes: - boxes_tensor = torch.tensor(aug_boxes, dtype=torch.float32) - labels_tensor = torch.tensor(aug_labels, dtype=torch.int64) - else: - boxes_tensor = torch.zeros((0, 4), dtype=torch.float32) - labels_tensor = torch.zeros((0,), dtype=torch.int64) - - return { - 'image': aug_img, - 'boxes': boxes_tensor, - 'labels': labels_tensor, - 'identifiers': aug_identifiers, - } - - @override - def __getitem__(self, index: int) -> dict[str, Any]: - resource = self.resources[index] - img_hwc = self._load_image(index) - - # Use pre-fetched annotations - anns = self.resource_annotations[index] - box_anns = [ - ann for ann in anns - if getattr(ann, 'annotation_type', None) == 'square' - and getattr(ann, 'geometry', None) is not None - ] - - # Skip boxes without an identifier - unnamed = [ann for ann in box_anns if not ann.identifier] - if unnamed: - _LOGGER.warning( - "Skipping %d box annotation(s) on resource '%s' that have no identifier. " - "Set a class name on each BoxAnnotation before training.", - len(unnamed), resource.id, - ) - box_anns = [ann for ann in box_anns if ann.identifier] - - raw_boxes = [ - (float(ann.geometry.point1[0]), float(ann.geometry.point1[1]), - float(ann.geometry.point2[0]), float(ann.geometry.point2[1]), - ann) - for ann in box_anns - ] - degenerate = [(x1, y1, x2, y2) for x1, y1, x2, y2, _ in raw_boxes if x2 <= x1 or y2 <= y1] - if degenerate: - _LOGGER.warning( - "Skipping %d degenerate box(es) on resource '%s' (x2<=x1 or y2<=y1): %s", - len(degenerate), resource.id, degenerate, - ) - box_anns = [ann for x1, y1, x2, y2, ann in raw_boxes if x2 > x1 and y2 > y1] - boxes = [ - (float(ann.geometry.point1[0]), float(ann.geometry.point1[1]), - float(ann.geometry.point2[0]), float(ann.geometry.point2[1])) - for ann in box_anns - ] - labels = [self._class_map.get(ann.identifier, 0) for ann in box_anns] - identifiers = [ann.identifier for ann in box_anns] - - if self.alb_transform is not None: - aug = self.apply_alb_transform(img_hwc, boxes, labels, identifiers) - img_tensor = aug['image'] - boxes_tensor = aug['boxes'] - labels_tensor = aug['labels'] - identifiers = aug['identifiers'] - else: - img_tensor = torch.from_numpy( - np.ascontiguousarray(img_hwc.transpose(2, 0, 1)).astype(np.float32) - ) - if boxes: - boxes_tensor = torch.tensor(boxes, dtype=torch.float32) - labels_tensor = torch.tensor(labels, dtype=torch.int64) - else: - boxes_tensor = torch.zeros((0, 4), dtype=torch.float32) - labels_tensor = torch.zeros((0,), dtype=torch.int64) - - return { - 'image': img_tensor, - 'boxes': boxes_tensor, - 'labels': labels_tensor, - 'resource_id': resource.id, - 'identifiers': identifiers, - } - - @override - def __repr__(self) -> str: - base = super(ImageDataset, self).__repr__() - return f"DetectionDataset\n{base}" - - -def detection_collate_fn(batch: list[dict]) -> dict: - """Collate a list of detection items into a batch. - - Images are stacked into a single tensor. Boxes and labels are kept as - lists because each image may have a different number of annotations. - """ - return { - 'image': torch.stack([item['image'] for item in batch]), - 'boxes': [item['boxes'] for item in batch], - 'labels': [item['labels'] for item in batch], - 'resource_id': [item['resource_id'] for item in batch], - 'identifiers': [item['identifiers'] for item in batch], - } diff --git a/datamint/dataset/image_dataset.py b/datamint/dataset/image_dataset.py index 84efd8b3..74984975 100644 --- a/datamint/dataset/image_dataset.py +++ b/datamint/dataset/image_dataset.py @@ -31,19 +31,19 @@ def __getitem__(self, index: int) -> dict[str, Any]: result['image'] = img if self.return_segmentations: - segmentations = result['segmentations'] - # convert segmentations shape to expected format: + masks = result['masks'] + # convert masks shape to expected format: # if semantic and no merge: dict[author -> (num_labels+1, H, W)] # if instance and no merge: dict[author -> (num_instances, H, W)] # if merged (semantic only): (num_labels+1, H, W) - if isinstance(segmentations, (Tensor, np.ndarray)): - _LOGGER.debug("squeezing merged segmentations of shape %s", segmentations.shape) - segmentations = segmentations.squeeze(1) + if isinstance(masks, (Tensor, np.ndarray)): + _LOGGER.debug("squeezing merged masks of shape %s", masks.shape) + masks = masks.squeeze(1) else: - for author in segmentations: - segmentations[author] = segmentations[author].squeeze(1) + for author in masks: + masks[author] = masks[author].squeeze(1) - result['segmentations'] = segmentations + result['masks'] = masks return result @@ -51,7 +51,7 @@ def __getitem__(self, index: int) -> dict[str, Any]: def apply_alb_transform( self, img: np.ndarray, - segmentations: dict[str, np.ndarray], + targets: dict[str, Any], ) -> dict[str, Any]: if self.alb_transform is None: raise ValueError("alb_transform is not set") @@ -63,44 +63,86 @@ def apply_alb_transform( raise ValueError(f"Expected 3D image array (C, H, W) or (C, 1, H, W), got shape {img.shape}") # transpose to (H, W, C) - img = np.transpose(img, (1, 2, 0)) - - replay_alb_transf = albumentations.ReplayCompose([self.alb_transform]) - - aug_data = replay_alb_transf(image=img) # First call - replay_data = aug_data['replay'] - aug_img = aug_data['image'] - - aug_segmentations = {} - for author, segs in segmentations.items(): + img_hwc = np.transpose(img, (1, 2, 0)) + + # Flatten masks from all authors into a single list for one transform call + segmentations = targets.get('masks', {}) + author_order = list(segmentations.keys()) + author_counts: dict[str, int] = {} + all_masks: list[np.ndarray] = [] + for author in author_order: + segs = segmentations[author] if segs.ndim == 4 and segs.shape[1] == 1: - segs = segs.squeeze(1) # (num_instances, 1, H, W) -> (num_instances, H, W) - # if segs.dtype == bool: - # segs = segs.astype(np.uint8) - aug_segs = replay_alb_transf.replay(replay_data, masks=segs)['masks'] - # store back with original shape - if segs.ndim == 3: - aug_segs = aug_segs[:, np.newaxis, :, :] # (num_instances, H, W) -> (num_instances, 1, H, W) - aug_segmentations[author] = aug_segs + segs = segs.squeeze(1) # (N, 1, H, W) -> (N, H, W) + author_counts[author] = len(segs) + all_masks.extend(list(segs)) + + alb_kwargs: dict[str, Any] = {'image': img_hwc} + if all_masks: + alb_kwargs['masks'] = all_masks + + boxes_tensor = targets.get('boxes') + box_labels_tensor = targets.get('box_labels') + if boxes_tensor is not None: + alb_kwargs['bboxes'] = boxes_tensor.tolist() if isinstance(boxes_tensor, torch.Tensor) else list(boxes_tensor) + alb_kwargs['box_labels'] = box_labels_tensor.tolist() if isinstance(box_labels_tensor, torch.Tensor) else list(box_labels_tensor) + + aug = self.alb_transform(**alb_kwargs) + aug_img = aug['image'] + result: dict[str, Any] = {} + + # Reconstruct per-author masks + if all_masks: + aug_masks_flat = aug['masks'] + start = 0 + aug_segmentations: dict[str, np.ndarray] = {} + for author in author_order: + count = author_counts[author] + segs_aug = np.array(aug_masks_flat[start:start + count]) + segs_aug = segs_aug[:, np.newaxis, :, :] # (N, H, W) -> (N, 1, H, W) + aug_segmentations[author] = segs_aug + start += count + result['masks'] = aug_segmentations + + # Reconstruct boxes + if boxes_tensor is not None: + aug_bboxes: list = list(aug.get('bboxes', [])) + aug_box_labels: list = list(aug.get('box_labels', [])) + if aug_bboxes: + result['boxes'] = torch.tensor(aug_bboxes, dtype=torch.float32) + result['box_labels'] = torch.tensor(aug_box_labels, dtype=torch.int64) + else: + result['boxes'] = torch.zeros((0, 4), dtype=torch.float32) + result['box_labels'] = torch.zeros((0,), dtype=torch.int64) # transpose back to (C, H, W) if isinstance(aug_img, np.ndarray): aug_img = np.transpose(aug_img, (2, 0, 1)) elif isinstance(aug_img, torch.Tensor): - # shape is (C, H, W), assuming albumentation transformation changed it - if aug_img.shape[0] == img.shape[-1]: # if C is in dim 0 + if aug_img.shape[0] == img_hwc.shape[-1]: # C already in dim 0 aug_img = aug_img.permute(0, 1, 2) else: aug_img = aug_img.permute(2, 0, 1) # back to (C, 1, H, W) aug_img = aug_img[:, np.newaxis, :, :] - - return { - 'image': aug_img, - 'segmentations': aug_segmentations, - } + result['image'] = aug_img + return result @override def __repr__(self) -> str: base = super(VolumeDataset, self).__repr__() return f"ImageDataset\n{base}" + + +def detection_collate_fn(batch: list[dict]) -> dict: + """Collate a list of detection items into a batch. + + Images are stacked into a single tensor. Boxes and box_labels are kept as + lists because each image may have a different number of annotations. + """ + import torch + return { + 'image': torch.stack([item['image'] for item in batch]), + 'boxes': [item['boxes'] for item in batch], + 'box_labels': [item['box_labels'] for item in batch], + } diff --git a/datamint/dataset/multiframe_dataset.py b/datamint/dataset/multiframe_dataset.py index 8d461039..2c8a67af 100644 --- a/datamint/dataset/multiframe_dataset.py +++ b/datamint/dataset/multiframe_dataset.py @@ -62,29 +62,31 @@ def _get_raw_item(self, index: int) -> dict[str, Any]: def apply_alb_transform( self, img: np.ndarray, - segmentations: dict[str, np.ndarray], + targets: dict[str, Any], ) -> dict[str, Any]: """Apply albumentations transform to 4D image and masks. Args: img: Image array of shape ``(C, depth, H, W)``. - segmentations: Dict of author -> mask arrays of shape - ``(#instances, depth, H, W)``. + targets: Dict of annotation targets. Supports ``'masks'`` (dict of + author -> mask arrays of shape ``(#instances, depth, H, W)``). Returns: - Dict with transformed ``'image'`` and ``'segmentations'``. + Dict with transformed ``'image'`` and ``'masks'``. """ if self.alb_transform is None: raise ValueError("alb_transform is not set") if img.ndim != 4: raise ValueError(f"Expected 4D image array (C, depth, H, W), got shape {img.shape}") + segmentations = targets.get('masks', {}) + # transpose to (depth, H, W, C) img = np.transpose(img, (1, 2, 3, 0)) replay_alb_transf = albumentations.ReplayCompose([self.alb_transform]) _LOGGER.debug( - f'before alb transform image shape: {img.shape} | segmentations shape: {[segmentations[a].shape for a in segmentations]}') + f'before alb transform image shape: {img.shape} | masks shape: {[segmentations[a].shape for a in segmentations]}') aug_data = replay_alb_transf(volume=img) # First call replay_data = aug_data['replay'] @@ -108,7 +110,7 @@ def apply_alb_transform( aug_img = aug_img.permute(3, 0, 1, 2) _LOGGER.debug(f"augmented image tensor shape after permute: {aug_img.shape}") - return { - 'image': aug_img, - 'segmentations': aug_segmentations, - } + result: dict[str, Any] = {'image': aug_img} + if aug_segmentations: + result['masks'] = aug_segmentations + return result diff --git a/datamint/lightning/trainers/detection_trainer.py b/datamint/lightning/trainers/detection_trainer.py index 4c60ee03..83998a60 100644 --- a/datamint/lightning/trainers/detection_trainer.py +++ b/datamint/lightning/trainers/detection_trainer.py @@ -6,7 +6,7 @@ from torch import nn -from datamint.dataset.detection_dataset import DetectionDataset, detection_collate_fn +from datamint.dataset.image_dataset import ImageDataset, detection_collate_fn from datamint.lightning.datamodule import DatamintDataModule from .base_trainer import BaseTrainer @@ -19,7 +19,7 @@ class DetectionTrainer(BaseTrainer): Provides shared defaults for all detection models: - * **Dataset** – :class:`~datamint.dataset.DetectionDataset` + * **Dataset** – :class:`~datamint.dataset.ImageDataset` with ``return_boxes=True`` * **Collate** – :func:`~datamint.dataset.detection_collate_fn` (variable-length boxes) * **Metrics** – Mean Average Precision (torchmetrics) * **Monitor** – ``val/map`` (maximise) @@ -32,12 +32,12 @@ def _build_dataset( self, project: 'str | Project', **kwargs: Any, - ) -> DetectionDataset: - return DetectionDataset(project=project, **kwargs) + ) -> ImageDataset: + return ImageDataset(project=project, return_boxes=True, **kwargs) def _build_datamodule( self, - dataset: DetectionDataset, + dataset: ImageDataset, train_transform: Any, eval_transform: Any, ) -> DatamintDataModule: diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_module.py b/datamint/lightning/trainers/lightning_modules/segmentation_module.py index 0958b2b5..4177e1c3 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_module.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_module.py @@ -63,7 +63,7 @@ def forward(self, x: Tensor) -> Tensor: def _common_step(self, batch: dict, stage: str) -> Tensor: images = batch['image'] # shape (B, C, H, W) - masks = batch['segmentations'][:, 1:] # exclude background channel + masks = batch['masks'][:, 1:] # exclude background channel # masks.shape is (B, C, H, W) where C is num_classes (excluding background) logits = self(images) @@ -140,7 +140,7 @@ def _compute_sample_confidence(self, logits: Tensor) -> dict[str, Tensor]: def _compute_sample_metrics(self, logits: Tensor, batch: dict) -> dict[str, Tensor]: """Per-sample IoU and Dice.""" - masks = batch['segmentations'][:, 1:].float() + masks = batch['masks'][:, 1:].float() preds = (logits > 0).float() intersection = (preds * masks).sum(dim=[1, 2, 3]) union = ((preds + masks) > 0).float().sum(dim=[1, 2, 3]) diff --git a/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py b/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py index e2c38923..741294c5 100644 --- a/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py +++ b/datamint/lightning/trainers/lightning_modules/segmentation_modules/unetrpp.py @@ -538,7 +538,7 @@ def _to_float_tensor(x: 'Tensor | np.ndarray') -> Tensor: def training_step(self, batch: dict, batch_idx: int) -> Tensor: device = next(self.parameters()).device images = self._to_float_tensor(batch['image']).to(device) - masks = self._to_float_tensor(batch['segmentations'])[:, 1:].to(device) + masks = self._to_float_tensor(batch['masks'])[:, 1:].to(device) images, masks = self._random_crop_3d(images, masks) @@ -565,7 +565,7 @@ def test_step(self, batch: dict, batch_idx: int) -> Tensor | None: def _eval_step(self, batch: dict, stage: str) -> Tensor | None: device = next(self.parameters()).device images = self._to_float_tensor(batch['image']).to(device) - masks = self._to_float_tensor(batch['segmentations'])[:, 1:].to(device) + masks = self._to_float_tensor(batch['masks'])[:, 1:].to(device) logits = self._sliding_window_inference(images) loss = self.criterion(logits, masks) if self.criterion else None diff --git a/datamint/lightning/trainers/specialized/yolox.py b/datamint/lightning/trainers/specialized/yolox.py index 3529b70e..70c6665b 100644 --- a/datamint/lightning/trainers/specialized/yolox.py +++ b/datamint/lightning/trainers/specialized/yolox.py @@ -26,10 +26,10 @@ class YOLOXTrainer(DetectionTrainer): results = trainer.fit() Args: - dataset: A pre-built :class:`~datamint.dataset.DetectionDataset`. + dataset: A pre-built :class:`~datamint.dataset.ImageDataset`. Mutually exclusive with *project*. project: Project name or :class:`~datamint.entities.Project` object. - A :class:`~datamint.dataset.DetectionDataset` is created automatically. + A :class:`~datamint.dataset.ImageDataset` is created automatically. model_size: YOLOX size variant — ``'nano'``, ``'tiny'``, ``'s'``, ``'m'``, ``'l'``, or ``'x'``. Defaults to ``'s'``, which balances speed and accuracy for most medical imaging tasks. @@ -53,7 +53,7 @@ class YOLOXTrainer(DetectionTrainer): split_as_of_timestamp: Historical timestamp for reproducible splits. auto_deploy_adapter: Auto-log a deploy adapter after training. trainer_kwargs: Extra kwargs forwarded to :class:`lightning.Trainer`. - dataset_kwargs: Extra kwargs forwarded to :class:`~datamint.dataset.DetectionDataset`. + dataset_kwargs: Extra kwargs forwarded to :class:`~datamint.dataset.ImageDataset`. """ def __init__( @@ -155,7 +155,7 @@ def _build_model( loss_fn: nn.Module | None, metrics: dict, ) -> 'DatamintLightningModule': - num_classes = len(self.dataset._class_map) + num_classes = len(self.dataset.box_class_map) if num_classes == 0: raise ValueError( "No box annotation classes found in the dataset. " @@ -163,7 +163,7 @@ def _build_model( "Make sure your project has BoxAnnotation objects with an identifier set." ) - class_names = sorted(self.dataset._class_map, key=self.dataset._class_map.__getitem__) + class_names = sorted(self.dataset.box_class_map, key=self.dataset.box_class_map.__getitem__) return YOLOXModule( num_classes=num_classes, model_size=self.model_size, diff --git a/tests/test_dataset_patient_split.py b/tests/test_dataset_patient_split.py index 9d399ae4..9c223b17 100644 --- a/tests/test_dataset_patient_split.py +++ b/tests/test_dataset_patient_split.py @@ -17,8 +17,8 @@ def _reinit_api(self) -> None: def _get_raw_item(self, index: int) -> dict: return {'image': None, 'metainfo': {}, 'annotations': []} - def apply_alb_transform(self, img, segmentations): - return {'image': img, 'segmentations': segmentations} + def apply_alb_transform(self, img, targets): + return {'image': img, **targets} class _MockResource: diff --git a/tests/test_detection_trainer.py b/tests/test_detection_trainer.py index 6760f828..b5919fc0 100644 --- a/tests/test_detection_trainer.py +++ b/tests/test_detection_trainer.py @@ -6,7 +6,7 @@ from albumentations.pytorch import ToTensorV2 from datamint.lightning.trainers.detection_trainer import DetectionTrainer -from datamint.dataset.detection_dataset import DetectionDataset, detection_collate_fn +from datamint.dataset.image_dataset import ImageDataset, detection_collate_fn from datamint.lightning.datamodule import DatamintDataModule @@ -33,20 +33,27 @@ def _bare_trainer(): # -- _build_dataset tests -def test_build_dataset_returns_detection_dataset(): +def test_build_dataset_returns_image_dataset(): trainer = _bare_trainer() - with patch.object(DetectionDataset, '__init__', return_value=None): + with patch.object(ImageDataset, '__init__', return_value=None): ds = trainer._build_dataset(project='thyroid') - assert isinstance(ds, DetectionDataset) + assert isinstance(ds, ImageDataset) def test_build_dataset_passes_project_kwarg(): trainer = _bare_trainer() - with patch.object(DetectionDataset, '__init__', return_value=None) as mock_init: + with patch.object(ImageDataset, '__init__', return_value=None) as mock_init: trainer._build_dataset(project='thyroid') assert mock_init.call_args.kwargs.get('project') == 'thyroid' +def test_build_dataset_sets_return_boxes(): + trainer = _bare_trainer() + with patch.object(ImageDataset, '__init__', return_value=None) as mock_init: + trainer._build_dataset(project='thyroid') + assert mock_init.call_args.kwargs.get('return_boxes') is True + + # -- _build_datamodule tests def test_datamodule_receives_detection_collate_fn(): diff --git a/tests/test_yolox_trainer.py b/tests/test_yolox_trainer.py index 6c79b08d..72480f37 100644 --- a/tests/test_yolox_trainer.py +++ b/tests/test_yolox_trainer.py @@ -5,14 +5,14 @@ import albumentations as A from datamint.lightning.trainers.specialized.yolox import YOLOXTrainer -from datamint.dataset.detection_dataset import DetectionDataset +from datamint.dataset.image_dataset import ImageDataset from datamint.lightning.trainers.lightning_modules.detection_modules.yolox_module import YOLOXModule @pytest.fixture() def trainer(): - ds = MagicMock(spec=DetectionDataset) - ds._class_map = {'cyst': 0, 'nodule': 1} + ds = MagicMock(spec=ImageDataset) + ds.box_class_map = {'cyst': 0, 'nodule': 1} t = YOLOXTrainer.__new__(YOLOXTrainer) t.dataset = ds t.model_size = 's' @@ -53,7 +53,7 @@ def test_build_model_forwards_thresholds(trainer): def test_build_model_raises_when_no_classes(trainer): - trainer.dataset._class_map = {} + trainer.dataset.box_class_map = {} with pytest.raises(ValueError, match='No box annotation classes'): trainer._build_model(loss_fn=None, metrics={}) @@ -61,7 +61,7 @@ def test_build_model_raises_when_no_classes(trainer): # -- image_size normalization test -- def test_image_size_int_becomes_tuple(): - with patch.object(DetectionDataset, '__init__', return_value=None): + with patch.object(ImageDataset, '__init__', return_value=None): t = YOLOXTrainer.__new__(YOLOXTrainer) t.image_size = 416 if isinstance(t.image_size, int):