Skip to content
Draft
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
150 changes: 150 additions & 0 deletions benchmarks/graph_array_downscale.py
Original file line number Diff line number Diff line change
@@ -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)"
)
6 changes: 4 additions & 2 deletions docs/examples/basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
24 changes: 23 additions & 1 deletion docs/faq.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
133 changes: 125 additions & 8 deletions src/tracksdata/array/_graph_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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__(
Expand All @@ -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()}'")
Expand All @@ -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):
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -346,14 +450,21 @@ 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],
)

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."""
Expand All @@ -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:
Expand Down
Loading
Loading