From 8f42e30cbfc0eb140561b099213111240c6d34d1 Mon Sep 17 00:00:00 2001 From: Ivan Ivanov Date: Wed, 3 Jun 2026 16:22:59 -0700 Subject: [PATCH 1/4] feat(cytoland): non-square rotation TTA + reuse helpers - Fix AugmentedPredictionVSUNet._predict_with_tta to crop to the augmented (post-transform) spatial shape instead of the original. Shape-changing transforms such as 90/270-degree rotations swap Y and X, so cropping to the original shape was wrong for non-square FOVs (mismatched shapes failed to reduce). Rotation TTA now works for non-square inputs with no square padding. - Add rotation_tta_transforms() factory for the standard 90-degree rotation forward/inverse transform set, and AugmentedPredictionVSUNet.with_rotation_tta() convenience constructor that uses it, so callers don't hand-build transforms. - Rename viscy_data._read_norm_meta to public read_norm_meta (updating all callers) and export it, so external pipelines can read precomputed normalization statistics from a FOV's zattrs. - Add regression tests for non-square rotation TTA and the helper. Co-Authored-By: Claude Opus 4.8 (1M context) --- .../cytoland/src/cytoland/__init__.py | 2 + applications/cytoland/src/cytoland/engine.py | 71 ++++++++++++++++++- applications/cytoland/tests/test_engine.py | 43 +++++++++++ .../dynaclr/src/dynaclr/data/dataset.py | 4 +- .../viscy-data/src/viscy_data/__init__.py | 4 ++ packages/viscy-data/src/viscy_data/_utils.py | 6 +- .../src/viscy_data/cell_classification.py | 4 +- packages/viscy-data/src/viscy_data/gpu_aug.py | 4 +- .../viscy-data/src/viscy_data/mmap_cache.py | 4 +- .../src/viscy_data/sliding_window.py | 4 +- packages/viscy-data/src/viscy_data/triplet.py | 4 +- 11 files changed, 134 insertions(+), 16 deletions(-) diff --git a/applications/cytoland/src/cytoland/__init__.py b/applications/cytoland/src/cytoland/__init__.py index 6df466c9e..d399e9802 100644 --- a/applications/cytoland/src/cytoland/__init__.py +++ b/applications/cytoland/src/cytoland/__init__.py @@ -5,6 +5,7 @@ FcmaeUNet, MaskedMSELoss, VSUNet, + rotation_tta_transforms, ) from cytoland.evaluation import SegmentationMetrics2D @@ -14,4 +15,5 @@ "MaskedMSELoss", "SegmentationMetrics2D", "VSUNet", + "rotation_tta_transforms", ] diff --git a/applications/cytoland/src/cytoland/engine.py b/applications/cytoland/src/cytoland/engine.py index 03b272c08..9a1be7b75 100644 --- a/applications/cytoland/src/cytoland/engine.py +++ b/applications/cytoland/src/cytoland/engine.py @@ -3,6 +3,7 @@ import inspect import logging import os +from functools import partial from typing import Callable, Literal, Sequence import numpy as np @@ -70,6 +71,34 @@ def _center_crop_to_shape(tensor: Tensor, spatial_shape: tuple[int, ...]) -> Ten return tensor[tuple(slices)] +def rotation_tta_transforms( + n: int = 4, +) -> tuple[list[Callable[[Tensor], Tensor]], list[Callable[[Tensor], Tensor]]]: + """Build forward/inverse 90-degree rotation transforms for test-time augmentation. + + Returns ``(forward_transforms, inverse_transforms)`` suitable for + :class:`AugmentedPredictionVSUNet`. Each forward transform rotates the YX + plane by ``k * 90`` degrees (``k = 0 .. n-1``) and the matching inverse + rotates back. Combined with ``reduction="median"`` this reproduces the + rotation TTA used by :meth:`VSUNet.perform_test_time_augmentations`, and + (unlike passing rotations through other code paths) works for non-square + fields of view. + + Parameters + ---------- + n : int, optional + Number of 90-degree rotations, by default 4 (0, 90, 180, 270 degrees). + + Returns + ------- + tuple[list[Callable], list[Callable]] + The forward and inverse transform lists. + """ + forward = [partial(torch.rot90, k=k, dims=(-2, -1)) for k in range(n)] + inverse = [partial(torch.rot90, k=-k, dims=(-2, -1)) for k in range(n)] + return forward, inverse + + class MaskedMSELoss(nn.Module): """Masked MSE loss for FCMAE pre-training.""" @@ -603,6 +632,40 @@ def __init__( self._inverse_transforms = inverse_transforms or [_identity] self._reduction = reduction + @classmethod + def with_rotation_tta( + cls, + model: nn.Module, + n_rotations: int = 4, + reduction: Literal["mean", "median"] = "median", + ) -> "AugmentedPredictionVSUNet": + """Build a predictor that applies 90-degree rotation test-time augmentation. + + Convenience constructor that wires :func:`rotation_tta_transforms` into + the forward/inverse transform lists, so callers do not have to build + them by hand. Works for non-square fields of view. + + Parameters + ---------- + model : nn.Module + The model to wrap. + n_rotations : int, optional + Number of 90-degree rotations, by default 4. + reduction : {"mean", "median"}, optional + How to aggregate the rotated predictions, by default "median". + + Returns + ------- + AugmentedPredictionVSUNet + """ + forward_transforms, inverse_transforms = rotation_tta_transforms(n_rotations) + return cls( + model=model, + forward_transforms=forward_transforms, + inverse_transforms=inverse_transforms, + reduction=reduction, + ) + def forward(self, x: Tensor) -> Tensor: """Run forward pass through the model. @@ -659,9 +722,15 @@ def _predict_with_tta(self, source: Tensor) -> Tensor: preds = [] for fwd_t, inv_t in zip(self._forward_transforms, self._inverse_transforms): aug_source = fwd_t(source) + # Crop back to the augmented (post-forward-transform) spatial shape, + # not the original one: a shape-changing transform such as a 90/270 + # degree rotation swaps Y and X, so the prediction lives in the + # augmented frame until ``inv_t`` undoes the transform. Cropping to + # ``source.shape[2:]`` here would be wrong for non-square inputs. + aug_shape = aug_source.shape[2:] aug_source = self._predict_pad(aug_source) pred = self.forward(aug_source) - pred = _center_crop_to_shape(pred, source.shape[2:]) + pred = _center_crop_to_shape(pred, aug_shape) preds.append(inv_t(pred)) if len(preds) == 1: return preds[0] diff --git a/applications/cytoland/tests/test_engine.py b/applications/cytoland/tests/test_engine.py index de1ed5138..ac2a03759 100644 --- a/applications/cytoland/tests/test_engine.py +++ b/applications/cytoland/tests/test_engine.py @@ -206,3 +206,46 @@ def test_predict_sliding_windows_missing_out_stack_depth(): vs = AugmentedPredictionVSUNet(model=model) with pytest.raises(ValueError, match="out_stack_depth"): vs.predict_sliding_windows(torch.randn(1, 1, 10, 4, 4)) + + +def test_rotation_tta_transforms(): + """Verify the rotation TTA factory builds matched forward/inverse rotations.""" + from cytoland.engine import rotation_tta_transforms + + forward, inverse = rotation_tta_transforms() + assert len(forward) == len(inverse) == 4 + x = torch.randn(1, 1, 5, 6, 8) # non-square YX + for fwd_t, inv_t in zip(forward, inverse): + # inverse(forward(x)) is the identity and restores the original shape + assert torch.allclose(inv_t(fwd_t(x)), x) + + +@pytest.mark.parametrize("yx", [(64, 64), (64, 48), (48, 64)]) +def test_predict_sliding_windows_rotation_tta_nonsquare(yx): + """Verify rotation TTA + sliding windows works for non-square FOVs. + + Regression test: ``_predict_with_tta`` must crop to the augmented (rotated) + shape, otherwise 90/270-degree rotations on non-square inputs produce + mismatched shapes and fail to reduce. + """ + z_window, depth, out_channels = 5, 8, 2 + height, width = yx + model = VSUNet( + architecture="fcmae", + model_config={ + "in_channels": 1, + "out_channels": out_channels, + "encoder_blocks": [2, 2, 2, 2], + "dims": [4, 8, 16, 32], + "decoder_conv_blocks": 1, + "stem_kernel_size": [z_window, 4, 4], + "in_stack_depth": z_window, + "pretraining": False, + }, + ) + vs = AugmentedPredictionVSUNet.with_rotation_tta(model.model, reduction="median").eval() + x = torch.randn(1, 1, depth, height, width) + with torch.inference_mode(): + output = vs.predict_sliding_windows(x, out_channel=out_channels, step=1) + assert output.shape == (1, out_channels, depth, height, width) + assert torch.isfinite(output).all() diff --git a/applications/dynaclr/src/dynaclr/data/dataset.py b/applications/dynaclr/src/dynaclr/data/dataset.py index a682b744e..06acf3b35 100644 --- a/applications/dynaclr/src/dynaclr/data/dataset.py +++ b/applications/dynaclr/src/dynaclr/data/dataset.py @@ -36,7 +36,7 @@ from dynaclr.data.index import MultiExperimentIndex from dynaclr.data.tau_sampling import sample_tau from viscy_data._typing import ULTRACK_INDEX_COLUMNS, NormMeta, SampleMeta -from viscy_data._utils import _read_norm_meta +from viscy_data._utils import read_norm_meta def _pick_temporal_candidate( @@ -742,7 +742,7 @@ def _build_norm_meta( cache_key = (store_path, fov_name) if cache_key not in self._norm_meta_cache: position = self._get_position(store_path, fov_name) - self._norm_meta_cache[cache_key] = _read_norm_meta(position) + self._norm_meta_cache[cache_key] = read_norm_meta(position) cached = self._norm_meta_cache[cache_key] if cached is None: return None diff --git a/packages/viscy-data/src/viscy_data/__init__.py b/packages/viscy-data/src/viscy_data/__init__.py index edbec8ed6..4d2d54214 100644 --- a/packages/viscy-data/src/viscy_data/__init__.py +++ b/packages/viscy-data/src/viscy_data/__init__.py @@ -72,6 +72,9 @@ except ImportError: pass +# Normalization metadata reader (from _utils.py) +from viscy_data._utils import read_norm_meta + # Channel dropout augmentation (from channel_dropout.py) from viscy_data.channel_dropout import ChannelDropout @@ -150,6 +153,7 @@ "ChannelDropout", # Utilities "FlexibleBatchSampler", + "read_norm_meta", "SelectWell", "ShardedDistributedSampler", # Core diff --git a/packages/viscy-data/src/viscy_data/_utils.py b/packages/viscy-data/src/viscy_data/_utils.py index e6a6523c7..945741752 100644 --- a/packages/viscy-data/src/viscy_data/_utils.py +++ b/packages/viscy-data/src/viscy_data/_utils.py @@ -2,7 +2,7 @@ This module centralizes helper functions that are used by multiple data modules: - From ``hcs.py``: ``_ensure_channel_list``, ``_search_int_in_str``, - ``_collate_samples``, ``_read_norm_meta`` + ``_collate_samples``, ``read_norm_meta`` - From ``triplet.py``: ``_scatter_channels``, ``_gather_channels``, ``_transform_channel_wise`` """ @@ -24,7 +24,7 @@ "_collate_samples", "_ensure_channel_list", "_gather_channels", - "_read_norm_meta", + "read_norm_meta", "_scatter_channels", "_search_int_in_str", "_transform_channel_wise", @@ -136,7 +136,7 @@ def _collate_samples(batch: Sequence[Sample]) -> Sample: return collated -def _read_norm_meta(fov: Position) -> NormMeta | None: +def read_norm_meta(fov: Position) -> NormMeta | None: """Read normalization metadata from the FOV. Convert to float32 tensors to avoid automatic casting to float64. diff --git a/packages/viscy-data/src/viscy_data/cell_classification.py b/packages/viscy-data/src/viscy_data/cell_classification.py index b78a72cc1..c378747a8 100644 --- a/packages/viscy-data/src/viscy_data/cell_classification.py +++ b/packages/viscy-data/src/viscy_data/cell_classification.py @@ -21,7 +21,7 @@ from torch.utils.data import DataLoader, Dataset from viscy_data._typing import ULTRACK_INDEX_COLUMNS, AnnotationColumns -from viscy_data._utils import _read_norm_meta +from viscy_data._utils import read_norm_meta class ClassificationDataset(Dataset): @@ -100,7 +100,7 @@ def __getitem__(self, idx) -> tuple[Tensor, Tensor] | tuple[Tensor, Tensor, dict slice(x - x_half, x + x_half), ] ).float()[None] - norm_meta = _read_norm_meta(fov) + norm_meta = read_norm_meta(fov) if norm_meta is None: raise ValueError(f"Normalization metadata not found for FOV '{fov_name}'.") norm_meta = norm_meta[self.channel_name]["fov_statistics"] diff --git a/packages/viscy-data/src/viscy_data/gpu_aug.py b/packages/viscy-data/src/viscy_data/gpu_aug.py index b661a1c08..0ea48d302 100644 --- a/packages/viscy-data/src/viscy_data/gpu_aug.py +++ b/packages/viscy-data/src/viscy_data/gpu_aug.py @@ -19,7 +19,7 @@ from torch.utils.data import DataLoader, Dataset from viscy_data._typing import DictTransform, NormMeta -from viscy_data._utils import _ensure_channel_list, _read_norm_meta +from viscy_data._utils import _ensure_channel_list, read_norm_meta from viscy_data.distributed import ShardedDistributedSampler from viscy_data.select import SelectWell @@ -163,7 +163,7 @@ def __init__( self._metadata_map: dict[int, _CacheMetadata] = {} for position in positions: img = position[array_key] - norm_meta = _read_norm_meta(position) + norm_meta = read_norm_meta(position) for time_idx in range(img.frames): cache_map[key] = None self._metadata_map[key] = (position, time_idx, norm_meta) diff --git a/packages/viscy-data/src/viscy_data/mmap_cache.py b/packages/viscy-data/src/viscy_data/mmap_cache.py index d39fefd27..954a70349 100644 --- a/packages/viscy-data/src/viscy_data/mmap_cache.py +++ b/packages/viscy-data/src/viscy_data/mmap_cache.py @@ -23,7 +23,7 @@ MemoryMappedTensor = None from viscy_data._typing import DictTransform, NormMeta -from viscy_data._utils import _ensure_channel_list, _read_norm_meta +from viscy_data._utils import _ensure_channel_list, read_norm_meta from viscy_data.gpu_aug import GPUTransformDataModule from viscy_data.select import SelectWell @@ -75,7 +75,7 @@ def __init__( self._metadata_map: dict[int, _CacheMetadata] = {} for position in positions: img = position[array_key] - norm_meta = _read_norm_meta(position) + norm_meta = read_norm_meta(position) for time_idx in range(img.frames): cache_map[key] = None self._metadata_map[key] = (position, time_idx, norm_meta) diff --git a/packages/viscy-data/src/viscy_data/sliding_window.py b/packages/viscy-data/src/viscy_data/sliding_window.py index 7e109f555..c66258da2 100644 --- a/packages/viscy-data/src/viscy_data/sliding_window.py +++ b/packages/viscy-data/src/viscy_data/sliding_window.py @@ -12,7 +12,7 @@ from torch.utils.data import Dataset from viscy_data._typing import ChannelMap, DictTransform, HCSStackIndex, NormMeta, Sample -from viscy_data._utils import _ensure_channel_list, _read_norm_meta, _search_int_in_str +from viscy_data._utils import _ensure_channel_list, _search_int_in_str, read_norm_meta from viscy_data.foreground_masks import ForegroundMaskSupport _logger = logging.getLogger("lightning.pytorch") @@ -134,7 +134,7 @@ def _get_windows(self) -> None: w += ts * zs self.window_keys.append(w) self.window_arrays.append(img_arr) - self.window_norm_meta.append(_read_norm_meta(fov)) + self.window_norm_meta.append(read_norm_meta(fov)) if self.fg_mask_support is not None: self.fg_mask_support.validate_and_store(fov, img_arr, self.target_ch_idx) self._max_window = w diff --git a/packages/viscy-data/src/viscy_data/triplet.py b/packages/viscy-data/src/viscy_data/triplet.py index deed9fe57..41991be0c 100644 --- a/packages/viscy-data/src/viscy_data/triplet.py +++ b/packages/viscy-data/src/viscy_data/triplet.py @@ -31,8 +31,8 @@ from viscy_data._typing import ULTRACK_INDEX_COLUMNS, NormMeta from viscy_data._utils import ( - _read_norm_meta, _transform_channel_wise, + read_norm_meta, ) from viscy_data.hcs import HCSDataModule from viscy_data.select import _filter_fovs, _filter_wells @@ -239,7 +239,7 @@ def _slice_patch(self, track_row: "pd.Series") -> "tuple[ts.TensorStore, NormMet slice(y_center - y_half, y_center + y_half), slice(x_center - x_half, x_center + x_half), ] - return patch, _read_norm_meta(position) + return patch, read_norm_meta(position) def _slice_patches(self, track_rows: "pd.DataFrame"): """Slice and stack patches for multiple track rows.""" From e0f2fc078a16a2b2ab05f66bb2e185dd22fcc3b3 Mon Sep 17 00:00:00 2001 From: Ivan Ivanov Date: Fri, 5 Jun 2026 09:53:45 -0700 Subject: [PATCH 2/4] Limit n>=1 in `rotation_tta_transforms` Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- applications/cytoland/src/cytoland/engine.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/applications/cytoland/src/cytoland/engine.py b/applications/cytoland/src/cytoland/engine.py index 9a1be7b75..83ca7ea66 100644 --- a/applications/cytoland/src/cytoland/engine.py +++ b/applications/cytoland/src/cytoland/engine.py @@ -94,6 +94,8 @@ def rotation_tta_transforms( tuple[list[Callable], list[Callable]] The forward and inverse transform lists. """ + if n < 1: + raise ValueError(f"n must be >= 1, got {n}") forward = [partial(torch.rot90, k=k, dims=(-2, -1)) for k in range(n)] inverse = [partial(torch.rot90, k=-k, dims=(-2, -1)) for k in range(n)] return forward, inverse From 0fd09229295cc30f0acefa515a0a209825c0b2a8 Mon Sep 17 00:00:00 2001 From: Ivan Ivanov Date: Fri, 5 Jun 2026 10:03:51 -0700 Subject: [PATCH 3/4] refactor(viscy-data): keep _read_norm_meta as alias for read_norm_meta Instead of renaming every caller, keep the private name working: define the public read_norm_meta and add `_read_norm_meta = read_norm_meta`. Revert the call sites (triplet, sliding_window, gpu_aug, mmap_cache, cell_classification, dynaclr) back to the private name so existing code is untouched. __all__ still exports the public read_norm_meta. Co-Authored-By: Claude Opus 4.8 (1M context) --- applications/dynaclr/src/dynaclr/data/dataset.py | 4 ++-- packages/viscy-data/src/viscy_data/_utils.py | 4 ++++ packages/viscy-data/src/viscy_data/cell_classification.py | 4 ++-- packages/viscy-data/src/viscy_data/gpu_aug.py | 4 ++-- packages/viscy-data/src/viscy_data/mmap_cache.py | 4 ++-- packages/viscy-data/src/viscy_data/sliding_window.py | 4 ++-- packages/viscy-data/src/viscy_data/triplet.py | 4 ++-- 7 files changed, 16 insertions(+), 12 deletions(-) diff --git a/applications/dynaclr/src/dynaclr/data/dataset.py b/applications/dynaclr/src/dynaclr/data/dataset.py index 06acf3b35..a682b744e 100644 --- a/applications/dynaclr/src/dynaclr/data/dataset.py +++ b/applications/dynaclr/src/dynaclr/data/dataset.py @@ -36,7 +36,7 @@ from dynaclr.data.index import MultiExperimentIndex from dynaclr.data.tau_sampling import sample_tau from viscy_data._typing import ULTRACK_INDEX_COLUMNS, NormMeta, SampleMeta -from viscy_data._utils import read_norm_meta +from viscy_data._utils import _read_norm_meta def _pick_temporal_candidate( @@ -742,7 +742,7 @@ def _build_norm_meta( cache_key = (store_path, fov_name) if cache_key not in self._norm_meta_cache: position = self._get_position(store_path, fov_name) - self._norm_meta_cache[cache_key] = read_norm_meta(position) + self._norm_meta_cache[cache_key] = _read_norm_meta(position) cached = self._norm_meta_cache[cache_key] if cached is None: return None diff --git a/packages/viscy-data/src/viscy_data/_utils.py b/packages/viscy-data/src/viscy_data/_utils.py index 945741752..ae4589821 100644 --- a/packages/viscy-data/src/viscy_data/_utils.py +++ b/packages/viscy-data/src/viscy_data/_utils.py @@ -165,6 +165,10 @@ def read_norm_meta(fov: Position) -> NormMeta | None: return norm_meta +# Backwards-compatible private alias: existing callers import ``_read_norm_meta``. +_read_norm_meta = read_norm_meta + + def _collate_norm_meta(norm_metas: list[NormMeta]) -> NormMeta: """Stack per-sample norm_meta dicts into batched tensors. diff --git a/packages/viscy-data/src/viscy_data/cell_classification.py b/packages/viscy-data/src/viscy_data/cell_classification.py index c378747a8..b78a72cc1 100644 --- a/packages/viscy-data/src/viscy_data/cell_classification.py +++ b/packages/viscy-data/src/viscy_data/cell_classification.py @@ -21,7 +21,7 @@ from torch.utils.data import DataLoader, Dataset from viscy_data._typing import ULTRACK_INDEX_COLUMNS, AnnotationColumns -from viscy_data._utils import read_norm_meta +from viscy_data._utils import _read_norm_meta class ClassificationDataset(Dataset): @@ -100,7 +100,7 @@ def __getitem__(self, idx) -> tuple[Tensor, Tensor] | tuple[Tensor, Tensor, dict slice(x - x_half, x + x_half), ] ).float()[None] - norm_meta = read_norm_meta(fov) + norm_meta = _read_norm_meta(fov) if norm_meta is None: raise ValueError(f"Normalization metadata not found for FOV '{fov_name}'.") norm_meta = norm_meta[self.channel_name]["fov_statistics"] diff --git a/packages/viscy-data/src/viscy_data/gpu_aug.py b/packages/viscy-data/src/viscy_data/gpu_aug.py index 0ea48d302..b661a1c08 100644 --- a/packages/viscy-data/src/viscy_data/gpu_aug.py +++ b/packages/viscy-data/src/viscy_data/gpu_aug.py @@ -19,7 +19,7 @@ from torch.utils.data import DataLoader, Dataset from viscy_data._typing import DictTransform, NormMeta -from viscy_data._utils import _ensure_channel_list, read_norm_meta +from viscy_data._utils import _ensure_channel_list, _read_norm_meta from viscy_data.distributed import ShardedDistributedSampler from viscy_data.select import SelectWell @@ -163,7 +163,7 @@ def __init__( self._metadata_map: dict[int, _CacheMetadata] = {} for position in positions: img = position[array_key] - norm_meta = read_norm_meta(position) + norm_meta = _read_norm_meta(position) for time_idx in range(img.frames): cache_map[key] = None self._metadata_map[key] = (position, time_idx, norm_meta) diff --git a/packages/viscy-data/src/viscy_data/mmap_cache.py b/packages/viscy-data/src/viscy_data/mmap_cache.py index 954a70349..d39fefd27 100644 --- a/packages/viscy-data/src/viscy_data/mmap_cache.py +++ b/packages/viscy-data/src/viscy_data/mmap_cache.py @@ -23,7 +23,7 @@ MemoryMappedTensor = None from viscy_data._typing import DictTransform, NormMeta -from viscy_data._utils import _ensure_channel_list, read_norm_meta +from viscy_data._utils import _ensure_channel_list, _read_norm_meta from viscy_data.gpu_aug import GPUTransformDataModule from viscy_data.select import SelectWell @@ -75,7 +75,7 @@ def __init__( self._metadata_map: dict[int, _CacheMetadata] = {} for position in positions: img = position[array_key] - norm_meta = read_norm_meta(position) + norm_meta = _read_norm_meta(position) for time_idx in range(img.frames): cache_map[key] = None self._metadata_map[key] = (position, time_idx, norm_meta) diff --git a/packages/viscy-data/src/viscy_data/sliding_window.py b/packages/viscy-data/src/viscy_data/sliding_window.py index c66258da2..7e109f555 100644 --- a/packages/viscy-data/src/viscy_data/sliding_window.py +++ b/packages/viscy-data/src/viscy_data/sliding_window.py @@ -12,7 +12,7 @@ from torch.utils.data import Dataset from viscy_data._typing import ChannelMap, DictTransform, HCSStackIndex, NormMeta, Sample -from viscy_data._utils import _ensure_channel_list, _search_int_in_str, read_norm_meta +from viscy_data._utils import _ensure_channel_list, _read_norm_meta, _search_int_in_str from viscy_data.foreground_masks import ForegroundMaskSupport _logger = logging.getLogger("lightning.pytorch") @@ -134,7 +134,7 @@ def _get_windows(self) -> None: w += ts * zs self.window_keys.append(w) self.window_arrays.append(img_arr) - self.window_norm_meta.append(read_norm_meta(fov)) + self.window_norm_meta.append(_read_norm_meta(fov)) if self.fg_mask_support is not None: self.fg_mask_support.validate_and_store(fov, img_arr, self.target_ch_idx) self._max_window = w diff --git a/packages/viscy-data/src/viscy_data/triplet.py b/packages/viscy-data/src/viscy_data/triplet.py index 41991be0c..deed9fe57 100644 --- a/packages/viscy-data/src/viscy_data/triplet.py +++ b/packages/viscy-data/src/viscy_data/triplet.py @@ -31,8 +31,8 @@ from viscy_data._typing import ULTRACK_INDEX_COLUMNS, NormMeta from viscy_data._utils import ( + _read_norm_meta, _transform_channel_wise, - read_norm_meta, ) from viscy_data.hcs import HCSDataModule from viscy_data.select import _filter_fovs, _filter_wells @@ -239,7 +239,7 @@ def _slice_patch(self, track_row: "pd.Series") -> "tuple[ts.TensorStore, NormMet slice(y_center - y_half, y_center + y_half), slice(x_center - x_half, x_center + x_half), ] - return patch, read_norm_meta(position) + return patch, _read_norm_meta(position) def _slice_patches(self, track_rows: "pd.DataFrame"): """Slice and stack patches for multiple track rows.""" From 2e9bf7b98d28b482bb2db3d58be18feca399b3c9 Mon Sep 17 00:00:00 2001 From: Ivan Ivanov Date: Tue, 9 Jun 2026 10:14:09 -0700 Subject: [PATCH 4/4] refactor(viscy-data): make _read_norm_meta canonical, read_norm_meta the alias Flip the alias direction so the canonical definition keeps its private name in the private _utils module (where all internal callers import _read_norm_meta) and read_norm_meta is the public re-exported alias. Co-Authored-By: Claude Opus 4.8 (1M context) --- packages/viscy-data/src/viscy_data/_utils.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/packages/viscy-data/src/viscy_data/_utils.py b/packages/viscy-data/src/viscy_data/_utils.py index ae4589821..cad79cb7d 100644 --- a/packages/viscy-data/src/viscy_data/_utils.py +++ b/packages/viscy-data/src/viscy_data/_utils.py @@ -136,7 +136,7 @@ def _collate_samples(batch: Sequence[Sample]) -> Sample: return collated -def read_norm_meta(fov: Position) -> NormMeta | None: +def _read_norm_meta(fov: Position) -> NormMeta | None: """Read normalization metadata from the FOV. Convert to float32 tensors to avoid automatic casting to float64. @@ -165,8 +165,10 @@ def read_norm_meta(fov: Position) -> NormMeta | None: return norm_meta -# Backwards-compatible private alias: existing callers import ``_read_norm_meta``. -_read_norm_meta = read_norm_meta +# Public alias re-exported as ``viscy_data.read_norm_meta``; the canonical +# definition keeps its private name since this is a private ``_utils`` module +# and all internal callers import ``_read_norm_meta``. +read_norm_meta = _read_norm_meta def _collate_norm_meta(norm_metas: list[NormMeta]) -> NormMeta: