Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 45 additions & 16 deletions src/virtual_stain_flow/datasets/base_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@ def __init__(
input_channel_keys: Optional[Union[str, Sequence[str]]] = None,
target_channel_keys: Optional[Union[str, Sequence[str]]] = None,
transforms: Optional[Sequence[LoggableTransform]] = None,
input_transforms: Optional[Sequence[LoggableTransform]] = None,
target_transforms: Optional[Sequence[LoggableTransform]] = None,
cache_capacity: Optional[int] = None,
file_state: Optional[FileState] = None,
):
Expand All @@ -49,6 +51,10 @@ def __init__(
:param target_channel_keys: Keys for target channels in the file index.
:param transforms: Optional sequence of LoggableTransform objects to apply
to the images before returning them.
:param input_transforms: Optional sequence of LoggableTransform objects to
apply to the input image only, after `transforms`.
:param target_transforms: Optional sequence of LoggableTransform objects to
apply to the target image only, after `transforms`.
:param cache_capacity: Optional capacity for caching loaded images.
When set to None, default caching behavior of caching at most
`file_index.shape[0]` images is used. When set to -1, unbounded
Expand Down Expand Up @@ -86,6 +92,33 @@ def __init__(
raise ValueError("All transforms must be instances of LoggableTransform.")
self.transforms = transforms

self.input_transforms = self._normalize_transforms(
input_transforms, "input_transforms"
)
self.target_transforms = self._normalize_transforms(
target_transforms, "target_transforms"
)

def _normalize_transforms(
self,
transforms: Optional[Sequence[LoggableTransform]],
name: str,
) -> Sequence[LoggableTransform]:
Comment thread
wli51 marked this conversation as resolved.
"""
Normalize and validate a sequence of transforms.

:param transforms: Sequence of LoggableTransform objects or None.
:param name: Name of the transform sequence for error messages.
:return: Normalized sequence of LoggableTransform objects.
"""
if not isinstance(transforms, Sequence):
transforms = [transforms] if transforms else []
if not all(isinstance(t, LoggableTransform) for t in transforms):
raise ValueError(
f"All {name} must be instances of LoggableTransform."
)
return transforms

def get_raw_item(
self,
idx: int
Expand Down Expand Up @@ -119,28 +152,24 @@ def __len__(self) -> int:
"""
return len(self.manifest)

def _apply_transforms(
self,
image: np.ndarray,
) -> np.ndarray:
"""
Applies the sequence of transforms to the input image.

:param image: Input image as a numpy array.
:return: Transformed image as a numpy array.
"""
for transform in self.transforms:
image = transform.apply(img=image)
return image

def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Overridden Dataset `__getitem__` method so class works with torch DataLoader.
"""
input_image_raw, target_image_raw = self.get_raw_item(idx)

return (torch.from_numpy(self._apply_transforms(input_image_raw)).float(),
torch.from_numpy(self._apply_transforms(target_image_raw)).float())
input_image = input_image_raw
for transform in (*self.transforms, *self.input_transforms):
input_image = transform.apply(img=input_image)

target_image = target_image_raw
for transform in (*self.transforms, *self.target_transforms):
target_image = transform.apply(img=target_image)

return (
torch.from_numpy(input_image).float(),
torch.from_numpy(target_image).float()
)

@property
def pil_image_mode(self) -> str:
Expand Down
25 changes: 23 additions & 2 deletions src/virtual_stain_flow/datasets/crop_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
crop_dataset.py
"""

from typing import Any, Dict, List, Sequence, Optional, Tuple, Union, Type
from typing import Any, Dict, List, Sequence, Optional, Tuple, Union

import numpy as np
import pandas as pd
Expand Down Expand Up @@ -30,6 +30,8 @@ def __init__(
input_channel_keys: Optional[Union[str, Sequence[str]]] = None,
target_channel_keys: Optional[Union[str, Sequence[str]]] = None,
transforms: Optional[Sequence[LoggableTransform]] = None,
input_transforms: Optional[Sequence[LoggableTransform]] = None,
target_transforms: Optional[Sequence[LoggableTransform]] = None,
crop_file_state: Optional[CropFileState] = None,
):
"""
Expand All @@ -53,6 +55,10 @@ def __init__(
:param input_channel_keys: Keys for input channels in the file index.
:param target_channel_keys: Keys for target channels in the file index.
:param transforms: Optional sequence of transformations to apply to the images.
:param input_transforms: Optional sequence of LoggableTransform objects to
apply to the input image only, after `transforms`.
:param target_transforms: Optional sequence of LoggableTransform objects to
apply to the target image only, after `transforms`.
:param crop_file_state: Optional pre-initialized CropFileState object. If provided,
it takes precedence over `file_index` and `crop_specs`. Intended
to be used by only .from_config class method and similar deserialization
Expand Down Expand Up @@ -84,6 +90,13 @@ def __init__(
raise ValueError("All transforms must be instances of LoggableTransform.")
self.transforms = transforms

self.input_transforms = self._normalize_transforms(
input_transforms, "input_transforms"
)
self.target_transforms = self._normalize_transforms(
target_transforms, "target_transforms"
)

@property
def pil_image_mode(self) -> str:
return self.manifest.pil_image_mode
Expand Down Expand Up @@ -157,7 +170,9 @@ def from_base_dataset(
cls,
base_dataset: BaseImageDataset,
transforms: Optional[Sequence[LoggableTransform]] = None,
how: Type[CropGenerator] = generate_center_crops,
input_transforms: Optional[Sequence[LoggableTransform]] = None,
target_transforms: Optional[Sequence[LoggableTransform]] = None,
how: CropGenerator = generate_center_crops,
**kwargs: Any
) -> 'CropImageDataset':
"""
Expand All @@ -166,6 +181,10 @@ def from_base_dataset(
:param base_dataset: The BaseImageDataset to convert.
:param how: A function that generates crop specifications from the base dataset.
Default is `generate_center_crops`.
:param input_transforms: Optional sequence of LoggableTransform objects to
apply to the input image only, after `transforms`.
:param target_transforms: Optional sequence of LoggableTransform objects to
apply to the target image only, after `transforms`.
:param kwargs: Additional keyword arguments for the `how` function.
"""

Expand All @@ -177,6 +196,8 @@ def from_base_dataset(
return cls(
file_index=base_dataset.file_index,
transforms=transforms,
input_transforms=input_transforms,
target_transforms=target_transforms,
crop_specs=crop_specs,
pil_image_mode=base_dataset.pil_image_mode,
input_channel_keys=base_dataset.input_channel_keys,
Expand Down
108 changes: 16 additions & 92 deletions src/virtual_stain_flow/datasets/ds_engine/crop_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,98 +3,22 @@

Utilities for generating crop coordinates from BaseImageDataset objects.
Designed for easy creation of CropImageDataset instances.
Made Facade to account for increased complexity and future expansion.
"""

from typing import Dict, List, Tuple, Any, Protocol

from ..base_dataset import BaseImageDataset
from .ds_utils import (
_get_active_channels,
_validate_same_dimensions_across_channels
from .crop_generators.protocol import CropSpec, CropMap, CropGenerator
from .crop_generators.center import (
generate_center_crops,
_compute_center_crop,
)

CropSpec = Tuple[Tuple[int, int], int, int]
CropMap = Dict[int, List[CropSpec]]


class CropGenerator(Protocol):
"""
Protocol for crop generator functions.
"""
def __call__(
self,
dataset: BaseImageDataset,
**kwargs: Any
) -> CropMap:
pass


def _compute_center_crop(
image_width: int,
image_height: int,
crop_size: int
) -> Tuple[int, int]:
"""
Compute top-left (x, y) coordinates for a center crop.

:param image_width: Width of the source image.
:param image_height: Height of the source image.
:param crop_size: Size of the square crop (width and height).
:return: Tuple of (x, y) for top-left corner of center crop.
:raises ValueError: If crop_size exceeds image dimensions.
"""
if crop_size > image_width or crop_size > image_height:
raise ValueError(
f"crop_size ({crop_size}) exceeds image dimensions "
f"({image_width}x{image_height})."
)

x = (image_width - crop_size) // 2
y = (image_height - crop_size) // 2
return x, y


def generate_center_crops(
dataset: BaseImageDataset,
crop_size: int,
) -> CropMap:
"""
Generate center crop coordinates for each sample in a BaseImageDataset.

:param dataset: A BaseImageDataset instance (or compatible object with
`file_state.manifest` attribute supporting `get_image_dimensions()`).
:param crop_size: Size of the square crop (same width and height).
:return: Dictionary mapping manifest indices to lists of crop specs.
Format: {manifest_idx: [((x, y), width, height), ...]}
:raises ValueError: If crop_size is non-positive, if no active channels
are configured, or if channel dimensions don't match for any sample.
"""
if crop_size <= 0:
raise ValueError(f"crop_size must be positive, got {crop_size}.")

active_channels = _get_active_channels(dataset)
if not active_channels:
raise ValueError(
"No active channels configured. Set input_channel_keys and/or "
"target_channel_keys on the dataset before generating crops."
)

manifest = dataset.file_state.manifest
crop_specs: Dict[int, List[Tuple[Tuple[int, int], int, int]]] = {}

for idx in range(len(dataset)):
# Get dimensions for all active channels
dims = manifest.get_image_dimensions(idx, channels=active_channels)

# Validate all channels have matching dimensions
width, height = _validate_same_dimensions_across_channels(
dims, active_channels, idx
)

# Compute center crop coordinates
x, y = _compute_center_crop(width, height, crop_size)

# Store as crop_specs format: {idx: [((x, y), w, h), ...]}
crop_specs[idx] = [((x, y), crop_size, crop_size)]

return crop_specs
from .crop_generators.point_centered import generate_point_centered_crops
from .crop_generators.tile import generate_tile_crops
__all__ = [
"CropSpec",
"CropMap",
"CropGenerator",
"generate_center_crops",
"generate_point_centered_crops",
"generate_tile_crops",
"_compute_center_crop",
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
from .protocol import CropMap, CropSpec, CropGenerator
from .center import generate_center_crops, _compute_center_crop
from .point_centered import generate_point_centered_crops
from .tile import generate_tile_crops

__all__ = [
"CropMap",
"CropSpec",
"CropGenerator",
"generate_center_crops",
"generate_point_centered_crops",
"generate_tile_crops",
"_compute_center_crop"
]
Loading
Loading