Skip to content
Open
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
46 changes: 36 additions & 10 deletions src/hoct/_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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")
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
65 changes: 65 additions & 0 deletions src/hoct/_tests/test_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down
34 changes: 34 additions & 0 deletions src/hoct/_tests/test_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down
33 changes: 28 additions & 5 deletions src/hoct/data/_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"]
Expand All @@ -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:]
Expand Down
9 changes: 9 additions & 0 deletions src/hoct/features/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,13 +186,22 @@ 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(
distance_threshold=distance_threshold,
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
Expand Down
Loading