diff --git a/benchmarks/graph_array_downscale.py b/benchmarks/graph_array_downscale.py new file mode 100644 index 00000000..5fe1d895 --- /dev/null +++ b/benchmarks/graph_array_downscale.py @@ -0,0 +1,150 @@ +"""ASV benchmarks for ``GraphArrayView``'s ``downscale`` parameter. + +Rendering a graph at full resolution allocates one whole-volume buffer per +timepoint and paints every mask at full resolution. ``downscale`` samples the +masks on a strided view instead, so both the buffer and the paint cost shrink +by roughly the product of the factors. + +``track_buffer_mbytes_*`` reports the cached buffer size, which is the metric +the parameter targets and is independent of machine speed. The ``time_*`` +methods measure the cold-fetch wall time that follows from it. + +The benchmarks skip themselves on revisions predating ``downscale`` so that +this file can live on ``main`` before the feature lands. +""" + +from __future__ import annotations + +import inspect +from typing import TYPE_CHECKING + +import numpy as np +import polars as pl + +from benchmarks.common import IS_CI +from tracksdata.array import GraphArrayView +from tracksdata.constants import DEFAULT_ATTR_KEYS +from tracksdata.nodes._mask import Mask + +if TYPE_CHECKING: + from tracksdata.graph import RustWorkXGraph + +FRAMES = 2 if IS_CI else 4 +# Anisotropic on purpose: matching factors to physical voxel size is the +# intended usage, so ``(1, 4, 4)`` is the case worth tracking. +SHAPE = (FRAMES, 32, 256, 256) +MASK_SIZE = (16, 56, 56) +NODES_PER_FRAME = 20 if IS_CI else 40 +COLD_FRAME = 0 + + +def _downscale_supported() -> bool: + """Whether the checked-out revision has the ``downscale`` parameter.""" + return "downscale" in inspect.signature(GraphArrayView.__init__).parameters + + +def _build_graph() -> RustWorkXGraph: + # In-memory on purpose: with SQLGraph, per-node mask decompression happens at + # full resolution whatever `downscale` is, which caps the ratio and obscures + # the paint cost this benchmark exists to track. + from tracksdata.graph import RustWorkXGraph + + graph = RustWorkXGraph() + graph.add_node_attr_key("label", dtype=pl.Int64) + graph.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 6)) + + rng = np.random.default_rng(0) + label = 1 + for t in range(FRAMES): + starts = np.stack( + [rng.integers(0, s - m, size=NODES_PER_FRAME) for s, m in zip(SHAPE[1:], MASK_SIZE, strict=True)], + axis=1, + ) + for start in starts: + bbox = np.concatenate([start, start + np.asarray(MASK_SIZE)]) + graph.add_node( + { + DEFAULT_ATTR_KEYS.T: int(t), + "label": label, + DEFAULT_ATTR_KEYS.MASK: Mask(np.ones(MASK_SIZE, dtype=bool), bbox=bbox), + DEFAULT_ATTR_KEYS.BBOX: bbox, + }, + validate_keys=False, + ) + label += 1 + return graph + + +class DownscaleBenchmark: + """Cold whole-frame fetches at several downscale factors.""" + + # Each timed call builds a fresh (empty-cache) view, so one invocation per + # sample is correct; batching them would measure the cache instead. + number = 1 + timeout = 300 + + # Nested on purpose: asv reads a flat `params` whose first entry is a sequence as a + # multi-parameter spec, so the tuple factors below would be split into separate + # parameters. Tuples also keep these immutable, as the other benchmarks here do. + params = ((None, (1, 2, 2), (1, 4, 4), (2, 2, 2), 4),) + param_names = ("downscale",) + + def setup(self, downscale) -> None: + if downscale is not None and not _downscale_supported(): + raise NotImplementedError("GraphArrayView has no `downscale` parameter on this revision") + self.graph = _build_graph() + + def _make_view(self, downscale) -> GraphArrayView: + kwargs = {} if downscale is None else {"downscale": downscale} + return GraphArrayView( + graph=self.graph, + shape=SHAPE, + attr_key="label", + dtype=np.uint32, + **kwargs, + ) + + def time_cold_whole_frame(self, downscale) -> None: + np.asarray(self._make_view(downscale)[COLD_FRAME]) + + def track_buffer_mbytes(self, downscale) -> float: + """Cached buffer size for one timepoint, in MB.""" + view = self._make_view(downscale) + np.asarray(view[COLD_FRAME]) + return view._cache._store[COLD_FRAME].buffer.nbytes / 1e6 + + def track_labels_rendered(self, downscale) -> int: + """Objects still visible, to catch a factor that silently drops them.""" + view = self._make_view(downscale) + return int(len(np.unique(np.asarray(view[COLD_FRAME]))) - 1) + + +if __name__ == "__main__": + import time + + print(f"shape={SHAPE} masks/frame={NODES_PER_FRAME} mask={MASK_SIZE}") + print(f"downscale supported on this revision: {_downscale_supported()}\n") + bench = DownscaleBenchmark() + base_t = base_mb = None + for downscale in DownscaleBenchmark.params[0]: + try: + bench.setup(downscale) + except NotImplementedError as e: + print(f" {downscale!s:>10s} skipped ({e})") + continue + mbytes = bench.track_buffer_mbytes(downscale) + labels = bench.track_labels_rendered(downscale) + reps = 7 + start = time.perf_counter() + for _ in range(reps): + bench.time_cold_whole_frame(downscale) + wall = (time.perf_counter() - start) * 1e3 / reps + if base_t is None: + base_t, base_mb = wall, mbytes + print(f" {downscale!s:>10s} {wall:8.1f} ms {mbytes:8.2f} MB labels={labels}") + else: + print( + f" {downscale!s:>10s} {wall:8.1f} ms {mbytes:8.2f} MB " + f"labels={labels} ({base_t / wall:.1f}x faster, {base_mb / mbytes:.0f}x smaller)" + ) diff --git a/docs/examples/basic.py b/docs/examples/basic.py index 269d78f1..1d1f373a 100644 --- a/docs/examples/basic.py +++ b/docs/examples/basic.py @@ -134,8 +134,10 @@ def basic_tracking_example(show_napari_viewer: bool = True) -> None: print("Opening napari viewer...") viewer = napari.Viewer() - # Add original segmented labels - viewer.add_labels(track_labels, name="Tracked Labels") + # Add original segmented labels. + # For large 3D data, pass `downscale=(1, 4, 4)` to `to_napari_format` above to + # render the labels coarsely; this example's data is small enough not to need it. + viewer.add_labels(track_labels, name="Tracked Labels", scale=track_labels.scale) # Add tracking trajectories with lineage information viewer.add_tracks(tracks_df, graph=track_graph, name="Tracks") diff --git a/docs/faq.md b/docs/faq.md index ad243543..ace89d43 100644 --- a/docs/faq.md +++ b/docs/faq.md @@ -76,7 +76,29 @@ tracks_df, track_graph, track_labels = td.functional.to_napari_format( ) viewer = napari.Viewer() -viewer.add_labels(track_labels) +viewer.add_labels(track_labels, scale=track_labels.scale) viewer.add_tracks(tracks_df, graph=track_graph) napari.run() ``` + +Large 3D data materializes a lot of voxels per timepoint, which is slow to render. Pass +`downscale` to draw the labels coarsely, and pick per-axis factors that equalize the +*physical* voxel size — for anisotropic data that usually means leaving `z` alone: + +```python +import numpy as np + +tracks_df, track_graph, track_labels = td.functional.to_napari_format( + solution_graph, + shape=labels.shape, + mask_key="mask", + downscale=(1, 4, 4), + dtype=np.uint32, +) +``` + +Downscaling is nearest-neighbor. Objects thinner than the factor are drawn as a single +voxel so that they stay visible and selectable, but their shape is meaningless, and a +voxel covered by several objects may show either of them. Use the full-resolution view +for anything quantitative, and for editing: coordinates taken from a downscaled layer +refer to downscaled voxels and must not be written back into the graph as-is. diff --git a/src/tracksdata/array/_graph_array.py b/src/tracksdata/array/_graph_array.py index 69837a11..9db2db85 100644 --- a/src/tracksdata/array/_graph_array.py +++ b/src/tracksdata/array/_graph_array.py @@ -32,6 +32,36 @@ def _validate_shape( return shape +def _normalize_downscale( + downscale: int | Sequence[int] | None, + ndim: int, +) -> np.ndarray: + """ + Normalize `downscale` to an int64 vector of per-spatial-axis factors. + + Unlike `chunk_shape`, a short sequence is rejected rather than padded: `downscale` + changes the physical meaning of every voxel, so a wrong-length sequence would show + up as misregistered data rather than as an error. + """ + if downscale is None: + return np.ones(ndim, dtype=np.int64) + + if not np.issubdtype(np.asarray(downscale).dtype, np.integer): + raise ValueError(f"`downscale` must be integer, got {downscale!r}") + + if np.isscalar(downscale): + factors = np.full(ndim, downscale, dtype=np.int64) + else: + factors = np.asarray(downscale, dtype=np.int64).reshape(-1) + if len(factors) != ndim: + raise ValueError(f"`downscale` must have length {ndim}, got {len(factors)}") + + if np.any(factors < 1): + raise ValueError(f"`downscale` factors must be >= 1, got {factors.tolist()}") + + return factors + + def chain_indices(slicing1: ArrayIndex | None, slicing2: ArrayIndex | None) -> ArrayIndex: """Chain two array indexing operations into a single one. @@ -136,10 +166,37 @@ class GraphArrayView(BaseReadOnlyArray): shape : tuple[int, ...] | None, optional The shape of the array. If None, the shape is inferred from the graph metadata `shape` key. chunk_shape : tuple[int] | None, optional - The chunk shape for the array. If None, the default chunk size is used. + The chunk shape for the array, counted in *downscaled* voxels. + If None, the default chunk size is used. buffer_cache_size : int, optional The maximum number of buffers to keep in the cache for the array. If None, the default buffer cache size is used. + downscale : int | tuple[int, ...] | None, optional + Per-spatial-axis integer factors used to render the graph at a reduced + resolution, trading accuracy for memory and speed. `shape` is always given at + full resolution; `shape` of the resulting array is `ceil(shape / downscale)`, + with the time axis left untouched. + + Prefer the tuple form and pick factors that equalize *physical* voxel size, i.e. + `scale[i] * downscale[i]` roughly constant across spatial axes. Microscopy data + is usually already anisotropic, so an isotropic factor coarsens the axis that + has least to give: for `scale=(1.97, 0.485, 0.485)`, `downscale=(1, 4, 4)` both + saves more memory than `(2, 2, 2)` and preserves z. + + Rendering is nearest-neighbor (see `Mask.paint_buffer`). Objects thinner than + the factor are drawn as a single voxel so that they stay visible and selectable, + but their shape and size are meaningless, and the value of a voxel covered by + more than one object is unspecified. Use `downscale=1` for anything quantitative + such as metrics or CTC export. + + Coordinates read out of this array are in downscaled coordinates. Writing them + back into the graph without rescaling corrupts it, so editing a graph through a + view with `downscale != 1` is unsupported. + + See Also + -------- + [Mask.paint_buffer][tracksdata.nodes.Mask.paint_buffer]: + The rendering primitive, which documents the sampling convention. """ def __init__( @@ -152,6 +209,7 @@ def __init__( chunk_shape: tuple[int, ...] | int | None = None, buffer_cache_size: int | None = None, dtype: np.dtype | None = None, + downscale: int | tuple[int, ...] | None = None, ): if attr_key not in graph.node_attr_keys(return_ids=True): raise ValueError(f"Attribute key '{attr_key}' not found in graph. Expected '{graph.node_attr_keys()}'") @@ -176,7 +234,21 @@ def __init__( dtype = np.uint8 self._dtype = dtype - self.original_shape = _validate_shape(shape, graph, "GraphArrayView") + + # `full_shape`, `Mask.bbox` and `offset` are in graph (full-resolution) coordinates. + # `original_shape`, `_indices`, `chunk_shape` and the cache buffers are in view + # (downscaled) coordinates. Only `_fill_array` and `_bbox_to_slices` cross over. + # normalized to a tuple: the graph metadata may hold a list, which would make + # the public shape attributes compare unequal depending on the graph backend + self.full_shape = tuple(int(s) for s in _validate_shape(shape, graph, "GraphArrayView")) + self._downscale = _normalize_downscale(downscale, len(self.full_shape) - 1) + # `reindex` shallow-copies, so derived views share this array; keep it immutable + self._downscale.flags.writeable = False + self._offset_vec = self._offset_as_array(len(self.full_shape) - 1) + self.original_shape = ( + self.full_shape[0], + *((np.asarray(self.full_shape[1:], dtype=np.int64) + self._downscale - 1) // self._downscale).tolist(), + ) chunk_shape = chunk_shape or get_options().gav_chunk_shape if isinstance(chunk_shape, int): @@ -226,6 +298,38 @@ def dtype(self) -> np.dtype: """Returns the dtype of the array.""" return np.dtype(self._dtype) + @property + def downscale(self) -> tuple[int, ...]: + """Per-spatial-axis downscaling factors, in the order of the spatial axes.""" + return tuple(int(f) for f in self._downscale) + + @property + def scale(self) -> tuple[float, ...]: + """ + Voxel size of this array, aligned with `shape`, for use as a napari `scale`. + + This is a convenience for callers that have no other source of physical scale. + A caller that already tracks one must instead multiply its own scale by + `downscale`: using this property alongside an existing scale double-counts the + downscaling. Note also that when the graph carries no `scale` metadata this + falls back to `downscale` alone, which looks like a valid scale but is not. + + Slice steps are deliberately not folded in. + """ + base = self.graph.metadata.get(DEFAULT_METADATA_KEYS.SCALE) + if base is None: + base = np.ones(len(self._downscale)) + else: + base = np.asarray(base, dtype=float).reshape(-1) + if len(base) != len(self._downscale): + raise ValueError( + f"Graph metadata '{DEFAULT_METADATA_KEYS.SCALE}' has length {len(base)}, " + f"expected {len(self._downscale)} (spatial axes only, without time)." + ) + scale = (1.0, *(base * self._downscale).tolist()) + # drop axes that were indexed away, so `scale` stays aligned with `shape` + return tuple(s for s, ind in zip(scale, self._indices, strict=True) if not np.isscalar(ind)) + def __getitem__(self, index: ArrayIndex) -> "GraphArrayView": """Return a sliced view of the GraphArrayView. @@ -346,6 +450,13 @@ def _fill_array(self, time: int, volume_slicing: Sequence[slice], buffer: np.nda np.ndarray The filled buffer. """ + # the window is in view coordinates, the spatial index is in graph coordinates. + # coarse `[a, b)` covers full-resolution `[a * f, b * f)`; the index treats the + # upper corner as closed, so this stays over-inclusive, which is safe. + volume_slicing = tuple( + slice(int(s.start) * int(f), int(s.stop) * int(f)) + for s, f in zip(volume_slicing, self._downscale, strict=True) + ) subgraph = self._spatial_filter[(slice(time, time), *volume_slicing)] df = subgraph.node_attrs( attr_keys=[self._attr_key, DEFAULT_ATTR_KEYS.MASK], @@ -353,7 +464,7 @@ def _fill_array(self, time: int, volume_slicing: Sequence[slice], buffer: np.nda for mask, value in zip(df[DEFAULT_ATTR_KEYS.MASK], df[self._attr_key], strict=True): mask: Mask - mask.paint_buffer(buffer, value, offset=self._offset) + mask.paint_buffer(buffer, value, offset=self._offset_vec, downscale=self._downscale) def _offset_as_array(self, ndim: int) -> np.ndarray: """Normalize `offset` to a vector for each spatial axis.""" @@ -372,21 +483,27 @@ def _bbox_to_slices(self, bbox: Any) -> tuple[slice, ...] | None: Returns `None` when the bbox does not overlap the current array volume. """ bbox = np.asarray(bbox, dtype=np.int64).reshape(-1) - ndim = len(self.original_shape) - 1 + ndim = len(self.full_shape) - 1 if len(bbox) != 2 * ndim: raise ValueError(f"`bbox` must have length {2 * ndim}, got {len(bbox)}") - offset = self._offset_as_array(ndim) - start = bbox[:ndim] + offset - stop = bbox[ndim:] + offset + start = bbox[:ndim] + self._offset_vec + stop = bbox[ndim:] + self._offset_vec - shape = np.asarray(self.original_shape[1:], dtype=np.int64) + shape = np.asarray(self.full_shape[1:], dtype=np.int64) start = np.clip(start, 0, shape) stop = np.clip(stop, 0, shape) if np.any(stop <= start): return None + # Floor the start and ceil the stop. Painting can touch output voxels + # `[ceil(start / f), ceil(stop / f))`, and a too-small object instead gets a + # fallback voxel at `center // f`, which can fall *below* `ceil(start / f)`. + # Flooring the stop would leave the last partially covered voxel stale. + start = start // self._downscale + stop = (stop + self._downscale - 1) // self._downscale + return tuple(slice(int(s), int(e)) for s, e in zip(start, stop, strict=True)) def _invalidate_bbox(self, time_values: Sequence[Any], bboxes: Sequence[np.ndarray | None]) -> None: diff --git a/src/tracksdata/array/_test/test_graph_array.py b/src/tracksdata/array/_test/test_graph_array.py index 8543cb33..a76be319 100644 --- a/src/tracksdata/array/_test/test_graph_array.py +++ b/src/tracksdata/array/_test/test_graph_array.py @@ -8,7 +8,7 @@ from tracksdata.array import GraphArrayView from tracksdata.array._graph_array import chain_indices -from tracksdata.constants import DEFAULT_ATTR_KEYS +from tracksdata.constants import DEFAULT_ATTR_KEYS, DEFAULT_METADATA_KEYS from tracksdata.graph import BaseGraph from tracksdata.nodes import RegionPropsNodes from tracksdata.nodes._mask import Mask @@ -749,3 +749,401 @@ def test_graph_array_view_no_invalidation_when_mask_unchanged(graph_backend: Bas assert n_regions == 0 np.testing.assert_array_equal(array_view._cache._store[0].ready, np.ones((2, 2), dtype=bool)) + + +@fixture( + params=[ + ((10, 100, 100), 2), + ((10, 100, 100), 3), + ((10, 100, 100, 100), 4), + ((10, 100, 100, 100), (1, 4, 4)), + ] +) +def downscaled_graph_from_image(request, graph_backend) -> tuple[GraphArrayView, np.ndarray, tuple[int, ...]]: + """ + A graph rendered at a reduced resolution, alongside its dense reference. + + Objects are wide enough on every axis to survive sampling and far enough apart that + no output voxel is claimed by two of them, so exact comparisons are well defined. + """ + shape, downscale = request.param + label = np.zeros(shape, dtype=np.uint8) + for i in range(shape[0]): + label[i, 8:24, 8:24] = i + 1 + + RegionPropsNodes(extra_properties=["label"]).add_nodes(graph_backend, labels=label) + array_view = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=downscale) + return array_view, label, array_view.downscale + + +def _strided_reference(label: np.ndarray, factors: tuple[int, ...]) -> np.ndarray: + """Global-grid strided subsample, i.e. output voxel `o` takes coordinate `o * f`.""" + return label[(slice(None), *(slice(None, None, f) for f in factors))] + + +def test_downscale_shape(downscaled_graph_from_image) -> None: + """The spatial shape must be the ceiling of the full-resolution shape.""" + array_view, label, factors = downscaled_graph_from_image + + assert array_view.full_shape == label.shape + assert array_view.shape == _strided_reference(label, factors).shape + assert array_view.ndim == label.ndim + + +@pytest.mark.parametrize( + ("shape", "downscale", "expected"), + [ + ((10, 100, 100), 4, (10, 25, 25)), + ((10, 100, 100), 3, (10, 34, 34)), # ceil, not 33 + ((10, 100, 100, 100), (1, 3, 7), (10, 100, 34, 15)), + ], +) +def test_downscale_shape_non_divisible( + graph_backend: BaseGraph, + shape: tuple[int, ...], + downscale: int | tuple[int, ...], + expected: tuple[int, ...], +) -> None: + """A non-divisible shape must round up, so trailing voxels are never dropped.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 6)) + array_view = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=downscale) + + assert array_view.shape == expected + + +def test_downscale_matches_strided_reference(downscaled_graph_from_image) -> None: + """Rendering must match a global-grid strided subsample of the dense labels.""" + array_view, label, factors = downscaled_graph_from_image + expected = _strided_reference(label, factors) + + for t in range(array_view.shape[0]): + np.testing.assert_array_equal(array_view[t], expected[t]) + np.testing.assert_array_equal(array_view, expected) + + +def test_downscale_slicing(downscaled_graph_from_image) -> None: + """Slicing happens in downscaled coordinates and needs no special handling.""" + array_view, label, factors = downscaled_graph_from_image + expected = _strided_reference(label, factors) + + np.testing.assert_array_equal(array_view[2], expected[2]) + np.testing.assert_array_equal(array_view[:3], expected[:3]) + np.testing.assert_array_equal(array_view[[1, 3]], expected[[1, 3]]) + np.testing.assert_array_equal(array_view[:, 4, 2:7], expected[:, 4, 2:7]) + + +def test_downscale_one_is_strict_noop(graph_backend: BaseGraph) -> None: + """`downscale` of 1 must be indistinguishable from not passing it at all.""" + shape = (4, 32, 32) + label = np.zeros(shape, dtype=np.uint8) + for i in range(shape[0]): + label[i, 5:12, 6:14] = i + 1 + + RegionPropsNodes(extra_properties=["label"]).add_nodes(graph_backend, labels=label) + baseline = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label") + + for downscale in (1, (1, 1)): + array_view = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=downscale) + assert array_view.shape == baseline.shape + assert array_view.downscale == (1, 1) + np.testing.assert_array_equal(np.asarray(array_view), np.asarray(baseline)) + + +def test_downscale_small_objects_survive(graph_backend: BaseGraph) -> None: + """ + Objects too small to be sampled must still appear, so they stay selectable. + + This is what distinguishes the rendering from a plain strided subsample. + """ + shape = (1, 64, 64) + label = np.zeros(shape, dtype=np.uint8) + for i, (y, x) in enumerate([(3, 5), (17, 22), (33, 41), (50, 7)]): + label[0, y, x] = i + 1 + + RegionPropsNodes(extra_properties=["label"]).add_nodes(graph_backend, labels=label) + array_view = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=8) + + # none of the objects lies on the sampling grid + assert not _strided_reference(label, (8, 8)).any() + np.testing.assert_array_equal(np.unique(np.asarray(array_view[0])), [0, 1, 2, 3, 4]) + + +def test_downscale_buffer_memory(graph_backend: BaseGraph) -> None: + """The cached buffer must shrink by the product of the factors.""" + shape = (2, 32, 32, 32) + label = np.zeros(shape, dtype=np.uint8) + label[:, 4:12, 4:12, 4:12] = 1 + + RegionPropsNodes(extra_properties=["label"]).add_nodes(graph_backend, labels=label) + kwargs = {"graph": graph_backend, "shape": shape, "attr_key": "label"} + full_res = GraphArrayView(**kwargs) + downscaled = GraphArrayView(**kwargs, downscale=4) + + _ = np.asarray(full_res[0]) + _ = np.asarray(downscaled[0]) + + assert downscaled._cache._store[0].buffer.nbytes * 4**3 == full_res._cache._store[0].buffer.nbytes + + +@pytest.mark.parametrize("offset", [2, (2, 2)]) +def test_downscale_with_offset(graph_backend: BaseGraph, offset: int | tuple[int, int]) -> None: + """ + `offset` is in graph coordinates, so it must be added before dividing. + + The geometry is chosen so that `ceil((bbox + offset) / f) != ceil(bbox / f)`. Most + bboxes do not satisfy this -- with `f=4` and `offset=2`, a bbox of `[2, 10)` samples + the same voxels whether the offset is applied before dividing, after dividing, or + dropped entirely -- so a carelessly placed object makes this test vacuous. + """ + shape = (1, 16, 16) + label = np.zeros(shape, dtype=np.uint8) + label[0, 3:11, 3:11] = 1 + + RegionPropsNodes(extra_properties=["label"]).add_nodes(graph_backend, labels=label) + array_view = GraphArrayView( + graph=graph_backend, + shape=shape, + attr_key="label", + offset=offset, + downscale=4, + ) + + # offset first (correct): [5, 13) -> [ceil(5/4), ceil(13/4)) = voxels 2, 3 + # offset after dividing: [ceil(3/4) + 2, ceil(11/4) + 2) = voxel 3 once clipped + # offset dropped: [ceil(3/4), ceil(11/4)) = voxels 1, 2 + painted = np.asarray(array_view[0]) + np.testing.assert_array_equal(np.unique(np.argwhere(painted)), [2, 3]) + + +def test_downscale_invalidation_chunk_grid(graph_backend: BaseGraph) -> None: + """Invalidation must target chunks in downscaled coordinates.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + array_view = GraphArrayView( + graph=graph_backend, + shape=(2, 16, 16), + attr_key="label", + chunk_shape=(4, 4), + downscale=2, + ) + assert array_view.shape == (2, 8, 8) + + _ = np.asarray(array_view[0]) + np.testing.assert_array_equal(array_view._cache._store[0].ready, np.ones((2, 2), dtype=bool)) + + # full-resolution bbox [10, 10, 12, 12] -> output [5, 5, 6, 6] -> chunk (1, 1) only + graph_backend.add_node( + { + DEFAULT_ATTR_KEYS.T: 0, + DEFAULT_ATTR_KEYS.BBOX: np.array([10, 10, 12, 12]), + DEFAULT_ATTR_KEYS.MASK: Mask(np.ones((2, 2), dtype=bool), bbox=np.array([10, 10, 12, 12])), + "label": 1, + } + ) + + np.testing.assert_array_equal( + array_view._cache._store[0].ready, + np.array([[True, True], [True, False]]), + ) + + +def test_downscale_invalidation_ceil_boundary(graph_backend: BaseGraph) -> None: + """ + The invalidated stop must be ceiled, not floored. + + A floored stop passes every shape and equivalence test and only shows up here, as a + stale label left behind by the partially covered trailing voxel. + """ + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + array_view = GraphArrayView( + graph=graph_backend, + shape=(1, 16, 16), + attr_key="label", + chunk_shape=(4, 4), + downscale=2, + ) + + node_id = graph_backend.add_node( + { + DEFAULT_ATTR_KEYS.T: 0, + DEFAULT_ATTR_KEYS.BBOX: np.array([2, 2, 9, 9]), + DEFAULT_ATTR_KEYS.MASK: Mask(np.ones((7, 7), dtype=bool), bbox=np.array([2, 2, 9, 9])), + "label": 1, + } + ) + rendered = np.asarray(array_view[0]) + # full-resolution stop 9 -> output stop ceil(9 / 2) = 5, crossing into chunk 1 + assert rendered[4, 4] == 1 + + graph_backend.remove_node(node_id) + np.testing.assert_array_equal(np.asarray(array_view[0]), np.zeros((8, 8))) + + +def test_downscale_invalidation_covers_fallback_voxel(graph_backend: BaseGraph) -> None: + """ + A fallback voxel can sit below `ceil(start / f)`, so the start must be floored. + + For a bbox of [2, 3) with a factor of 4 the sampled range is empty while the + fallback lands at `2 // 4 == 0`. + """ + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + array_view = GraphArrayView(graph=graph_backend, shape=(1, 16, 16), attr_key="label", downscale=4) + + node_id = graph_backend.add_node( + { + DEFAULT_ATTR_KEYS.T: 0, + DEFAULT_ATTR_KEYS.BBOX: np.array([2, 2, 3, 3]), + DEFAULT_ATTR_KEYS.MASK: Mask(np.ones((1, 1), dtype=bool), bbox=np.array([2, 2, 3, 3])), + "label": 1, + } + ) + assert np.asarray(array_view[0])[0, 0] == 1 + + graph_backend.remove_node(node_id) + assert not np.asarray(array_view[0]).any() + + +def test_downscale_scale_property(graph_backend: BaseGraph) -> None: + """`scale` must combine the graph metadata with the factors and stay aligned with `shape`.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 6)) + shape = (2, 16, 16, 16) + + without_metadata = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=(1, 2, 2)) + assert without_metadata.scale == (1.0, 1.0, 2.0, 2.0) + + graph_backend.metadata[DEFAULT_METADATA_KEYS.SCALE] = (2.0, 0.5, 0.5) + array_view = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=(1, 2, 2)) + assert array_view.scale == (1.0, 2.0, 1.0, 1.0) + + # indexing away the time axis must drop its entry too + assert len(array_view[0].scale) == len(array_view[0].shape) + assert array_view[0].scale == (2.0, 1.0, 1.0) + + +@pytest.mark.parametrize("downscale", [0, -1, 1.5, (2, 2, 2)]) +def test_downscale_validation(graph_backend: BaseGraph, downscale: int | float | tuple[int, ...]) -> None: + """Invalid factors must be rejected at construction.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + + with pytest.raises(ValueError, match="downscale"): + GraphArrayView(graph=graph_backend, shape=(2, 16, 16), attr_key="label", downscale=downscale) + + +def test_downscale_larger_than_axis(graph_backend: BaseGraph) -> None: + """A factor larger than an axis must collapse it to one voxel, keeping the object.""" + shape = (2, 3, 3) + label = np.zeros(shape, dtype=np.uint8) + label[:, 1, 1] = 1 + + RegionPropsNodes(extra_properties=["label"]).add_nodes(graph_backend, labels=label) + array_view = GraphArrayView(graph=graph_backend, shape=shape, attr_key="label", downscale=8) + + assert array_view.shape == (2, 1, 1) + np.testing.assert_array_equal(np.asarray(array_view[0]), [[1]]) + + +def test_downscale_scale_metadata_length_mismatch(graph_backend: BaseGraph) -> None: + """A wrong-length `scale` metadata entry must raise, not broadcast silently.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 6)) + graph_backend.metadata[DEFAULT_METADATA_KEYS.SCALE] = (0.5,) + + array_view = GraphArrayView(graph=graph_backend, shape=(2, 16, 16, 16), attr_key="label", downscale=2) + + with pytest.raises(ValueError, match="expected 3"): + _ = array_view.scale + + +def test_downscale_shape_attrs_are_tuples(graph_backend: BaseGraph) -> None: + """Shape attributes must not depend on how the backend stores `shape` metadata.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + graph_backend.metadata.update(shape=(4, 16, 16)) + + array_view = GraphArrayView(graph=graph_backend, shape=None, attr_key="label", downscale=2) + + assert array_view.full_shape == (4, 16, 16) + assert array_view.original_shape == (4, 8, 8) + + +def test_downscale_is_immutable(downscaled_graph_from_image) -> None: + """`downscale` must not be mutable: derived views share the underlying array.""" + array_view, _, factors = downscaled_graph_from_image + + with pytest.raises(ValueError, match="read-only"): + array_view._downscale[0] = 99 + + assert array_view[0].downscale == factors + + +def test_downscale_invalidation_on_move(graph_backend: BaseGraph) -> None: + """Moving a node must clear the coarse voxels it vacated and paint the new ones.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + + # chunks must be small enough that the old and new regions land in different ones, + # otherwise any invalidation clears the whole frame and the test proves nothing + array_view = GraphArrayView( + graph=graph_backend, + shape=(1, 32, 32), + attr_key="label", + chunk_shape=(2, 2), + downscale=4, + ) + node_id = graph_backend.add_node( + { + DEFAULT_ATTR_KEYS.T: 0, + DEFAULT_ATTR_KEYS.BBOX: np.array([4, 4, 12, 12]), + DEFAULT_ATTR_KEYS.MASK: Mask(np.ones((8, 8), dtype=bool), bbox=np.array([4, 4, 12, 12])), + "label": 1, + } + ) + np.testing.assert_array_equal(np.argwhere(np.asarray(array_view[0])), [[1, 1], [1, 2], [2, 1], [2, 2]]) + + new_bbox = np.array([20, 20, 28, 28]) + graph_backend.update_node_attrs( + attrs={ + DEFAULT_ATTR_KEYS.BBOX: [new_bbox], + DEFAULT_ATTR_KEYS.MASK: [Mask(np.ones((8, 8), dtype=bool), bbox=new_bbox)], + }, + node_ids=[node_id], + ) + + # the vacated voxels must be gone, not merely joined by the new ones + np.testing.assert_array_equal(np.argwhere(np.asarray(array_view[0])), [[5, 5], [5, 6], [6, 5], [6, 6]]) + + +def test_downscale_invalidation_on_attr_change(graph_backend: BaseGraph) -> None: + """Changing the displayed attribute must repaint the coarse voxels in place.""" + graph_backend.add_node_attr_key("label", dtype=pl.Int64) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph_backend.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + + array_view = GraphArrayView(graph=graph_backend, shape=(1, 32, 32), attr_key="label", downscale=4) + bbox = np.array([4, 4, 12, 12]) + node_id = graph_backend.add_node( + { + DEFAULT_ATTR_KEYS.T: 0, + DEFAULT_ATTR_KEYS.BBOX: bbox, + DEFAULT_ATTR_KEYS.MASK: Mask(np.ones((8, 8), dtype=bool), bbox=bbox), + "label": 1, + } + ) + assert np.asarray(array_view[0])[1, 1] == 1 + + graph_backend.update_node_attrs(attrs={"label": [7]}, node_ids=[node_id]) + assert np.asarray(array_view[0])[1, 1] == 7 diff --git a/src/tracksdata/functional/_napari.py b/src/tracksdata/functional/_napari.py index df1046b1..4771fab6 100644 --- a/src/tracksdata/functional/_napari.py +++ b/src/tracksdata/functional/_napari.py @@ -9,26 +9,40 @@ from tracksdata.graph._base_graph import BaseGraph if TYPE_CHECKING: + import numpy as np + from tracksdata.array._graph_array import GraphArrayView @overload def to_napari_format( graph: BaseGraph, - shape: tuple[int, ...] | None, - solution_key: str | None, - output_tracklet_id_key: str, - mask_key: None, + shape: tuple[int, ...] | None = ..., + solution_key: str | None = ..., + output_tracklet_id_key: str = ..., + *, + mask_key: None = ..., + chunk_shape: tuple[int] | None = ..., + buffer_cache_size: int | None = ..., + allow_frame_skip: bool = ..., + downscale: int | tuple[int, ...] | None = ..., + dtype: "np.dtype | None" = ..., ) -> tuple[pl.DataFrame, dict[int, int]]: ... @overload def to_napari_format( graph: BaseGraph, - shape: tuple[int, ...] | None, - solution_key: str | None, - output_tracklet_id_key: str, + shape: tuple[int, ...] | None = ..., + solution_key: str | None = ..., + output_tracklet_id_key: str = ..., + *, mask_key: str, + chunk_shape: tuple[int] | None = ..., + buffer_cache_size: int | None = ..., + allow_frame_skip: bool = ..., + downscale: int | tuple[int, ...] | None = ..., + dtype: "np.dtype | None" = ..., ) -> tuple[pl.DataFrame, dict[int, int], "GraphArrayView"]: ... @@ -41,6 +55,8 @@ def to_napari_format( chunk_shape: tuple[int] | None = None, buffer_cache_size: int | None = None, allow_frame_skip: bool = False, + downscale: int | tuple[int, ...] | None = None, + dtype: "np.dtype | None" = None, ) -> ( tuple[ pl.DataFrame, @@ -82,6 +98,16 @@ def to_napari_format( allow_frame_skip : bool, optional Whether to allow frame skipping when assigning tracklet ids. If True, tracklets can skip + downscale : int | tuple[int, ...] | None, optional + Render the labels layer at a reduced resolution. `shape` stays at full resolution. + See [GraphArrayView][tracksdata.array.GraphArrayView] for the trade-offs; in + short, prefer per-axis factors matched to the physical voxel size, and do not + edit a graph through a downscaled layer. + dtype : np.dtype | None, optional + The dtype of the labels layer. Worth passing together with `downscale`: the + inferred dtype follows the attribute column, which is often wider than tracklet + ids need, and passing it explicitly also skips the inference query. + Examples -------- @@ -89,6 +115,16 @@ def to_napari_format( labels = ... graph = ... tracks_data, dict_graph, array_view = to_napari_format(graph, labels.shape, mask_key="mask") + viewer.add_labels(array_view, scale=array_view.scale) + ``` + + For large 3D data, render coarsely. Pick factors that equalize physical voxel size, + which for anisotropic data means leaving `z` alone: + + ```python + tracks_data, dict_graph, array_view = to_napari_format( + graph, labels.shape, mask_key="mask", downscale=(1, 4, 4), dtype=np.uint32 + ) ``` Returns @@ -130,6 +166,8 @@ def to_napari_format( attr_key=output_tracklet_id_key, chunk_shape=chunk_shape, buffer_cache_size=buffer_cache_size, + downscale=downscale, + dtype=dtype, ) return tracks_data, dict_graph, array_view diff --git a/src/tracksdata/functional/_test/test_napari.py b/src/tracksdata/functional/_test/test_napari.py index 712cf53d..9e37e861 100644 --- a/src/tracksdata/functional/_test/test_napari.py +++ b/src/tracksdata/functional/_test/test_napari.py @@ -10,6 +10,27 @@ @pytest.mark.parametrize("metadata_shape", [True, False]) def test_napari_conversion(metadata_shape: bool) -> None: + _run_napari_conversion(metadata_shape) + + +def test_napari_conversion_downscale() -> None: + """`downscale` and `dtype` must reach the array view, with `shape` staying full-resolution.""" + array_view = _run_napari_conversion( + metadata_shape=False, + downscale=(1, 2, 2), + dtype=np.uint32, + ) + + assert array_view.full_shape == (2, 10, 22, 32) + assert array_view.shape == (2, 10, 11, 16) + assert array_view.downscale == (1, 2, 2) + assert array_view.dtype == np.uint32 + assert array_view.scale == (1.0, 1.0, 2.0, 2.0) + # the disks are still rendered, just coarsely + assert np.asarray(array_view[0]).any() + + +def _run_napari_conversion(metadata_shape: bool, **kwargs): positions = np.asarray( [ [0, 5, 10, 20], # t=0, z=5, y=10, x=20 @@ -55,6 +76,7 @@ def test_napari_conversion(metadata_shape: bool) -> None: graph, shape=arg_shape, mask_key=DEFAULT_ATTR_KEYS.MASK, + **kwargs, ) assert dict_graph == tracklet_id_graph @@ -66,7 +88,9 @@ def test_napari_conversion(metadata_shape: bool) -> None: positions, ) - assert array_view.shape == (2, 10, 22, 32) + if not kwargs: + assert array_view.shape == (2, 10, 22, 32) + np.testing.assert_equal(np.unique(array_view[0]), [0, 1]) + np.testing.assert_equal(np.unique(array_view[1]), [0, 2, 3]) - np.testing.assert_equal(np.unique(array_view[0]), [0, 1]) - np.testing.assert_equal(np.unique(array_view[1]), [0, 2, 3]) + return array_view diff --git a/src/tracksdata/graph/_graph_view.py b/src/tracksdata/graph/_graph_view.py index dfe7d8c8..84f8970f 100644 --- a/src/tracksdata/graph/_graph_view.py +++ b/src/tracksdata/graph/_graph_view.py @@ -13,6 +13,7 @@ from tracksdata.graph._mapped_graph_mixin import MappedGraphMixin from tracksdata.graph._rustworkx_graph import IndexedRXGraph, RustWorkXGraph, RXFilter from tracksdata.graph.filters._indexed_filter import IndexRXFilter +from tracksdata.utils._dataframe import unpack_array_attrs from tracksdata.utils._dtypes import AttrSchema from tracksdata.utils._signal import ( emit_node_added_events, @@ -107,6 +108,13 @@ class GraphView(MappedGraphMixin, RustWorkXGraph): sync : bool, default True Whether to automatically synchronize changes in the view. By default only the root graph is updated. + root_fallback : bool, default False + When True, reading a node attribute key this view does not hold is + served from the root graph instead of raising. This is what makes a + view built with a reduced `node_attr_keys` usable for the keys it left + out -- at the cost of going to the root's storage for them, which for a + SQL root is a query per read. Read path only -- writes already go + straight to the root for keys the view does not hold. Attributes ---------- @@ -157,6 +165,7 @@ def __init__( mode: ViewMode = ViewMode.WRITE_THROUGH, node_attr_keys: list[str] | None = None, edge_attr_keys: list[str] | None = None, + root_fallback: bool = False, ) -> None: # Initialize RustWorkXGraph RustWorkXGraph.__init__(self, rx_graph=None) # rx_graph is not used to avoid initialization @@ -211,6 +220,10 @@ def __init__( ] self._edge_attr_keys = list(dict.fromkeys(self._edge_attr_keys)) + # When set, a read for a node attribute key this view does not hold is + # served from the root instead of raising. See `_node_attrs_from_node_ids`. + self._root_fallback = root_fallback + # use parent graph overlaps self._overlaps = None @@ -363,6 +376,35 @@ def node_attr_keys(self, return_ids: bool = False) -> list[str]: pass return keys + def _validate_attr_keys( + self, + attr_keys: Sequence[str] | str | None, + mode: Literal["node", "edge"], + ) -> None: + """ + Same validation as any other graph, plus a pointer when the key is only + missing because this view was built without it. + + A view built with an explicit `node_attr_keys` list is otherwise a trap: + the key exists, the root has it, and the error says only that this graph + does not. + """ + try: + super()._validate_attr_keys(attr_keys, mode) + except KeyError as err: + if mode != "node" or self._root_fallback: + raise + if isinstance(attr_keys, str): + attr_keys = [attr_keys] + missing = set(attr_keys) - set(self.node_attr_keys(return_ids=True)) + on_root = sorted(missing & set(self._root.node_attr_keys(return_ids=True))) + if not on_root: + raise + raise KeyError( + f"{err.args[0]} The root graph does hold {on_root}; build the view with " + "`root_fallback=True` to read those through it." + ) from None + def edge_attr_keys(self, return_ids: bool = False) -> list[str]: """ Get the keys of the attributes of the edges. @@ -1037,6 +1079,26 @@ def predecessors( raise RuntimeError("Out of sync graph view cannot be used to get predecessors") return super().predecessors(node_ids, attr_keys, return_attrs=return_attrs) + def _split_node_attr_keys(self, attr_keys: Sequence[str]) -> tuple[list[str], list[str]]: + """ + Partition requested node attribute keys into those this view holds and + those only the root holds. + + Only ever returns root keys when the view was built with + ``root_fallback=True``; otherwise every key is reported as local and the + usual validation in the local read path rejects the ones that are not. + Keys that exist on neither the view nor the root still raise, here via + the root's own validation. + """ + if not self._root_fallback: + return list(attr_keys), [] + + local = set(self.node_attr_keys(return_ids=True)) + root_keys = [k for k in attr_keys if k not in local] + if root_keys: + self._root._validate_attr_keys(root_keys, "node") + return [k for k in attr_keys if k in local], root_keys + def _node_attrs_from_node_ids( self, *, @@ -1044,14 +1106,57 @@ def _node_attrs_from_node_ids( attr_keys: Sequence[str] | str | None = None, unpack: bool = False, ) -> pl.DataFrame: - node_dfs = super()._node_attrs_from_node_ids( - node_ids=self._map_to_local(node_ids), - attr_keys=attr_keys, - unpack=unpack, - ) - node_dfs = self._map_df_to_external( - node_dfs, [DEFAULT_ATTR_KEYS.NODE_ID, DEFAULT_ATTR_KEYS.EDGE_SOURCE, DEFAULT_ATTR_KEYS.EDGE_TARGET] + id_keys = [DEFAULT_ATTR_KEYS.NODE_ID, DEFAULT_ATTR_KEYS.EDGE_SOURCE, DEFAULT_ATTR_KEYS.EDGE_TARGET] + + if attr_keys is None: + # "every key" means every key this view holds -- a fallback would + # silently pull in the columns the view was built to leave out. + node_dfs = super()._node_attrs_from_node_ids( + node_ids=self._map_to_local(node_ids), + attr_keys=None, + unpack=unpack, + ) + return self._map_df_to_external(node_dfs, id_keys) + + if isinstance(attr_keys, str): + attr_keys = [attr_keys] + attr_keys = list(dict.fromkeys(attr_keys)) + + local_keys, root_keys = self._split_node_attr_keys(attr_keys) + + if not root_keys: + node_dfs = super()._node_attrs_from_node_ids( + node_ids=self._map_to_local(node_ids), + attr_keys=attr_keys, + unpack=unpack, + ) + return self._map_df_to_external(node_dfs, id_keys) + + node_id_key = DEFAULT_ATTR_KEYS.NODE_ID + external_ids = self.node_ids() if node_ids is None else list(node_ids) + + # NODE_ID is needed as the join key even when the caller did not ask for + # it, and `unpack` is deferred until after the join so that it sees the + # same columns it would have on a view holding all of them. + local_df = super()._node_attrs_from_node_ids( + node_ids=self._map_to_local(external_ids), + attr_keys=[node_id_key, *local_keys], + unpack=False, ) + local_df = self._map_df_to_external(local_df, id_keys) + + # Ask the root for exactly the rows this view covers -- never for all of + # its rows, which include the ones this view deliberately dropped. + root_df = self._root.filter(node_ids=external_ids).node_attrs(attr_keys=[node_id_key, *root_keys]) + + # `maintain_order="left"` is not optional: callers zip columns + # positionally, so a reordered join would pair the wrong values. + node_dfs = local_df.join(root_df, on=node_id_key, how="left", maintain_order="left") + node_dfs = node_dfs.select(attr_keys) + + if unpack: + node_dfs = unpack_array_attrs(node_dfs) + return node_dfs def node_attrs( diff --git a/src/tracksdata/graph/_rustworkx_graph.py b/src/tracksdata/graph/_rustworkx_graph.py index 98247e7c..9df10c0a 100644 --- a/src/tracksdata/graph/_rustworkx_graph.py +++ b/src/tracksdata/graph/_rustworkx_graph.py @@ -368,6 +368,7 @@ def subgraph( edge_attr_keys: Sequence[str] | None = None, *, mode: "ViewMode | None" = None, + root_fallback: bool = False, ) -> "GraphView": from tracksdata.graph._graph_view import GraphView, ViewMode @@ -387,6 +388,7 @@ def subgraph( mode=mode if mode is not None else ViewMode.WRITE_THROUGH, node_attr_keys=node_attr_keys, edge_attr_keys=edge_attr_keys, + root_fallback=root_fallback, ) return graph_view diff --git a/src/tracksdata/graph/_sql_graph.py b/src/tracksdata/graph/_sql_graph.py index d4ab8663..433dad73 100644 --- a/src/tracksdata/graph/_sql_graph.py +++ b/src/tracksdata/graph/_sql_graph.py @@ -524,6 +524,7 @@ def subgraph( edge_attr_keys: Sequence[str] | None = None, *, mode: "ViewMode | None" = None, + root_fallback: bool = False, ) -> "GraphView": from tracksdata.graph._graph_view import GraphView, ViewMode @@ -582,6 +583,7 @@ def subgraph( mode=mode if mode is not None else ViewMode.WRITE_THROUGH, node_attr_keys=node_attr_keys, edge_attr_keys=edge_attr_keys, + root_fallback=root_fallback, ) return graph diff --git a/src/tracksdata/graph/_test/test_root_fallback.py b/src/tracksdata/graph/_test/test_root_fallback.py new file mode 100644 index 00000000..feef566a --- /dev/null +++ b/src/tracksdata/graph/_test/test_root_fallback.py @@ -0,0 +1,353 @@ +"""Tests for `GraphView(root_fallback=True)`: reading node attribute keys a +partial view was built without, by falling back to its root graph.""" + +import numpy as np +import polars as pl +import pytest + +from tracksdata.array import GraphArrayView +from tracksdata.attrs import NodeAttr +from tracksdata.constants import DEFAULT_ATTR_KEYS +from tracksdata.graph import BaseGraph, GraphView +from tracksdata.nodes._mask import Mask + +LEAN_KEYS = ["x", "y"] + + +def _populate(graph: BaseGraph) -> BaseGraph: + """A graph with a cheap key set and one key a lean view will leave out.""" + graph.add_node_attr_key("x", dtype=pl.Float64) + graph.add_node_attr_key("y", dtype=pl.Float64) + graph.add_node_attr_key("label", dtype=pl.String, default_value="") + graph.add_node_attr_key("solution", dtype=pl.Boolean, default_value=False) + graph.add_node_attr_key("vec", dtype=pl.Array(pl.Int64, 2)) + + for i in range(6): + graph.add_node( + { + DEFAULT_ATTR_KEYS.T: i, + "x": float(i), + "y": float(i * 10), + "label": f"n{i}", + "solution": i % 2 == 0, + "vec": np.array([i, -i]), + } + ) + return graph + + +def _views(graph: BaseGraph, **kwargs) -> tuple[GraphView, GraphView]: + """A lean view and the equivalent full view, over the same nodes.""" + lean = graph.filter(NodeAttr("solution") == True).subgraph(node_attr_keys=LEAN_KEYS, root_fallback=True, **kwargs) + full = graph.filter(NodeAttr("solution") == True).subgraph() + return lean, full + + +# --- correctness of the fallback ------------------------------------------ + + +def test_fallback_matches_full_view(graph_backend: BaseGraph) -> None: + lean, full = _views(_populate(graph_backend)) + + assert "label" not in lean.node_attr_keys() + + keys = [DEFAULT_ATTR_KEYS.NODE_ID, "label"] + assert lean.node_attrs(attr_keys=keys).equals(full.node_attrs(attr_keys=keys)) + + +def test_mixed_request_keeps_requested_order(graph_backend: BaseGraph) -> None: + lean, full = _views(_populate(graph_backend)) + + keys = ["label", "x", DEFAULT_ATTR_KEYS.NODE_ID, "y"] + df = lean.node_attrs(attr_keys=keys) + + assert df.columns == keys + assert df.equals(full.node_attrs(attr_keys=keys)) + + +def test_fallback_without_node_id_requested(graph_backend: BaseGraph) -> None: + """NODE_ID is the join key but must not leak into the result.""" + lean, full = _views(_populate(graph_backend)) + + df = lean.node_attrs(attr_keys=["label", "x"]) + + assert df.columns == ["label", "x"] + assert df.equals(full.node_attrs(attr_keys=["label", "x"])) + + +def test_fallback_single_str_attr_key(graph_backend: BaseGraph) -> None: + lean, full = _views(_populate(graph_backend)) + + df = lean.node_attrs(attr_keys="label") + + assert df.columns == ["label"] + assert df.equals(full.node_attrs(attr_keys="label")) + + +def test_fallback_unpack(graph_backend: BaseGraph) -> None: + lean, full = _views(_populate(graph_backend)) + + keys = ["x", "vec"] + df = lean.node_attrs(attr_keys=keys, unpack=True) + + assert df.equals(full.node_attrs(attr_keys=keys, unpack=True)) + assert "vec" not in df.columns + + +def test_fallback_node_ids_subset(graph_backend: BaseGraph) -> None: + lean, full = _views(_populate(graph_backend)) + subset = lean.node_ids()[:2] + + keys = [DEFAULT_ATTR_KEYS.NODE_ID, "label"] + df = lean.filter(node_ids=subset).node_attrs(attr_keys=keys) + + assert df[DEFAULT_ATTR_KEYS.NODE_ID].to_list() == subset + assert df.equals(full.filter(node_ids=subset).node_attrs(attr_keys=keys)) + + +def test_fallback_never_returns_rows_outside_the_view(graph_backend: BaseGraph) -> None: + """The root holds the `solution=False` nodes the view dropped.""" + graph = _populate(graph_backend) + lean, _ = _views(graph) + + assert graph.num_nodes() > lean.num_nodes() + + df = lean.node_attrs(attr_keys=[DEFAULT_ATTR_KEYS.NODE_ID, "label"]) + + assert df.height == lean.num_nodes() + assert set(df[DEFAULT_ATTR_KEYS.NODE_ID].to_list()) == set(lean.node_ids()) + + +def test_fallback_row_order_matches_local_columns(graph_backend: BaseGraph) -> None: + """Callers zip columns positionally; the join must not reorder rows.""" + lean, _ = _views(_populate(graph_backend)) + + df = lean.node_attrs(attr_keys=["x", "label"]) + + for x, label in zip(df["x"], df["label"], strict=True): + assert label == f"n{int(x)}" + + +def test_fallback_after_local_update_is_not_stale(graph_backend: BaseGraph) -> None: + lean, _ = _views(_populate(graph_backend)) + node_ids = lean.node_ids() + + lean.update_node_attrs(node_ids=node_ids, attrs={"x": 99.0}) + + df = lean.node_attrs(attr_keys=["x", "label"]) + assert df["x"].to_list() == [99.0] * len(node_ids) + assert df["label"].to_list() == sorted(df["label"].to_list()) + + +def test_fallback_skips_node_removed_from_view(graph_backend: BaseGraph) -> None: + graph = _populate(graph_backend) + lean, _ = _views(graph) + + removed = lean.node_ids()[0] + lean.remove_node_from_view(removed) + + df = lean.node_attrs(attr_keys=[DEFAULT_ATTR_KEYS.NODE_ID, "label"]) + + assert removed not in df[DEFAULT_ATTR_KEYS.NODE_ID].to_list() + assert graph.has_node(removed) + + +def test_attr_keys_none_stays_lean(graph_backend: BaseGraph) -> None: + """`None` means "what this view holds", not "everything the root has".""" + lean, _ = _views(_populate(graph_backend)) + + assert "label" not in lean.node_attrs().columns + + +# --- error behaviour preserved -------------------------------------------- + + +def test_unknown_key_still_raises(graph_backend: BaseGraph) -> None: + lean, _ = _views(_populate(graph_backend)) + + with pytest.raises(KeyError, match="not found"): + lean.node_attrs(attr_keys=["nowhere"]) + + +def test_without_fallback_missing_key_raises_with_hint(graph_backend: BaseGraph) -> None: + graph = _populate(graph_backend) + view = graph.filter(NodeAttr("solution") == True).subgraph(node_attr_keys=LEAN_KEYS) + + with pytest.raises(KeyError, match="root_fallback=True"): + view.node_attrs(attr_keys=["label"]) + + +def test_without_fallback_unknown_key_has_no_hint(graph_backend: BaseGraph) -> None: + graph = _populate(graph_backend) + view = graph.filter(NodeAttr("solution") == True).subgraph(node_attr_keys=LEAN_KEYS) + + with pytest.raises(KeyError) as excinfo: + view.node_attrs(attr_keys=["nowhere"]) + + assert "root_fallback" not in str(excinfo.value) + + +def test_update_on_missing_key_writes_to_root_and_reads_back(graph_backend: BaseGraph) -> None: + """A write to a non-local key already goes to the root (pre-existing + behaviour: `_update_local_node_attrs` skips keys the view does not hold). + Fallback cannot make that stale -- there is no local copy to disagree.""" + graph = _populate(graph_backend) + lean, _ = _views(graph) + node_ids = lean.node_ids() + + lean.update_node_attrs(node_ids=node_ids, attrs={"label": "written"}) + + assert lean.node_attrs(attr_keys=["label"])["label"].to_list() == ["written"] * len(node_ids) + assert graph.filter(node_ids=node_ids).node_attrs(attr_keys=["label"])["label"].to_list() == ["written"] * len( + node_ids + ) + + +def test_node_attr_keys_reports_only_local_keys(graph_backend: BaseGraph) -> None: + lean, _ = _views(_populate(graph_backend)) + + assert sorted(lean.node_attr_keys()) == sorted([*LEAN_KEYS, DEFAULT_ATTR_KEYS.T]) + + +def test_fallback_is_off_by_default(graph_backend: BaseGraph) -> None: + view = _populate(graph_backend).filter(NodeAttr("solution") == True).subgraph(node_attr_keys=LEAN_KEYS) + + assert view._root_fallback is False + + +# --- end to end through GraphArrayView ------------------------------------- + + +def test_graph_array_view_on_lean_view(graph_backend: BaseGraph) -> None: + """The case the feature exists for: render masks the view does not hold.""" + graph = graph_backend + graph.add_node_attr_key("label", dtype=pl.Int64) + graph.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object) + graph.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, pl.Array(pl.Int64, 4)) + graph.add_node_attr_key("solution", dtype=pl.Boolean, default_value=False) + + for i, solution in enumerate([True, False, True]): + mask = Mask( + np.ones((2, 2), dtype=bool), + bbox=np.array([10 * i, 20, 10 * i + 2, 22]), + ) + graph.add_node( + { + DEFAULT_ATTR_KEYS.T: 0, + "label": i + 1, + DEFAULT_ATTR_KEYS.MASK: mask, + DEFAULT_ATTR_KEYS.BBOX: mask.bbox, + "solution": solution, + } + ) + + node_filter = graph.filter(NodeAttr("solution") == True) + # bbox stays local: the view's own R-tree is built from it. + lean = node_filter.subgraph(node_attr_keys=["label", DEFAULT_ATTR_KEYS.BBOX], root_fallback=True) + full = node_filter.subgraph() + + shape = (1, 100, 100) + lean_array = np.asarray(GraphArrayView(graph=lean, shape=shape, attr_key="label")[0]) + full_array = np.asarray(GraphArrayView(graph=full, shape=shape, attr_key="label")[0]) + + assert lean_array.max() > 0 + np.testing.assert_array_equal(lean_array, full_array) + + +# --- regression guard: the excluded column is never materialized ------------ + + +class _ExplodingBlob: + """Unpickles only if something actually reads the column it lives in.""" + + def __setstate__(self, state: dict) -> None: + raise AssertionError("excluded blob column was read while building the lean view") + + def __getstate__(self) -> dict: + return {} + + +def test_lean_view_does_not_read_the_excluded_column() -> None: + from tracksdata.graph import SQLGraph + + graph = SQLGraph(drivername="sqlite", database=":memory:") + graph.add_node_attr_key("x", dtype=pl.Float64) + graph.add_node_attr_key("blob", dtype=pl.Object) + + for i in range(4): + graph.add_node({DEFAULT_ATTR_KEYS.T: i, "x": float(i), "blob": _ExplodingBlob()}) + + # Would raise if `blob` were fetched and unpickled here. + lean = graph.filter().subgraph(node_attr_keys=["x"], root_fallback=True) + + assert "blob" not in lean.node_attr_keys() + # ... and it is genuinely absent from the view's own storage, not just hidden. + assert all("blob" not in payload for payload in lean.rx_graph.nodes()) + + with pytest.raises(AssertionError, match="excluded blob column"): + lean.node_attrs(attr_keys=["blob"]) + + +# --- a view of a lean view --------------------------------------------------- + + +def test_nested_subgraph_inherits_the_lean_key_set(graph_backend: BaseGraph) -> None: + """A child's rx payloads are copied from its parent, so it holds no more + than the parent did -- and must not advertise otherwise.""" + lean, _ = _views(_populate(graph_backend)) + + nested = lean.filter(node_ids=lean.node_ids()[:2]).subgraph() + + assert "label" not in nested.node_attr_keys() + assert nested.node_attrs(attr_keys=["label"])["label"].to_list() == ["n0", "n2"] + + +def test_nested_subgraph_without_fallback_raises(graph_backend: BaseGraph) -> None: + """The silent-defaults case: it must fail loudly, not return schema defaults.""" + graph = _populate(graph_backend) + lean = graph.filter(NodeAttr("solution") == True).subgraph(node_attr_keys=LEAN_KEYS) + + nested = lean.filter(node_ids=lean.node_ids()[:2]).subgraph() + + with pytest.raises(KeyError, match="root_fallback=True"): + nested.node_attrs(attr_keys=["label"]) + + +def test_nested_subgraph_explicit_key_the_parent_lacks(graph_backend: BaseGraph) -> None: + """Naming an excluded key explicitly must not re-claim it locally.""" + lean, _ = _views(_populate(graph_backend)) + + nested = lean.filter(node_ids=lean.node_ids()[:2]).subgraph(node_attr_keys=["label"]) + + assert "label" not in nested.node_attr_keys() + assert nested.node_attrs(attr_keys=["label"])["label"].to_list() == ["n0", "n2"] + + +def test_nested_subgraph_of_full_view_is_unclamped(graph_backend: BaseGraph) -> None: + """A parent holding everything clamps nothing.""" + graph = _populate(graph_backend) + full = graph.filter(NodeAttr("solution") == True).subgraph() + + nested = full.filter(node_ids=full.node_ids()[:2]).subgraph(node_attr_keys=["label"]) + + assert "label" in nested.node_attr_keys() + assert nested.node_attrs(attr_keys=["label"])["label"].to_list() == ["n0", "n2"] + + +def test_nested_subgraph_key_on_neither_parent_nor_root(graph_backend: BaseGraph) -> None: + """A key that exists nowhere raises plainly -- pointing at `root_fallback` + would suggest a flag that cannot help.""" + graph = _populate(graph_backend) + + for fallback in (True, False): + lean = graph.filter(NodeAttr("solution") == True).subgraph( + node_attr_keys=LEAN_KEYS, + root_fallback=fallback, + ) + nested = lean.filter(node_ids=lean.node_ids()[:2]).subgraph() + + with pytest.raises(KeyError) as excinfo: + nested.node_attrs(attr_keys=["nowhere"]) + + assert "nowhere" in str(excinfo.value) + assert "root_fallback" not in str(excinfo.value) diff --git a/src/tracksdata/graph/filters/_base_filter.py b/src/tracksdata/graph/filters/_base_filter.py index e6ea6b77..acafb4c2 100644 --- a/src/tracksdata/graph/filters/_base_filter.py +++ b/src/tracksdata/graph/filters/_base_filter.py @@ -51,6 +51,7 @@ def subgraph( edge_attr_keys: list[str] | None = None, *, mode: "ViewMode | None" = None, + root_fallback: bool = False, ) -> "GraphView": """ Get a subgraph of the graph resulting from the filter. @@ -59,6 +60,11 @@ def subgraph( How the resulting view relates to its root graph (write-through only, or live-updating). Defaults to ``ViewMode.WRITE_THROUGH`` when None. + root_fallback : bool + When True, reading a node attribute key the resulting view does not + hold is served from the root graph instead of raising. Off by + default: the fetch goes to the root's storage, which for a SQL root + means a query per read. """ @abc.abstractmethod diff --git a/src/tracksdata/graph/filters/_indexed_filter.py b/src/tracksdata/graph/filters/_indexed_filter.py index 8efea0d9..521b5cfe 100644 --- a/src/tracksdata/graph/filters/_indexed_filter.py +++ b/src/tracksdata/graph/filters/_indexed_filter.py @@ -16,6 +16,30 @@ from tracksdata.graph._graph_view import GraphView, ViewMode +def _clamp_to_parent( + requested: Sequence[str] | str | None, + parent_keys: list[str] | None, +) -> Sequence[str] | str | None: + """ + Limit the keys a child view claims to hold to the ones its parent held. + + ``parent_keys is None`` means the parent held everything its own root has, + so there is nothing to clamp. Otherwise requested keys the parent did not + hold are dropped rather than claimed: the child's rx payloads are copied + from the parent, so a key the parent lacked has no values here either. + Claiming it makes reads return schema defaults instead of raising or + falling back to the root. + """ + if parent_keys is None: + return requested + if requested is None: + return parent_keys + if isinstance(requested, str): + requested = [requested] + held = set(parent_keys) + return [k for k in requested if k in held] + + class IndexRXFilter(RXFilter): _graph: "GraphView | IndexedRXGraph" @@ -54,6 +78,7 @@ def subgraph( edge_attr_keys: Sequence[str] | str | None = None, *, mode: "ViewMode | None" = None, + root_fallback: bool = False, ) -> "GraphView": from tracksdata.graph._graph_view import GraphView, ViewMode @@ -68,7 +93,15 @@ def subgraph( root = self._graph if hasattr(self._graph, "_root"): + # A view of a partial view is flattened onto the shared root, but it + # holds only what its parent held -- the rx payloads above were + # copied from the parent, not re-read from the root. Both the key + # lists and the fallback flag have to be inherited, or the child + # advertises the root's columns while holding no values for them. root = self._graph._root + node_attr_keys = _clamp_to_parent(node_attr_keys, self._graph._node_attr_keys) + edge_attr_keys = _clamp_to_parent(edge_attr_keys, self._graph._edge_attr_keys) + root_fallback = root_fallback or self._graph._root_fallback graph_view = GraphView( rx_graph, @@ -77,6 +110,7 @@ def subgraph( mode=mode if mode is not None else ViewMode.WRITE_THROUGH, node_attr_keys=node_attr_keys, edge_attr_keys=edge_attr_keys, + root_fallback=root_fallback, ) return graph_view diff --git a/src/tracksdata/metrics/_traccuracy.py b/src/tracksdata/metrics/_traccuracy.py index 05452bc1..15e9e8ee 100644 --- a/src/tracksdata/metrics/_traccuracy.py +++ b/src/tracksdata/metrics/_traccuracy.py @@ -26,6 +26,8 @@ def to_traccuracy_graph( The graph to convert. array_view_kwargs : dict[str, Any] | None Additional keyword arguments to pass to the `GraphArrayView` constructor used to create the segmentation. + Do not pass `downscale`: metrics need exact segmentation ids, and a downscaled + view leaves the value of a voxel covered by several objects unspecified. location_keys : list[str] | None The keys of the location attributes to use for the segmentation. If None, the location keys are inferred from the intersection of the graph node attributes and diff --git a/src/tracksdata/nodes/_mask.py b/src/tracksdata/nodes/_mask.py index 2c942139..edb160c3 100644 --- a/src/tracksdata/nodes/_mask.py +++ b/src/tracksdata/nodes/_mask.py @@ -37,6 +37,20 @@ def _nd_sphere( raise ValueError(f"Spherical is only implemented for 2D and 3D, got ndim={ndim}") +def _as_axis_vector(value: ArrayLike | int, ndim: int, name: str) -> NDArray[np.int64]: + """Normalize a scalar or per-axis integer sequence to an int64 vector of length `ndim`.""" + if not np.issubdtype(np.asarray(value).dtype, np.integer): + raise ValueError(f"`{name}` must be integer, got {value!r}") + + if isinstance(value, int | np.integer): + return np.full(ndim, value, dtype=np.int64) + + vector = np.asarray(value, dtype=np.int64).reshape(-1) + if len(vector) != ndim: + raise ValueError(f"`{name}` must have length {ndim}, got {len(vector)}") + return vector + + class Mask: """ Object used to store an individual segmentation mask of a single instance (object) @@ -154,7 +168,7 @@ def mask_indices( tuple[NDArray[np.integer], ...] The indices of the pixels that are part of the object. """ - if isinstance(offset, int): + if isinstance(offset, int | np.integer): offset = np.full(self._mask.ndim, offset) indices = list(np.nonzero(self._mask)) @@ -169,6 +183,7 @@ def paint_buffer( buffer: np.ndarray, value: int | float, offset: NDArray[np.integer] | int = 0, + downscale: NDArray[np.integer] | int | None = None, ) -> None: """ Paint object into a buffer. @@ -181,14 +196,36 @@ def paint_buffer( The value to paint the object. offset : NDArray[np.integer] | int, optional The offset to add to the indices, should be used with bounding box information. - """ + downscale : NDArray[np.integer] | int | None, optional + Per-axis integer downscaling factors. When given, `buffer` is assumed to be a + downscaled volume and the object is painted by nearest-neighbor sampling: + output voxel `o` takes the value of full-resolution coordinate `o * downscale`. + + Sampling is anchored to the global output grid rather than to this mask's + bounding box, so neighboring objects always sample the same phase. + + Objects thinner than `downscale` would vanish under plain sampling, so they + instead get a single voxel painted at their bounding box center. They stay + visible and selectable, but their rendered shape and size are meaningless. + The exception is an object whose bounding box reaches past `buffer` and whose + center therefore falls outside it: such an object is dropped even though it + overlaps a sampled voxel's footprint. + + Never interpolates: averaging label values would invent values that belong to + no object. When several full-resolution voxels map to one output voxel the last + write wins, so an output voxel covered by more than one object is unspecified. + """ + if downscale is not None and not np.all(np.asarray(downscale) == 1): + self._paint_buffer_downscaled(buffer, value, offset, downscale) + return + ndim = self._mask.ndim bbox = self._bbox shape = buffer.shape - if isinstance(offset, int): - starts = [int(bbox[i]) + offset for i in range(ndim)] - stops = [int(bbox[i + ndim]) + offset for i in range(ndim)] + if isinstance(offset, int | np.integer): + starts = [int(bbox[i]) + int(offset) for i in range(ndim)] + stops = [int(bbox[i + ndim]) + int(offset) for i in range(ndim)] else: starts = [int(bbox[i]) + int(offset[i]) for i in range(ndim)] stops = [int(bbox[i + ndim]) + int(offset[i]) for i in range(ndim)] @@ -211,6 +248,68 @@ def paint_buffer( window = tuple(slice(clipped_start[i], clipped_stop[i]) for i in range(ndim)) buffer[window][self._mask[mask_slicing]] = value + def _paint_buffer_downscaled( + self, + buffer: np.ndarray, + value: int | float, + offset: NDArray[np.integer] | int, + downscale: NDArray[np.integer] | int, + ) -> None: + """ + Paint object into a downscaled buffer by nearest-neighbor sampling. + + See `paint_buffer` for the sampling convention and its consequences. + """ + ndim = self._mask.ndim + factors = _as_axis_vector(downscale, ndim, "downscale") + if np.any(factors < 1): + raise ValueError(f"`downscale` factors must be >= 1, got {factors.tolist()}") + + offset = _as_axis_vector(offset, ndim, "offset") + bbox = self._bbox + shape = buffer.shape + + # Plain Python integers throughout, as in the full-resolution path above: these + # arrays have one entry per axis, so numpy bookkeeping costs more than the paint. + window: list[slice] = [] + mask_slicing: list[slice] = [] + sampled_empty = False + + for i in range(ndim): + factor = int(factors[i]) + start = int(bbox[i]) + int(offset[i]) + stop = int(bbox[i + ndim]) + int(offset[i]) + + # Output voxel `o` samples full-resolution coordinate `o * factor`, anchored to + # the global grid, so `o` ranges over [ceil(start / f), ceil(stop / f)). Clip to + # the buffer by trimming the sampled window rather than filtering afterwards. + lo = max(-(-start // factor), 0) + hi = min(-(-stop // factor), shape[i]) + if hi <= lo: + sampled_empty = True + break + + # the sampling phase, plus whatever the clipping above skipped + local = lo * factor - start + window.append(slice(lo, hi)) + mask_slicing.append(slice(local, local + (hi - lo) * factor, factor)) + + if not sampled_empty: + sampled = self._mask[tuple(mask_slicing)] + if sampled.any(): + # boolean-mask assignment into a sliced view, as in the fast path above; + # `np.nonzero` plus fancy indexing is measurably slower for the same voxels + buffer[tuple(window)][sampled] = value + return + + # Object fell between samples: keep it visible as a single voxel so that it stays + # selectable. `//` floors towards -inf, so out-of-bounds stays out of bounds. + center = tuple( + (int(bbox[i]) + int(bbox[i + ndim]) + 2 * int(offset[i])) // 2 // int(factors[i]) for i in range(ndim) + ) + if all(0 <= c < s for c, s in zip(center, shape, strict=True)): + buffer[center] = value + def iou(self, other: "Mask") -> float: """ Compute the Intersection over Union (IoU) between two masks diff --git a/src/tracksdata/nodes/_test/test_mask.py b/src/tracksdata/nodes/_test/test_mask.py index 38ff8fdb..9f76cce7 100644 --- a/src/tracksdata/nodes/_test/test_mask.py +++ b/src/tracksdata/nodes/_test/test_mask.py @@ -829,3 +829,162 @@ def test_mask_sub_matches_canvas_subtraction_3d(seed: int) -> None: result.paint_buffer(painted, True) np.testing.assert_array_equal(painted, _subtract_via_canvas(mask1, mask2, canvas_shape)) + + +def _downscaled_shape(shape: tuple[int, ...], factors: tuple[int, ...]) -> tuple[int, ...]: + return tuple(-(-s // f) for s, f in zip(shape, factors, strict=True)) + + +def test_paint_buffer_downscale_none_and_one_are_identical() -> None: + """`downscale` of None, 1 and all-ones must all take the full-resolution fast path.""" + mask = Mask(np.ones((3, 3), dtype=bool), [1, 1, 4, 4]) + + buffers = [np.zeros((10, 10), dtype=np.uint32) for _ in range(3)] + mask.paint_buffer(buffers[0], value=5) + mask.paint_buffer(buffers[1], value=5, downscale=1) + mask.paint_buffer(buffers[2], value=5, downscale=(1, 1)) + + assert np.array_equal(buffers[0], buffers[1]) + assert np.array_equal(buffers[0], buffers[2]) + + +def test_paint_buffer_downscale_matches_strided_reference() -> None: + """Objects wider than the factor must match a global-grid strided subsample.""" + dense = np.zeros((40, 40), dtype=np.uint32) + dense[5:15, 6:16] = 7 + dense[21:33, 23:35] = 9 + + factors = (4, 4) + buffer = np.zeros(_downscaled_shape(dense.shape, factors), dtype=np.uint32) + for value, slicing in ((7, (slice(5, 15), slice(6, 16))), (9, (slice(21, 33), slice(23, 35)))): + bbox = [slicing[0].start, slicing[1].start, slicing[0].stop, slicing[1].stop] + Mask(dense[slicing] == value, bbox).paint_buffer(buffer, value, downscale=factors) + + np.testing.assert_array_equal(buffer, dense[:: factors[0], :: factors[1]]) + + +def test_paint_buffer_downscale_anisotropic() -> None: + """Per-axis factors must be applied independently.""" + dense = np.zeros((8, 40, 40), dtype=np.uint32) + dense[2:6, 8:20, 8:20] = 4 + + factors = (1, 4, 4) + buffer = np.zeros(_downscaled_shape(dense.shape, factors), dtype=np.uint32) + Mask(dense[2:6, 8:20, 8:20] == 4, [2, 8, 8, 6, 20, 20]).paint_buffer(buffer, 4, downscale=factors) + + assert buffer.shape == (8, 10, 10) + np.testing.assert_array_equal(buffer, dense[:: factors[0], :: factors[1], :: factors[2]]) + + +def test_paint_buffer_downscale_sampled_extent() -> None: + """The painted extent must span [ceil(start / f), ceil(stop / f)).""" + buffer = np.zeros((4,), dtype=np.uint32) + Mask(np.ones((7,), dtype=bool), [2, 9]).paint_buffer(buffer, 3, downscale=4) + + # start=2, stop=9, f=4 -> output voxels 1 and 2 sample coordinates 4 and 8 + np.testing.assert_array_equal(buffer, [0, 3, 3, 0]) + + +def test_paint_buffer_downscale_global_anchoring() -> None: + """ + Sampling is anchored to the output grid, not to each mask's bounding box. + + The mask below is only set at the voxels the output grid samples, so an + implementation that starts sampling at the bounding box reads the wrong voxels and + finds nothing. The bbox start must not be a multiple of the factor for this to bite. + """ + # bbox [2, 10) with f=4 samples full-resolution 4 and 8, i.e. mask offsets 2 and 6 + mask = np.zeros((8,), dtype=bool) + mask[[2, 6]] = True + + buffer = np.zeros((4,), dtype=np.uint32) + Mask(mask, [2, 10]).paint_buffer(buffer, 1, downscale=4) + + # anchored to the bbox instead, offsets 0 and 4 are read, both False + np.testing.assert_array_equal(buffer, [0, 1, 1, 0]) + + +def test_paint_buffer_downscale_phase_selects_correct_voxels() -> None: + """Only the globally sampled voxels are read, not merely *some* voxel per output.""" + # inverse of the test above: the sampled offsets are the only False ones + mask = np.ones((8,), dtype=bool) + mask[[2, 6]] = False + + buffer = np.zeros((4,), dtype=np.uint32) + Mask(mask, [2, 10]).paint_buffer(buffer, 1, downscale=4) + + # nothing sampled, so only the fallback voxel at (2 + 10) // 2 // 4 == 1 is painted + np.testing.assert_array_equal(buffer, [0, 1, 0, 0]) + + +def test_paint_buffer_downscale_fallback_keeps_small_object() -> None: + """An object falling between samples must survive as a single voxel.""" + buffer = np.zeros((4, 4), dtype=np.uint32) + Mask(np.ones((1, 1), dtype=bool), [7, 9, 8, 10]).paint_buffer(buffer, 11, downscale=4) + + # plain striding samples coordinates 4 and 8 on each axis, missing (7, 9) entirely + np.testing.assert_array_equal(np.argwhere(buffer), [[1, 2]]) + + +def test_paint_buffer_downscale_fallback_all_false_mask() -> None: + """An all-False mask still marks its location, consistently with the fallback.""" + buffer = np.zeros((4, 4), dtype=np.uint32) + Mask(np.zeros((2, 2), dtype=bool), [4, 4, 6, 6]).paint_buffer(buffer, 3, downscale=4) + + np.testing.assert_array_equal(np.argwhere(buffer), [[1, 1]]) + + +def test_paint_buffer_downscale_fallback_out_of_bounds() -> None: + """A fallback voxel outside the buffer must be dropped, not wrapped.""" + buffer = np.zeros((4, 4), dtype=np.uint32) + Mask(np.ones((1, 1), dtype=bool), [40, 40, 41, 41]).paint_buffer(buffer, 3, downscale=4) + + assert not buffer.any() + + +def test_paint_buffer_downscale_negative_offset() -> None: + """Negative coordinates must floor away from the buffer, never onto index 0.""" + buffer = np.zeros((4, 4), dtype=np.uint32) + mask = Mask(np.ones((2, 2), dtype=bool), [0, 0, 2, 2]) + + mask.paint_buffer(buffer, 6, offset=np.int64(-3), downscale=4) + + assert not buffer.any() + + +def test_paint_buffer_downscale_clips_positive_overflow() -> None: + """A bbox extending past the buffer must be clipped to the buffer.""" + buffer = np.zeros((4, 4), dtype=np.uint32) + Mask(np.ones((8, 8), dtype=bool), [12, 12, 20, 20]).paint_buffer(buffer, 6, downscale=4) + + np.testing.assert_array_equal(np.argwhere(buffer), [[3, 3]]) + + +def test_paint_buffer_downscale_numpy_scalar_offset() -> None: + """A numpy scalar offset must behave like a Python int.""" + expected = np.zeros((10, 10), dtype=np.uint32) + actual = np.zeros((10, 10), dtype=np.uint32) + mask = Mask(np.ones((2, 2), dtype=bool), [1, 1, 3, 3]) + + mask.paint_buffer(expected, 5, offset=2) + mask.paint_buffer(actual, 5, offset=np.int64(2)) + + np.testing.assert_array_equal(actual, expected) + + +@pytest.mark.parametrize("downscale", [0, -1, 1.5, (2, 2, 2)]) +def test_paint_buffer_downscale_invalid(downscale: int | float | tuple[int, ...]) -> None: + """Invalid factors must be rejected rather than silently mangling coordinates.""" + mask = Mask(np.ones((2, 2), dtype=bool), [0, 0, 2, 2]) + + with pytest.raises(ValueError, match="downscale"): + mask.paint_buffer(np.zeros((4, 4), dtype=np.uint32), 1, downscale=downscale) + + +@pytest.mark.parametrize("offset", [1.5, np.float64(1.5), [1.5, 2.5]]) +def test_paint_buffer_downscale_non_integer_offset(offset: float | list[float]) -> None: + """A non-integer offset must be rejected rather than silently truncated.""" + mask = Mask(np.ones((2, 2), dtype=bool), [0, 0, 2, 2]) + + with pytest.raises(ValueError, match="offset"): + mask.paint_buffer(np.zeros((6, 6), dtype=np.uint32), 5, offset=offset, downscale=2)