diff --git a/src/hoct/_api.py b/src/hoct/_api.py index 8b266ed..98c2cf3 100644 --- a/src/hoct/_api.py +++ b/src/hoct/_api.py @@ -139,6 +139,7 @@ def _create_dataset( tiling_scheme: TilingScheme | None = None, window_size: int = 5, test_time_augs: int = 0, + scale: tuple[float, ...] | None = None, ) -> FrameDataset | TiledRoiDataset | GraphConcatDataset: """ Create a dataset from a graph. @@ -153,23 +154,47 @@ def _create_dataset( The window size to use for the dataset. test_time_augs : int, default=0 The number of test time augmentations to use for the dataset. + scale : tuple[float, ...] | None, default=None + Physical spacing (t, [z,] y, x) to apply to dataset spatial features. + The scaling is deterministic; ``t`` is used for graph construction and + is not applied to spatial dataframe features. Returns ------- FrameDataset | TiledRoiDataset | GraphConcatDataset The created dataset. """ - if test_time_augs > 0: - df_transforms = [ - Flip(columns=["z", "y", "x"], p=0.5), + df_transforms = [] + if scale is not None: + spatial_scale: tuple[float, ...] = scale[1:] + if len(spatial_scale) == 2: + # 2D+t inputs have a singleton z axis in the graph. + spatial_scale = (1.0, *spatial_scale) + elif len(spatial_scale) != 3: + raise ValueError(f"Scale must have 3 or 4 elements (t, [z,] y, x), got {len(scale)}") + + # Scaling is a data transform rather than just an edge-construction + # parameter. Use fixed ranges so every dataset item receives the same + # physical scaling before any optional random test-time augmentation. + df_transforms.append( Affine( - degree_range=(-180, 180), - scale_range=(1, 1), + degree_range=(0, 0), + scale_range=[(value, value) for value in spatial_scale], shear_range=((0, 0), (0, 0)), - ), - ] - else: - df_transforms = [] + ) + ) + + if test_time_augs > 0: + df_transforms.extend( + [ + Flip(columns=["z", "y", "x"], p=0.5), + Affine( + degree_range=(-180, 180), + scale_range=[(1, 1), (1, 1), (1, 1)], + shear_range=((0, 0), (0, 0)), + ), + ] + ) if tiling_scheme is not None: LOG.info("Creating tiled ROI dataset") @@ -240,6 +265,7 @@ def predict( Maximum temporal gap for edges. scale : tuple[float, ...] | None Physical spacing (t, [z,] y, x). If None, uses isotropic spacing. + If provided `distance_threshold` is in physical units and features are scaled to physical units. window_size : int Temporal window size for the frame dataset. Only used if tiling_scheme is None. tiling_scheme : TilingScheme | None @@ -321,7 +347,7 @@ def predict( LOG.info(f"Created graph with {graph.num_nodes()} nodes and {graph.num_edges()} edges") - dataset = _create_dataset(graph, tiling_scheme, window_size, test_time_augs) + dataset = _create_dataset(graph, tiling_scheme, window_size, test_time_augs, scale) LOG.info("Running model inference and solving tracking") solution_graph = model_predict(model, dataset, solver_config=solver_config, return_solution=return_solution) diff --git a/src/hoct/_tests/test_api.py b/src/hoct/_tests/test_api.py index a8172dd..7ea0d44 100644 --- a/src/hoct/_tests/test_api.py +++ b/src/hoct/_tests/test_api.py @@ -4,8 +4,10 @@ """ import numpy as np +import polars as pl import pytest +import hoct._api as api from hoct.features.constants import REGIONPROPS from hoct.features.graph import create_graph from hoct.tracking import ILPSolverConfig @@ -99,6 +101,69 @@ def test_inference_mode_no_gt_features(self, synthetic_2d_labels): edge_attrs = graph.edge_attr_keys() assert "edge_is_gt" not in edge_attrs + def test_scaled_position_attributes_are_registered(self, synthetic_2d_labels): + """Test that scaled positions are available to candidate edge creation.""" + graph = create_graph( + labels=synthetic_2d_labels, + distance_threshold=300.0, + n_neighbors=5, + delta_t=3, + scale=(2.0, 3.0, 4.0), + ) + + node_attrs = graph.node_attr_keys() + assert {"scaled_t", "scaled_z", "scaled_y", "scaled_x"}.issubset(node_attrs) + + attrs = graph.node_attrs(attr_keys=["t", "z", "y", "x", "scaled_t", "scaled_z", "scaled_y", "scaled_x"]) + assert attrs["scaled_t"].to_list() == pytest.approx([2.0 * value for value in attrs["t"].to_list()]) + assert attrs["scaled_z"].to_list() == pytest.approx([0.0] * len(attrs)) + assert attrs["scaled_y"].to_list() == pytest.approx([3.0 * value for value in attrs["y"].to_list()]) + assert attrs["scaled_x"].to_list() == pytest.approx([4.0 * value for value in attrs["x"].to_list()]) + + +class TestCreateDataset: + """Tests for dataset-level data transforms.""" + + def test_scale_is_applied_deterministically(self, monkeypatch: pytest.MonkeyPatch) -> None: + class DummyDataset: + def __init__(self, **kwargs: object) -> None: + self.df_transforms = kwargs["df_transforms"] + + monkeypatch.setattr(api, "FrameDataset", DummyDataset) + dataset = api._create_dataset(graph=object(), scale=(1.0, 2.0, 3.0, 4.0)) + + assert len(dataset.df_transforms) == 1 + data = pl.DataFrame( + { + "z": [0.0, 1.0, 2.0], + "y": [0.0, 1.0, 2.0], + "x": [0.0, 1.0, 2.0], + "area": [1.0, 2.0, 3.0], + } + ) + transformed = dataset.df_transforms[0](data) + + assert transformed["z"].to_list() == pytest.approx([0.0, 2.0, 4.0]) + assert transformed["y"].to_list() == pytest.approx([0.0, 3.0, 6.0]) + assert transformed["x"].to_list() == pytest.approx([0.0, 4.0, 8.0]) + assert transformed["area"].to_list() == pytest.approx([24.0, 48.0, 72.0]) + assert transformed.equals(dataset.df_transforms[0](data)) + + def test_scale_adds_singleton_z_for_2d_data(self, monkeypatch: pytest.MonkeyPatch) -> None: + class DummyDataset: + def __init__(self, **kwargs: object) -> None: + self.df_transforms = kwargs["df_transforms"] + + monkeypatch.setattr(api, "FrameDataset", DummyDataset) + dataset = api._create_dataset(graph=object(), scale=(1.0, 2.0, 3.0)) + data = pl.DataFrame({"y": [0.0, 1.0], "x": [0.0, 1.0], "area": [1.0, 2.0]}) + + transformed = dataset.df_transforms[0](data) + + assert transformed["y"].to_list() == pytest.approx([0.0, 2.0]) + assert transformed["x"].to_list() == pytest.approx([0.0, 3.0]) + assert transformed["area"].to_list() == pytest.approx([6.0, 12.0]) + class TestSolverConfig: """Tests for ILPSolverConfig validation and immutability.""" diff --git a/src/hoct/_tests/test_transforms.py b/src/hoct/_tests/test_transforms.py index edbaec8..90487a2 100644 --- a/src/hoct/_tests/test_transforms.py +++ b/src/hoct/_tests/test_transforms.py @@ -95,6 +95,40 @@ def test_3d_coordinates() -> None: # z coordinate should also change due to rotation assert not result["z"].equals(df["z"]) + @staticmethod + def test_per_axis_scaling() -> None: + df = pl.DataFrame( + { + "z": [0.0, 1.0, 2.0], + "y": [0.0, 1.0, 2.0], + "x": [0.0, 1.0, 2.0], + "area": [1.0, 2.0, 3.0], + } + ) + transform = Affine( + degree_range=(0, 0), + scale_range=[(2.0, 2.0), (3.0, 3.0), (4.0, 4.0)], + shear_range=None, + ) + + result = transform(df) + + assert result["z"].to_list() == pytest.approx([0.0, 2.0, 4.0]) + assert result["y"].to_list() == pytest.approx([0.0, 3.0, 6.0]) + assert result["x"].to_list() == pytest.approx([0.0, 4.0, 8.0]) + assert result["area"].to_list() == pytest.approx([24.0, 48.0, 72.0]) + + @staticmethod + def test_per_axis_scaling_is_deterministic() -> None: + df = pl.DataFrame({"y": [0.0, 1.0], "x": [0.0, 1.0], "area": [1.0, 2.0]}) + transform = Affine( + degree_range=(0, 0), + scale_range=[(2.0, 2.0), (3.0, 3.0)], + shear_range=None, + ) + + assert transform(df).equals(transform(df)) + @staticmethod def test_affine_transformations_correctness() -> None: """Test mathematical correctness of affine transformations using static method.""" diff --git a/src/hoct/data/_transforms.py b/src/hoct/data/_transforms.py index 39c1b56..192382a 100644 --- a/src/hoct/data/_transforms.py +++ b/src/hoct/data/_transforms.py @@ -76,11 +76,24 @@ class Affine(BaseTransform): def __init__( self, degree_range: tuple[float, float] | None, - scale_range: tuple[float, float] | None, + scale_range: Sequence[tuple[float, float]] | tuple[float, float] | None, shear_range: tuple[tuple[float, float], tuple[float, float]] | None, ): self._degree_range = degree_range or (0, 0) - self._scale_range = scale_range or (1, 1) + self._scale_range: tuple[tuple[float, float], ...] + if scale_range is None: + self._scale_range = ((1.0, 1.0),) * 3 + elif len(scale_range) == 2 and all(np.isscalar(value) for value in scale_range): + # Keep accepting the old single range as a uniform range for every + # spatial axis. A sequence of ranges can then specify each axis + # independently (z, y, x). + self._scale_range = (tuple(float(value) for value in scale_range),) * 3 + else: + self._scale_range = tuple(tuple(float(value) for value in axis_range) for axis_range in scale_range) + if len(self._scale_range) not in (2, 3): + raise ValueError(f"Expected two or three scale ranges, got {len(self._scale_range)}") + if any(len(axis_range) != 2 for axis_range in self._scale_range): + raise ValueError("Each scale range must contain exactly two values") self._shear_range = shear_range or ((0, 0), (0, 0)) @staticmethod @@ -138,9 +151,6 @@ def _apply_affine_per_column(df: pl.DataFrame, affine: torch.Tensor, column: str def __call__(self, df: pl.DataFrame) -> pl.DataFrame: degrees = _uniform(1, self._degree_range).item() rad = np.deg2rad(degrees) - scales = _uniform(3, self._scale_range) - shear_y = _uniform(1, self._shear_range[0]).item() - shear_x = _uniform(1, self._shear_range[1]).item() if "z" in df.columns: sp_columns = ["z", "y", "x"] @@ -149,6 +159,19 @@ def __call__(self, df: pl.DataFrame) -> pl.DataFrame: sp_columns = ["y", "x"] ndim = 2 + if len(self._scale_range) == 3: + # The 3D affine matrix is sliced to its trailing dimensions for + # 2D data, so sample all three ranges and retain the corresponding + # trailing scales below. + sampled_scales = [_uniform(1, axis_range).item() for axis_range in self._scale_range] + scales = torch.tensor(sampled_scales) + elif len(self._scale_range) == ndim: + spatial_scales = [_uniform(1, axis_range).item() for axis_range in self._scale_range] + scales = torch.tensor([1.0, *spatial_scales]) + else: + raise ValueError(f"Expected {ndim} or 3 scale ranges, got {len(self._scale_range)}") + shear_y = _uniform(1, self._shear_range[0]).item() + shear_x = _uniform(1, self._shear_range[1]).item() affine = self._rotation_matrix(rad) @ np.diag(scales) @ self._shear_matrix(shear_y, shear_x) # original affine is 3D affine = affine[-ndim:, -ndim:] diff --git a/src/hoct/features/graph.py b/src/hoct/features/graph.py index 64de854..72f550c 100644 --- a/src/hoct/features/graph.py +++ b/src/hoct/features/graph.py @@ -186,6 +186,14 @@ def create_graph( if "intensity" in prop: graph.add_node_attr_key(prop, pl.Float32, 0.0) + for column in [f"scaled_{c}" for c in cols]: + graph.add_node_attr_key(column, pl.Float32, 0.0) + + graph.update_node_attrs( + attrs={f"scaled_{c}": node_attrs[f"scaled_{c}"].to_list() for c in cols}, + node_ids=node_attrs[td.DEFAULT_ATTR_KEYS.NODE_ID].to_list(), + ) + # Add candidate edges with td.options.Options(n_workers=1): td.edges.DistanceEdges( @@ -193,6 +201,7 @@ def create_graph( n_neighbors=n_neighbors, delta_t=delta_t, neighbors_per_frame=True, + attr_keys=[f"scaled_{c}" for c in cols], ).add_edges(graph) # Add required features