diff --git a/pyproject.toml b/pyproject.toml index 6d0fb1e17..39ba15170 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -16,7 +16,8 @@ classifiers = [ "Programming Language :: Python :: 3.13", ] dependencies = [ - "iohub>=0.2.0b0", + "iohub[tensorstore]>=0.2.2rc0", + "kornia", "torch>=2.4.1", "timm>=0.9.5", "tensorboard>=2.13.0", @@ -61,7 +62,6 @@ dev = [ "pytest-cov", "hypothesis", "ruff", - "profilehooks", "onnxruntime", ] diff --git a/viscy/cli.py b/viscy/cli.py index 1a82ad505..0c07787ad 100644 --- a/viscy/cli.py +++ b/viscy/cli.py @@ -64,3 +64,7 @@ def main() -> None: "description": "Computer vision models for single-cell phenotyping." }, ) + + +if __name__ == "__main__": + main() diff --git a/viscy/data/combined.py b/viscy/data/combined.py index 870367230..d4f2c8330 100644 --- a/viscy/data/combined.py +++ b/viscy/data/combined.py @@ -1,14 +1,20 @@ +import bisect +import logging +from collections import defaultdict from enum import Enum from typing import Literal, Sequence import torch from lightning.pytorch import LightningDataModule from lightning.pytorch.utilities.combined_loader import CombinedLoader +from monai.data import ThreadDataLoader from torch.utils.data import ConcatDataset, DataLoader, Dataset from viscy.data.distributed import ShardedDistributedSampler from viscy.data.hcs import _collate_samples +_logger = logging.getLogger("lightning.pytorch") + class CombineMode(Enum): MIN_SIZE = "min_size" @@ -82,6 +88,37 @@ def predict_dataloader(self): ) +class BatchedConcatDataset(ConcatDataset): + def __getitem__(self, idx): + raise NotImplementedError + + def _get_sample_indices(self, idx: int) -> tuple[int, int]: + if idx < 0: + if -idx > len(self): + raise ValueError( + "absolute value of index should not exceed dataset length" + ) + idx = len(self) + idx + dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx) + if dataset_idx == 0: + sample_idx = idx + else: + sample_idx = idx - self.cumulative_sizes[dataset_idx - 1] + return dataset_idx, sample_idx + + def __getitems__(self, indices: list[int]) -> list: + grouped_indices = defaultdict(list) + for idx in indices: + dataset_idx, sample_indices = self._get_sample_indices(idx) + grouped_indices[dataset_idx].append(sample_indices) + _logger.debug(f"Grouped indices: {grouped_indices}") + sub_batches = [] + for dataset_idx, sample_indices in grouped_indices.items(): + sub_batch = self.datasets[dataset_idx].__getitems__(sample_indices) + sub_batches.extend(sub_batch) + return sub_batches + + class ConcatDataModule(LightningDataModule): """ Concatenate multiple data modules. @@ -96,11 +133,16 @@ class ConcatDataModule(LightningDataModule): Data modules to concatenate. """ + _ConcatDataset = ConcatDataset + def __init__(self, data_modules: Sequence[LightningDataModule]): super().__init__() self.data_modules = data_modules self.num_workers = data_modules[0].num_workers self.batch_size = data_modules[0].batch_size + self.persistent_workers = data_modules[0].persistent_workers + self.prefetch_factor = data_modules[0].prefetch_factor + self.pin_memory = data_modules[0].pin_memory for dm in data_modules: if dm.num_workers != self.num_workers: raise ValueError("Inconsistent number of workers") @@ -124,28 +166,62 @@ def setup(self, stage: Literal["fit", "validate", "test", "predict"]): raise ValueError("Inconsistent patches per stack") if stage != "fit": raise NotImplementedError("Only fit stage is supported") - self.train_dataset = ConcatDataset( + self.train_dataset = self._ConcatDataset( [dm.train_dataset for dm in self.data_modules] ) - self.val_dataset = ConcatDataset([dm.val_dataset for dm in self.data_modules]) + self.val_dataset = self._ConcatDataset( + [dm.val_dataset for dm in self.data_modules] + ) + + def _dataloader_kwargs(self) -> dict: + return { + "num_workers": self.num_workers, + "persistent_workers": self.persistent_workers, + "prefetch_factor": self.prefetch_factor if self.num_workers else None, + "pin_memory": self.pin_memory, + } def train_dataloader(self): return DataLoader( self.train_dataset, - batch_size=self.batch_size // self.train_patches_per_stack, - num_workers=self.num_workers, shuffle=True, - persistent_workers=bool(self.num_workers), + batch_size=self.batch_size // self.train_patches_per_stack, collate_fn=_collate_samples, + drop_last=True, + **self._dataloader_kwargs(), ) def val_dataloader(self): return DataLoader( self.val_dataset, + shuffle=False, + batch_size=self.batch_size, + drop_last=False, + **self._dataloader_kwargs(), + ) + + +class BatchedConcatDataModule(ConcatDataModule): + _ConcatDataset = BatchedConcatDataset + + def train_dataloader(self): + return ThreadDataLoader( + self.train_dataset, + use_thread_workers=True, + batch_size=self.batch_size, + shuffle=True, + drop_last=True, + **self._dataloader_kwargs(), + ) + + def val_dataloader(self): + return ThreadDataLoader( + self.val_dataset, + use_thread_workers=True, batch_size=self.batch_size, - num_workers=self.num_workers, shuffle=False, - persistent_workers=bool(self.num_workers), + drop_last=False, + **self._dataloader_kwargs(), ) diff --git a/viscy/data/triplet.py b/viscy/data/triplet.py index c25a0fc74..59bd1410a 100644 --- a/viscy/data/triplet.py +++ b/viscy/data/triplet.py @@ -2,10 +2,14 @@ from pathlib import Path from typing import Literal, Sequence +import numpy as np import pandas as pd +import tensorstore as ts import torch from iohub.ngff import ImageArray, Position, open_ome_zarr -from monai.transforms import Compose, MapTransform +from monai.data import ThreadDataLoader +from monai.data.utils import collate_meta_tensor +from monai.transforms import Compose, MapTransform, ToDeviced from torch import Tensor from torch.utils.data import Dataset @@ -31,19 +35,23 @@ def _scatter_channels( channel_names: list[str], patch: Tensor, norm_meta: NormMeta | None ) -> dict[str, Tensor | NormMeta] | dict[str, Tensor]: - channels = {name: data[None] for name, data in zip(channel_names, patch)} + channels = { + name: patch[:, c : c + 1] + for name, c in zip(channel_names, range(patch.shape[1])) + } if norm_meta is not None: - channels |= {"norm_meta": norm_meta} + channels["norm_meta"] = collate_meta_tensor(norm_meta) return channels -def _gather_channels(patch_channels: dict[str, Tensor | NormMeta]) -> Tensor: - """ - :param dict[str, Tensor | NormMeta] patch_channels: dictionary of single-channel tensors - :return Tensor: Multi-channel tensor - """ - patch_channels.pop("norm_meta", None) - return torch.cat(list(patch_channels.values()), dim=0) +def _gather_channels( + patch_channels: list[dict[str, Tensor | NormMeta]], +) -> list[Tensor]: + samples = [] + for sample in patch_channels: + sample.pop("norm_meta", None) + samples.append(torch.cat(list(sample.values()), dim=0)) + return samples def _transform_channel_wise( @@ -51,10 +59,10 @@ def _transform_channel_wise( channel_names: list[str], patch: Tensor, norm_meta: NormMeta | None, -) -> Tensor: - return _gather_channels( - transform(_scatter_channels(channel_names, patch, norm_meta)) - ) +) -> list[Tensor]: + scattered_channels = _scatter_channels(channel_names, patch, norm_meta) + transformed_channels = transform(scattered_channels) + return _gather_channels(transformed_channels) class TripletDataset(Dataset): @@ -198,14 +206,11 @@ def _specific_cells(self, tracks: pd.DataFrame) -> pd.DataFrame: def __len__(self) -> int: return len(self.valid_anchors) - def _sample_positive(self, anchor_row: pd.Series) -> pd.Series: + def _sample_positives(self, anchor_rows: pd.DataFrame) -> pd.DataFrame: """Select a positive sample from the same track in the next time point.""" - same_track = self.tracks[ - (self.tracks["global_track_id"] == anchor_row["global_track_id"]) - ] - return same_track[ - same_track["t"] == (anchor_row["t"] + self.time_interval) - ].iloc[0] + query = anchor_rows[["global_track_id", "t"]].copy() + query["t"] += self.time_interval + return query.merge(self.tracks, on=["global_track_id", "t"], how="inner") def _sample_negative(self, anchor_row: pd.Series) -> pd.Series: """Select a negative sample from a different track in the next time point @@ -226,9 +231,17 @@ def _sample_negative(self, anchor_row: pd.Series) -> pd.Series: # reproducibility relies on setting a global seed for numpy return candidates.sample(n=1).iloc[0] - def _slice_patch(self, track_row: pd.Series) -> tuple[Tensor, NormMeta | None]: + def _sample_negatives(self, anchor_rows: pd.DataFrame) -> pd.DataFrame: + return pd.concat( + [self._sample_negative(row) for _, row in anchor_rows.iterrows()], + axis=1, + ) + + def _slice_patch( + self, track_row: pd.Series + ) -> tuple[ts.TensorStore, NormMeta | None]: position: Position = track_row["position"] - image = position["0"] + image = position["0"].tensorstore() time = track_row["t"] y_center = track_row["y"] x_center = track_row["x"] @@ -240,60 +253,74 @@ def _slice_patch(self, track_row: pd.Series) -> tuple[Tensor, NormMeta | None]: slice(y_center - y_half, y_center + y_half), slice(x_center - x_half, x_center + x_half), ] - return torch.from_numpy(patch), _read_norm_meta(position) - - def __getitem__(self, index: int) -> TripletSample: - anchor_row = self.valid_anchors.iloc[index] - anchor_patch, anchor_norm = self._slice_patch(anchor_row) + return patch, _read_norm_meta(position) + + def _slice_patches(self, track_rows: pd.DataFrame): + patches = [] + norms = [] + with ts.Batch() as batch: + for _, row in track_rows.iterrows(): + patch, norm = self._slice_patch(row) + patches.append(patch.read(batch=batch)) + norms.append(norm) + results = [p.result() for p in patches] + return torch.from_numpy(np.stack(results, axis=0)), norms + + def __getitems__(self, indices: list[int]) -> list[TripletSample]: + anchor_rows = self.valid_anchors.iloc[indices] + anchor_patches, anchor_norms = self._slice_patches(anchor_rows) if self.fit: if self.time_interval == "any": - positive_patch = anchor_patch.clone() - positive_norm = anchor_norm + positive_patches = anchor_patches.clone() + positive_norms = anchor_norms else: - positive_row = self._sample_positive(anchor_row) - positive_patch, positive_norm = self._slice_patch(positive_row) + positive_rows = self._sample_positives(anchor_rows) + positive_patches, positive_norms = self._slice_patches(positive_rows) if self.positive_transform: - positive_patch = _transform_channel_wise( + positive_patches = _transform_channel_wise( transform=self.positive_transform, channel_names=self.channel_names, - patch=positive_patch, - norm_meta=positive_norm, + patch=positive_patches, + norm_meta=positive_norms, ) if self.return_negative: - negative_row = self._sample_negative(anchor_row) - negative_patch, negative_norm = self._slice_patch(negative_row) + negative_rows = self._sample_negatives(anchor_rows) + negative_patches, negative_norms = self._slice_patches(negative_rows) if self.negative_transform: - negative_patch = _transform_channel_wise( + negative_patches = _transform_channel_wise( transform=self.negative_transform, channel_names=self.channel_names, - patch=negative_patch, - norm_meta=negative_norm, + patch=negative_patches, + norm_meta=negative_norms, ) if self.anchor_transform: - anchor_patch = _transform_channel_wise( + anchor_patches = _transform_channel_wise( transform=self.anchor_transform, channel_names=self.channel_names, - patch=anchor_patch, - norm_meta=anchor_norm, + patch=anchor_patches, + norm_meta=anchor_norms, ) - sample = {"anchor": anchor_patch} + samples: list[TripletSample] = [ + {"anchor": anchor_patch} for anchor_patch in anchor_patches + ] if self.fit: + for sample, positive_patch in zip(samples, positive_patches): + sample["positive"] = positive_patch if self.return_negative: - sample.update({"positive": positive_patch, "negative": negative_patch}) - else: - sample.update({"positive": positive_patch}) + for sample, negative_patch in zip(samples, negative_patches): + sample["negative"] = negative_patch else: - # For new predictions, ensure all INDEX_COLUMNS are included - index_dict = {} - for col in INDEX_COLUMNS: - if col in anchor_row.index: - index_dict[col] = anchor_row[col] - else: - # Skip y and x for legacy data - they weren't part of INDEX_COLUMNS - if col not in ["y", "x", "z"]: + for sample, (_, anchor_row) in zip(samples, anchor_rows.iterrows()): + # For new predictions, ensure all INDEX_COLUMNS are included + index_dict = {} + for col in INDEX_COLUMNS: + if col in anchor_row.index: + index_dict[col] = anchor_row[col] + elif col not in ["y", "x", "z"]: + # Skip y and x for legacy data - they weren't part of INDEX_COLUMNS raise KeyError(f"Required column '{col}' not found in data") - sample.update({"index": index_dict}) - return sample + sample["index"] = index_dict + return samples class TripletDataModule(HCSDataModule): @@ -435,7 +462,16 @@ def _base_dataset_settings(self) -> dict: "time_interval": self.time_interval, } + def _update_to_device_transform(self): + "Make sure that GPU transforms are set to the current device." + for transform in self.normalizations + self.augmentations: + if isinstance(transform, ToDeviced): + transform.converter.device = torch.device( + f"cuda:{torch.cuda.current_device()}" + ) + def _setup_fit(self, dataset_settings: dict): + self._update_to_device_transform() augment_transform, no_aug_transform = self._fit_transform() positions, tracks_tables = self._align_tracks_tables_with_positions() shuffled_indices = self._set_fit_global_state(len(positions)) @@ -495,3 +531,42 @@ def _setup_predict(self, dataset_settings: dict): def _setup_test(self, *args, **kwargs): raise NotImplementedError("Self-supervised model does not support testing") + + def train_dataloader(self): + return ThreadDataLoader( + self.train_dataset, + use_thread_workers=True, + batch_size=self.batch_size, + num_workers=self.num_workers, + shuffle=True, + prefetch_factor=self.prefetch_factor if self.num_workers else None, + persistent_workers=self.persistent_workers, + drop_last=True, + pin_memory=self.pin_memory, + ) + + def val_dataloader(self): + return ThreadDataLoader( + self.val_dataset, + use_thread_workers=True, + batch_size=self.batch_size, + num_workers=self.num_workers, + shuffle=False, + prefetch_factor=self.prefetch_factor if self.num_workers else None, + persistent_workers=self.persistent_workers, + drop_last=False, + pin_memory=self.pin_memory, + ) + + def predict_dataloader(self): + return ThreadDataLoader( + self.predict_dataset, + use_thread_workers=True, + batch_size=self.batch_size, + num_workers=self.num_workers, + shuffle=False, + prefetch_factor=self.prefetch_factor if self.num_workers else None, + persistent_workers=self.persistent_workers, + drop_last=False, + pin_memory=self.pin_memory, + ) diff --git a/viscy/data/typing.py b/viscy/data/typing.py index c824b9416..d6a70488c 100644 --- a/viscy/data/typing.py +++ b/viscy/data/typing.py @@ -84,7 +84,7 @@ class TripletSample(TypedDict): Triplet sample type for mini-batches. """ - index: TrackingIndex anchor: Tensor positive: NotRequired[Tensor] negative: NotRequired[Tensor] + index: NotRequired[TrackingIndex] diff --git a/viscy/scripts/profiling.py b/viscy/scripts/profiling.py index a0c3ca6d8..5b978a09f 100644 --- a/viscy/scripts/profiling.py +++ b/viscy/scripts/profiling.py @@ -1,34 +1,169 @@ # script to profile dataloading +# use with a sampling profiler like py-spy +from monai.transforms import ( + Decollated, + RandAdjustContrastd, + RandGaussianSmoothd, + RandScaleIntensityd, + ToDeviced, +) +from pytorch_metric_learning.losses import NTXentLoss + +from viscy.data.combined import BatchedConcatDataModule +from viscy.data.triplet import TripletDataModule +from viscy.representation.engine import ContrastiveEncoder, ContrastiveModule +from viscy.transforms import ( + NormalizeSampled, +) +from viscy.transforms._transforms import ( + BatchedRandAffined, + BatchedScaleIntensityRangePercentilesd, + RandGaussianNoiseTensord, +) -from profilehooks import profile -from viscy.data.hcs import HCSDataModule +def model( + input_channel_number: int = 1, + z_stack_depth: int = 30, + patch_size: int = 192, + temperature: float = 0.5, +): + return ContrastiveModule( + encoder=ContrastiveEncoder( + backbone="convnext_tiny", + in_channels=input_channel_number, + in_stack_depth=z_stack_depth, + stem_kernel_size=(5, 4, 4), + embedding_dim=768, + projection_dim=32, + drop_path_rate=0.0, + ), + loss_function=NTXentLoss(temperature=temperature), + lr=0.00002, + log_batches_per_epoch=3, + log_samples_per_batch=3, + example_input_array_shape=[ + 1, + input_channel_number, + z_stack_depth, + patch_size, + patch_size, + ], + ) -dataset = "/path/to/dataset.zarr" +def channel_augmentations(processing_channel: str): + return [ + BatchedRandAffined( + keys=[processing_channel], + prob=0.8, + scale_range=((1.0, 1.0), (0.8, 1.2), (0.8, 1.2)), + rotate_range=[1.0, 0.0, 0.0], + shear_range=(0.2, 0.2, 0.0, 0.2, 0.0, 0.2), + ), + Decollated(keys=[processing_channel]), + RandAdjustContrastd( + keys=[processing_channel], + prob=0.5, + gamma=[0.8, 1.2], + ), + RandScaleIntensityd( + keys=[processing_channel], + prob=0.5, + factors=0.5, + ), + RandGaussianSmoothd( + keys=[processing_channel], + prob=0.5, + sigma_x=[0.25, 0.75], + sigma_y=[0.25, 0.75], + sigma_z=[0.0, 0.0], + ), + RandGaussianNoiseTensord( + keys=[processing_channel], + prob=0.5, + mean=0.0, + std=0.2, + ), + ] -dm = HCSDataModule( - dataset, - "Phase3D", - "Deconvolved-Nuc", - 5, - 0.8, - batch_size=32, - num_workers=32, - augment=None, - caching=False, -) -dm.setup("fit") +def channel_normalization( + phase_channel: str = None, + fl_channel: str = None, +): + if phase_channel: + return [ + NormalizeSampled( + keys=[phase_channel], + level="fov_statistics", + subtrahend="mean", + divisor="std", + ) + ] + elif fl_channel: + return [ + ToDeviced(keys=[fl_channel], device="cuda"), + BatchedScaleIntensityRangePercentilesd( + keys=[fl_channel], + lower=50, + upper=99, + b_min=0.0, + b_max=1.0, + ), + ] + else: + raise NotImplementedError("Either phase_channel or fl_channel must be provided") + +if __name__ == "__main__": + dm1 = TripletDataModule( + data_path="/hpc/projects/organelle_phenotyping/datasets/organelle/SEC61B/2024_10_16_A549_SEC61_ZIKV_DENV/2024_10_16_A549_SEC61_ZIKV_DENV_2.zarr", + tracks_path="/hpc/projects/intracellular_dashboard/organelle_dynamics/rerun/2024_10_16_A549_SEC61_ZIKV_DENV/1-preprocess/label-free/3-track/2024_10_16_A549_SEC61_ZIKV_DENV_cropped.zarr", + source_channel=["raw GFP EX488 EM525-45"], + z_range=[5, 35], + initial_yx_patch_size=(384, 384), + final_yx_patch_size=(192, 192), + batch_size=16, + num_workers=4, + time_interval=1, + augmentations=channel_augmentations("raw GFP EX488 EM525-45"), + normalizations=channel_normalization( + phase_channel=None, fl_channel="raw GFP EX488 EM525-45" + ), + fit_include_wells=["B/3", "B/4", "C/3", "C/4"], + return_negative=False, + ) + dm2 = TripletDataModule( + data_path="/hpc/projects/organelle_phenotyping/datasets/organelle/SEC61B/2024_10_16_A549_SEC61_ZIKV_DENV/2024_10_16_A549_SEC61_ZIKV_DENV_2.zarr", + tracks_path="/hpc/projects/intracellular_dashboard/organelle_dynamics/rerun/2024_10_16_A549_SEC61_ZIKV_DENV/1-preprocess/label-free/3-track/2024_10_16_A549_SEC61_ZIKV_DENV_cropped.zarr", + source_channel=["raw mCherry EX561 EM600-37"], + z_range=[5, 35], + initial_yx_patch_size=(384, 384), + final_yx_patch_size=(192, 192), + batch_size=16, + num_workers=4, + time_interval=1, + augmentations=channel_augmentations("raw mCherry EX561 EM600-37"), + normalizations=channel_normalization( + phase_channel=None, fl_channel="raw mCherry EX561 EM600-37" + ), + fit_include_wells=["B/3", "B/4", "C/3", "C/4"], + return_negative=False, + ) + dm = BatchedConcatDataModule(data_modules=[dm1, dm2]) + dm.setup("fit") -@profile(immediate=True, sort="time", dirs=True) -def load_batch(n=1): + print(len(dm1.train_dataset), len(dm2.train_dataset), len(dm.train_dataset)) + n = 1 + + print("Training batches:") for i, batch in enumerate(dm.train_dataloader()): - print(batch["source"].shape) - print(dm.on_before_batch_transfer(batch, 0)["target"].shape) + print(i, batch["anchor"].shape, batch["positive"].device) + if i == n - 1: + break + print("Validation batches:") + for i, batch in enumerate(dm.val_dataloader()): + print(i, batch["anchor"].shape, batch["positive"].device) if i == n - 1: break - - -load_batch(3) diff --git a/viscy/transforms/__init__.py b/viscy/transforms/__init__.py index 12177b64b..40712f9e1 100644 --- a/viscy/transforms/__init__.py +++ b/viscy/transforms/__init__.py @@ -1,5 +1,6 @@ from viscy.transforms._redef import ( CenterSpatialCropd, + Decollated, RandAdjustContrastd, RandAffined, RandFlipd, @@ -9,23 +10,31 @@ RandSpatialCropd, RandWeightedCropd, ScaleIntensityRangePercentilesd, + ToDeviced, ) from viscy.transforms._transforms import ( + BatchedRandAffined, + BatchedScaleIntensityRangePercentilesd, BatchedZoom, NormalizeSampled, + RandGaussianNoiseTensord, RandInvertIntensityd, StackChannelsd, TiledSpatialCropSamplesd, ) __all__ = [ + "BatchedRandAffined", + "BatchedScaleIntensityRangePercentilesd", "BatchedZoom", "CenterSpatialCropd", + "Decollated", "NormalizeSampled", "RandAdjustContrastd", "RandAffined", "RandFlipd", "RandGaussianNoised", + "RandGaussianNoiseTensord", "RandGaussianSmoothd", "RandInvertIntensityd", "RandScaleIntensityd", @@ -34,4 +43,5 @@ "ScaleIntensityRangePercentilesd", "StackChannelsd", "TiledSpatialCropSamplesd", + "ToDeviced", ] diff --git a/viscy/transforms/_redef.py b/viscy/transforms/_redef.py index fe168603a..696c81abc 100644 --- a/viscy/transforms/_redef.py +++ b/viscy/transforms/_redef.py @@ -4,6 +4,7 @@ from monai.transforms import ( CenterSpatialCropd, + Decollated, RandAdjustContrastd, RandAffined, RandFlipd, @@ -13,10 +14,34 @@ RandSpatialCropd, RandWeightedCropd, ScaleIntensityRangePercentilesd, + ToDeviced, ) from numpy.typing import DTypeLike +class Decollated(Decollated): + def __init__( + self, + keys: Sequence[str] | str, + detach: bool = True, + pad_batch: bool = True, + fill_value: float | None = None, + **kwargs, + ): + super().__init__( + keys=keys, + detach=detach, + pad_batch=pad_batch, + fill_value=fill_value, + **kwargs, + ) + + +class ToDeviced(ToDeviced): + def __init__(self, keys: Sequence[str] | str, **kwargs): + super().__init__(keys=keys, **kwargs) + + class RandWeightedCropd(RandWeightedCropd): def __init__( self, diff --git a/viscy/transforms/_transforms.py b/viscy/transforms/_transforms.py index 0b99a2ac4..c9418f1b6 100644 --- a/viscy/transforms/_transforms.py +++ b/viscy/transforms/_transforms.py @@ -1,13 +1,20 @@ +from warnings import warn + import numpy as np import torch +from kornia.augmentation import RandomAffine3D from monai.transforms import ( MapTransform, MultiSampleTrait, + RandGaussianNoise, + RandGaussianNoised, RandomizableTransform, + ScaleIntensityRangePercentiles, Transform, ) +from numpy.typing import DTypeLike from torch import Tensor -from typing_extensions import Iterable, Literal +from typing_extensions import Iterable, Literal, Sequence from viscy.data.typing import ChannelMap, Sample @@ -44,12 +51,20 @@ def __init__( self.level = level self.remove_meta = remove_meta + @staticmethod + def _match_image(tensor: Tensor, target: Tensor) -> Tensor: + return tensor.reshape(tensor.shape + (1,) * (target.ndim - tensor.ndim)).to( + device=target.device + ) + # TODO: need to implement the case where the preprocessing already exists def __call__(self, sample: Sample) -> Sample: for key in self.keys: level_meta = sample["norm_meta"][key][self.level] subtrahend_val = level_meta[self.subtrahend] + subtrahend_val = self._match_image(subtrahend_val, sample[key]) divisor_val = level_meta[self.divisor] + 1e-8 # avoid div by zero + divisor_val = self._match_image(divisor_val, sample[key]) sample[key] = (sample[key] - subtrahend_val) / divisor_val if self.remove_meta: sample.pop("norm_meta") @@ -184,3 +199,169 @@ def __call__(self, sample: Tensor) -> Tensor: recompute_scale_factor=self.recompute_scale_factor, antialias=self.antialias, ) + + +class BatchedScaleIntensityRangePercentiles(ScaleIntensityRangePercentiles): + def _normalize(self, img: Tensor) -> Tensor: + q_low = self.lower / 100.0 + q_high = self.upper / 100.0 + batch_size, *_ = img.shape + # TODO: address pytorch#64947 to improve performance + a_min, a_max = torch.quantile( + img.view(batch_size, -1), + torch.tensor([q_low, q_high], dtype=img.dtype, device=img.device), + dim=1, + ).reshape(2, batch_size, 1, 1, 1, 1) + b_min = self.b_min + b_max = self.b_max + + if self.relative: + if (self.b_min is None) or (self.b_max is None): + raise ValueError( + "If it is relative, b_min and b_max should not be None." + ) + b_min = ((self.b_max - self.b_min) * (q_low)) + self.b_min + b_max = ((self.b_max - self.b_min) * (q_high)) + self.b_min + + if (a_min == a_max).any(): + warn("Divide by zero (a_min == a_max)") + if b_min is None: + return img - a_min + return img - a_min + b_min + + img = (img - a_min) / (a_max - a_min) + if (b_min is not None) and (b_max is not None): + img = img * (b_max - b_min) + b_min + if self.clip: + img = img.clip(b_min, b_max) + + return img + + def __call__(self, img: Tensor) -> Tensor: + if self.channel_wise: + channels = [self._normalize(img[:, c : c + 1]) for c in range(img.shape[1])] + return torch.cat(channels, dim=1) + else: + return self._normalize(img=img) + + +class BatchedScaleIntensityRangePercentilesd(MapTransform): + def __init__( + self, + keys: str | Iterable[str], + lower: float, + upper: float, + b_min: float | None, + b_max: float | None, + clip: bool = False, + relative: bool = False, + channel_wise: bool = False, + dtype: DTypeLike = np.float32, + allow_missing_keys: bool = False, + ) -> None: + super().__init__(keys, allow_missing_keys) + self.scaler = BatchedScaleIntensityRangePercentiles( + lower, upper, b_min, b_max, clip, relative, channel_wise, dtype + ) + + def __call__(self, data: dict[str, Tensor]) -> dict[str, Tensor]: + d = dict(data) + for key in self.key_iterator(d): + d[key] = self.scaler(d[key]) + return d + + +class BatchedRandAffined(MapTransform): + def __init__( + self, + keys: str | Iterable[str], + prob: float = 0.1, + rotate_range: Sequence[tuple[float, float] | float] | float | None = None, + shear_range: Sequence[tuple[float, float] | float] | float | None = None, + translate_range: Sequence[tuple[float, float] | float] | float | None = None, + scale_range: Sequence[tuple[float, float] | float] | float | None = None, + mode: str = "bilinear", + allow_missing_keys: bool = False, + ) -> None: + super().__init__(keys, allow_missing_keys) + rotate_range = self._radians_to_degrees( + self._maybe_invert_sequence(rotate_range) + ) + if rotate_range is None: + rotate_range = (0.0, 0.0, 0.0) + shear_range = self._radians_to_degrees(self._maybe_invert_sequence(shear_range)) + translate_range = self._maybe_invert_sequence(translate_range) + scale_range = self._maybe_invert_sequence(scale_range) + self.random_affine = RandomAffine3D( + degrees=rotate_range, + translate=translate_range, + scale=scale_range, + shears=shear_range, + resample=mode, + p=prob, + ) + # disable unnecessary transfer to CPU + self.random_affine.disable_features = True + + @staticmethod + def _maybe_invert_sequence( + value: Sequence[tuple[float, float] | float] | float | None, + ) -> Sequence[tuple[float, float] | float] | float | None: + """Translate MONAI's ZYX order to Kornia's XYZ order.""" + if isinstance(value, Sequence): + return tuple(reversed(value)) + return value + + @staticmethod + def _radians_to_degrees( + rotate_range: Sequence[tuple[float, float] | float] | float | None, + ) -> Sequence[tuple[float, float] | float] | float | None: + if rotate_range is None: + return None + return torch.from_numpy(np.rad2deg(rotate_range)) + + @torch.no_grad() + def __call__(self, sample: dict[str, Tensor]) -> dict[str, Tensor]: + d = dict(sample) + for key in self.key_iterator(d): + data = d[key] + try: + d[key] = self.random_affine(data) + except RuntimeError: + # retry + d[key] = self.random_affine(data) + assert d[key].device == data.device + return d + + +class RandGaussianNoiseTensor(RandGaussianNoise): + def randomize(self, img: Tensor, mean: float | None = None) -> None: + self._do_transform = self.R.rand() < self.prob + if not self._do_transform: + return None + std = self.R.uniform(0, self.std) if self.sample_std else self.std + self.noise = torch.normal( + self.mean if mean is None else mean, + std, + size=img.shape, + device=img.device, + dtype=img.dtype, + ) + + +class RandGaussianNoiseTensord(RandGaussianNoised): + def __init__( + self, + keys: str | Iterable[str], + prob: float = 0.1, + mean: float = 0.0, + std: float = 0.1, + dtype: DTypeLike = np.float32, + allow_missing_keys: bool = False, + sample_std: bool = True, + ) -> None: + MapTransform.__init__(self, keys, allow_missing_keys) + RandomizableTransform.__init__(self, prob) + self.rand_gaussian_noise = RandGaussianNoiseTensor( + mean=mean, std=std, prob=1.0, dtype=dtype, sample_std=sample_std + )