Skip to content
Open
Show file tree
Hide file tree
Changes from 46 commits
Commits
Show all changes
47 commits
Select commit Hold shift + click to select a range
29177f3
added private metadata machinery
yfukai Feb 17, 2026
d8292f1
before adding private
yfukai Feb 17, 2026
cff5898
added private metadata view
yfukai Feb 17, 2026
68b01d4
renamed func
yfukai Feb 17, 2026
1ae2426
implementation of saving and loading dtypes as metadata
yfukai Feb 17, 2026
c50a07b
lint
yfukai Feb 17, 2026
e9bf28f
restricted dtype metadata to sqlgraph
yfukai Feb 18, 2026
9aa9c3a
udpated serialization strategies
yfukai Feb 18, 2026
7e61ac3
solved failing tests
yfukai Feb 18, 2026
e5968bf
added test for shape-less pl.Array (xfail)
yfukai Feb 18, 2026
b4acde3
working
yfukai Feb 19, 2026
cc55976
simplified code
yfukai Feb 19, 2026
e76d8e5
initial try
yfukai Feb 19, 2026
7bec369
saving private metadata
yfukai Feb 20, 2026
852f717
rustworkx reviewed
yfukai Feb 26, 2026
4c151bb
Merge branch 'from_other_roundtrip' into struct_attr
yfukai Feb 26, 2026
4af9904
working with clean code?
yfukai Feb 26, 2026
19055ab
Merge branch 'main' into struct_attr
JoOkuma Feb 27, 2026
6c69e76
updated impl
yfukai Apr 10, 2026
d9bee26
removed codex config wrongly added
yfukai Apr 10, 2026
cc0beb4
issue fixes
yfukai Apr 14, 2026
007d4c7
Merge branch 'main' into struct_attr
yfukai Apr 14, 2026
ffea2ec
rolled back unncessary change
yfukai Apr 14, 2026
0ad6c60
Merge remote-tracking branch 'upstream/main' into struct_attr
yfukai May 28, 2026
5e7331f
additional comments
yfukai May 28, 2026
7e07801
Merge branch 'main' into struct_attr
JoOkuma Jun 1, 2026
37f8fc9
Fix lint: remove whitespace from blank lines
JoOkuma Jun 1, 2026
1842c55
fixes
yfukai Jun 4, 2026
db35287
refactor aligning main
yfukai Jun 4, 2026
3759d9c
Restore scratch-table machinery and tests from main
yfukai Jun 5, 2026
f5b7cc0
Merge branch 'main' of https://github.com/royerlab/tracksdata into st…
yfukai Jun 8, 2026
0d76262
ignored the devcontaienr
yfukai Jun 8, 2026
87707d0
bugfix
yfukai Jun 8, 2026
9ba99e1
Store Mask as a struct attribute instead of pickled pl.Object
yfukai Jun 10, 2026
e7f99af
Merge upstream/main into mask_struct_attr
yfukai Jun 17, 2026
7d02b80
Store binary attributes as raw BLOB in SQL instead of pickling
yfukai Jun 17, 2026
00f41eb
Clean up mask struct-attribute branch for review
yfukai Jun 17, 2026
07b68c6
Document why as_mask imports are function-local
yfukai Jun 17, 2026
1a6ad08
updating name and adding comments
yfukai Jun 22, 2026
a21dd31
adding test
yfukai Jun 22, 2026
81a473e
Merge remote-tracking branch 'upstream/main' into mask_struct_attr
yfukai Jul 27, 2026
0821359
Replace Mask class with a TypedDict and module-level functions
yfukai Jul 27, 2026
61c4ba2
Make Mask a frozen dataclass instead of a TypedDict
yfukai Jul 27, 2026
1231a9e
Fix CI lint against newer ruff versions
yfukai Jul 27, 2026
607d258
Pin witty away from 0.3.2 to fix the C extension build
yfukai Jul 27, 2026
594d3bf
Merge branch 'main' into mask_struct_attr
JoOkuma Jul 29, 2026
e84def5
Update src/tracksdata/graph/_sql_graph.py
yfukai Jul 30, 2026
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
10 changes: 9 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -60,10 +60,15 @@ dependencies = [
]

[project.optional-dependencies]
spatial = ["spatial-graph"]
# witty 0.3.2 changed `source_files` from "files to watch for cache
# invalidation" to "files to compile and link". spatial-graph still lists its
# .h headers there, so building its C extensions fails with
# `UnknownFileType: unknown file type '.h'`. Allow a future fixed release.
spatial = ["spatial-graph", "witty!=0.3.2"]
motile = ["motile"]
test = [
"spatial-graph",
"witty!=0.3.2",
"motile",
"pytest>=7.0",
"pytest-cov",
Expand Down Expand Up @@ -129,6 +134,9 @@ select = [

# https://docs.astral.sh/ruff/formatter/
[tool.ruff.format]
# ruff >= 0.16 formats markdown code blocks, which rewrites the mkdocs
# snippet directives (`--8<-- "..."`) in docs/ into invalid syntax.
exclude = ["*.md"]
docstring-code-format = true
skip-magic-trailing-comma = false # default is false

Expand Down
19 changes: 12 additions & 7 deletions src/tracksdata/array/_graph_array.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
from collections.abc import Sequence
from copy import copy
from typing import TYPE_CHECKING, Any
from typing import Any

import numpy as np

Expand All @@ -11,9 +11,6 @@
from tracksdata.options import get_options
from tracksdata.utils._dtypes import polars_dtype_to_numpy_dtype

if TYPE_CHECKING:
from tracksdata.nodes._mask import Mask


def _validate_shape(
shape: tuple[int, ...] | None,
Expand Down Expand Up @@ -346,14 +343,19 @@ def _fill_array(self, time: int, volume_slicing: Sequence[slice], buffer: np.nda
np.ndarray
The filled buffer.
"""
# Local import: avoids the graph <-> nodes package import cycle (importing
# tracksdata.nodes re-enters the partially-initialized graph package).
from tracksdata.nodes._mask import mask_paint_buffer, masks_from_column

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)
masks = masks_from_column(df[DEFAULT_ATTR_KEYS.MASK])

for mask, value in zip(masks, df[self._attr_key], strict=True):
mask_paint_buffer(mask, buffer, value, offset=self._offset)

def _offset_as_array(self, ndim: int) -> np.ndarray:
"""Normalize `offset` to a vector for each spatial axis."""
Expand Down Expand Up @@ -489,4 +491,7 @@ def _mask_changed(old_attr: dict, new_attr: dict) -> bool:
return False
elif old_mask is None or new_mask is None:
return True

# `Mask` compares by value; struct values are dicts of scalars plus the
# compressed blob, which also compare directly.
return old_mask != new_mask
18 changes: 9 additions & 9 deletions src/tracksdata/array/_test/test_graph_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,7 @@ def test_graph_array_view_getitem_with_nodes(graph_backend: BaseGraph) -> None:

# Create a mask
mask_data = np.array([[True, True], [True, False]], dtype=bool)
mask = Mask(mask_data, bbox=np.array([10, 20, 12, 22])) # y_min, x_min, y_max, x_max
mask = Mask(bbox=np.array([10, 20, 12, 22]), mask=mask_data) # y_min, x_min, y_max, x_max

# Add a node with mask and label
graph_backend.add_node(
Expand Down Expand Up @@ -136,10 +136,10 @@ def test_graph_array_view_getitem_multiple_nodes(graph_backend: BaseGraph) -> No

# Create two masks at different locations
mask1_data = np.array([[True, True]], dtype=bool)
mask1 = Mask(mask1_data, bbox=np.array([10, 20, 11, 22]))
mask1 = Mask(bbox=np.array([10, 20, 11, 22]), mask=mask1_data)

mask2_data = np.array([[True]], dtype=bool)
mask2 = Mask(mask2_data, bbox=np.array([30, 40, 31, 41]))
mask2 = Mask(bbox=np.array([30, 40, 31, 41]), mask=mask2_data)

# Add nodes with different labels
graph_backend.add_node(
Expand Down Expand Up @@ -191,7 +191,7 @@ def test_graph_array_view_getitem_boolean_dtype(graph_backend: BaseGraph) -> Non

# Create a mask
mask_data = np.array([[True]], dtype=bool)
mask = Mask(mask_data, bbox=np.array([10, 20, 11, 21]))
mask = Mask(bbox=np.array([10, 20, 11, 21]), mask=mask_data)

# Add a node with boolean attribute
graph_backend.add_node(
Expand Down Expand Up @@ -228,7 +228,7 @@ def test_graph_array_view_dtype_inference(graph_backend: BaseGraph) -> None:

# Create a mask
mask_data = np.array([[True]], dtype=bool)
mask = Mask(mask_data, bbox=np.array([10, 20, 11, 21]))
mask = Mask(bbox=np.array([10, 20, 11, 21]), mask=mask_data)

# Add a node with float attribute
graph_backend.add_node(
Expand Down Expand Up @@ -407,7 +407,7 @@ def test_graph_array_raise_error_on_non_scalar_attr_key(graph_backend: BaseGraph
{
DEFAULT_ATTR_KEYS.T: 0,
"label": np.array([1, 2]), # Non-scalar value
DEFAULT_ATTR_KEYS.MASK: Mask(np.array([[True]], dtype=bool), bbox=np.array([0, 0, 1, 1])),
DEFAULT_ATTR_KEYS.MASK: Mask(bbox=np.array([0, 0, 1, 1]), mask=np.array([[True]], dtype=bool)),
}
)

Expand All @@ -422,7 +422,7 @@ def _add_graph_array_node_attrs(graph_backend: BaseGraph) -> None:


def _make_square_mask(y: int, x: int, size: int = 2) -> Mask:
return Mask(np.ones((size, size), dtype=bool), bbox=np.array([y, x, y + size, x + size]))
return Mask(bbox=np.array([y, x, y + size, x + size]), mask=np.ones((size, size), dtype=bool))


def test_graph_array_view_invalidates_only_affected_chunk_on_add(graph_backend: BaseGraph) -> None:
Expand Down Expand Up @@ -695,7 +695,7 @@ def test_graph_array_view_invalidates_once_when_mask_changes_but_bbox_unchanged(
np.testing.assert_array_equal(array_view._cache._store[0].ready, np.ones((2, 2), dtype=bool))

# Same bbox [1, 1, 3, 3], but only the diagonal pixels are set.
new_mask = Mask(np.array([[True, False], [False, True]], dtype=bool), bbox=np.array([1, 1, 3, 3]))
new_mask = Mask(bbox=np.array([1, 1, 3, 3]), mask=np.array([[True, False], [False, True]], dtype=bool))
mock_invalidate = MagicMock(wraps=array_view._invalidate_bbox)
with patch.object(array_view, "_invalidate_bbox", mock_invalidate):
graph_backend.update_node_attrs(
Expand Down Expand Up @@ -737,7 +737,7 @@ def test_graph_array_view_no_invalidation_when_mask_unchanged(graph_backend: Bas
np.testing.assert_array_equal(array_view._cache._store[0].ready, np.ones((2, 2), dtype=bool))

# A fresh Mask object with identical bbox and pixels: no rendered change.
same_mask = Mask(np.ones((2, 2), dtype=bool), bbox=np.array([1, 1, 3, 3]))
same_mask = Mask(bbox=np.array([1, 1, 3, 3]), mask=np.ones((2, 2), dtype=bool))
mock_invalidate = MagicMock(wraps=array_view._invalidate_bbox)
with patch.object(array_view, "_invalidate_bbox", mock_invalidate):
graph_backend.update_node_attrs(
Expand Down
11 changes: 9 additions & 2 deletions src/tracksdata/edges/_iou_edges.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,13 @@
from typing import Any

from tracksdata.constants import DEFAULT_ATTR_KEYS
from tracksdata.edges._generic_edges import GenericFuncEdgeAttrs
from tracksdata.nodes._mask import Mask
from tracksdata.nodes._mask import Mask, _decode_mask, mask_iou


def _mask_iou(source_mask: "Mask | dict[str, Any]", target_mask: "Mask | dict[str, Any]") -> float:
"""IoU between two mask attribute values (struct dicts or `Mask` values)."""
return mask_iou(_decode_mask(source_mask), _decode_mask(target_mask))


class IoUEdgeAttr(GenericFuncEdgeAttrs):
Expand All @@ -22,7 +29,7 @@ def __init__(
mask_key: str = DEFAULT_ATTR_KEYS.MASK,
):
super().__init__(
func=Mask.iou,
func=_mask_iou,
attr_keys=mask_key,
output_key=output_key,
)
23 changes: 12 additions & 11 deletions src/tracksdata/edges/_test/test_iou_edges.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from tracksdata.constants import DEFAULT_ATTR_KEYS
from tracksdata.edges import IoUEdgeAttr
from tracksdata.edges._iou_edges import _mask_iou
from tracksdata.graph import RustWorkXGraph
from tracksdata.nodes import Mask
from tracksdata.options import get_options, options_context
Expand All @@ -15,7 +16,7 @@ def test_iou_edges_init_default() -> None:

assert operator.output_key == "iou_score"
assert operator.attr_keys == DEFAULT_ATTR_KEYS.MASK
assert operator.func == Mask.iou
assert operator.func == _mask_iou


def test_iou_edges_init_custom() -> None:
Expand All @@ -24,7 +25,7 @@ def test_iou_edges_init_custom() -> None:

assert operator.output_key == "custom_iou"
assert operator.attr_keys == "custom_mask"
assert operator.func == Mask.iou
assert operator.func == _mask_iou


@pytest.mark.parametrize("n_workers", [1, 2])
Expand All @@ -38,13 +39,13 @@ def test_iou_edges_add_weights(n_workers: int) -> None:

# Create test masks
mask1_data = np.array([[True, True], [True, False]], dtype=bool)
mask1 = Mask(mask1_data, bbox=np.array([0, 0, 2, 2]))
mask1 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask1_data)

mask2_data = np.array([[True, False], [False, False]], dtype=bool)
mask2 = Mask(mask2_data, bbox=np.array([0, 0, 2, 2]))
mask2 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask2_data)

mask3_data = np.array([[True, True], [True, True]], dtype=bool)
mask3 = Mask(mask3_data, bbox=np.array([0, 0, 2, 2]))
mask3 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask3_data)

# Add nodes with masks
node1 = graph.add_node({DEFAULT_ATTR_KEYS.T: 0, DEFAULT_ATTR_KEYS.MASK: mask1})
Expand Down Expand Up @@ -88,10 +89,10 @@ def test_iou_edges_no_overlap() -> None:

# Create non-overlapping masks
mask1_data = np.array([[True, True], [False, False]], dtype=bool)
mask1 = Mask(mask1_data, bbox=np.array([0, 0, 2, 2]))
mask1 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask1_data)

mask2_data = np.array([[False, False], [True, True]], dtype=bool)
mask2 = Mask(mask2_data, bbox=np.array([0, 0, 2, 2]))
mask2 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask2_data)

# Add nodes with masks
node1 = graph.add_node({DEFAULT_ATTR_KEYS.T: 0, DEFAULT_ATTR_KEYS.MASK: mask1})
Expand Down Expand Up @@ -127,8 +128,8 @@ def test_iou_edges_perfect_overlap() -> None:

# Create identical masks
mask_data = np.array([[True, True], [True, False]], dtype=bool)
mask1 = Mask(mask_data, bbox=np.array([0, 0, 2, 2]))
mask2 = Mask(mask_data, bbox=np.array([0, 0, 2, 2]))
mask1 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask_data)
mask2 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask_data)

# Add nodes with masks
node1 = graph.add_node({DEFAULT_ATTR_KEYS.T: 0, DEFAULT_ATTR_KEYS.MASK: mask1})
Expand Down Expand Up @@ -163,10 +164,10 @@ def test_iou_edges_custom_mask_key() -> None:

# Create test masks
mask1_data = np.array([[True, True], [True, True]], dtype=bool)
mask1 = Mask(mask1_data, bbox=np.array([0, 0, 2, 2]))
mask1 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask1_data)

mask2_data = np.array([[True, True], [False, False]], dtype=bool)
mask2 = Mask(mask2_data, bbox=np.array([0, 0, 2, 2]))
mask2 = Mask(bbox=np.array([0, 0, 2, 2]), mask=mask2_data)

# Add nodes with custom mask key
node1 = graph.add_node({DEFAULT_ATTR_KEYS.T: 0, "custom_mask": mask1})
Expand Down
8 changes: 4 additions & 4 deletions src/tracksdata/functional/_test/test_division.py
Original file line number Diff line number Diff line change
Expand Up @@ -635,10 +635,10 @@ def _make_graph_with_mask() -> tuple[td.graph.RustWorkXGraph, dict[str, int]]:
g = td.graph.RustWorkXGraph()
g.add_node_attr_key(DEFAULT_ATTR_KEYS.MASK, pl.Object)

mask_p = Mask(np.ones((4, 4), dtype=bool), bbox=np.array([0, 0, 4, 4]))
mask_d = Mask(np.ones((4, 4), dtype=bool), bbox=np.array([0, 0, 4, 4]))
mask_c1 = Mask(np.ones((2, 2), dtype=bool), bbox=np.array([0, 0, 2, 2]))
mask_c2 = Mask(np.ones((2, 2), dtype=bool), bbox=np.array([2, 2, 4, 4]))
mask_p = Mask(bbox=np.array([0, 0, 4, 4]), mask=np.ones((4, 4), dtype=bool))
mask_d = Mask(bbox=np.array([0, 0, 4, 4]), mask=np.ones((4, 4), dtype=bool))
mask_c1 = Mask(bbox=np.array([0, 0, 2, 2]), mask=np.ones((2, 2), dtype=bool))
mask_c2 = Mask(bbox=np.array([2, 2, 4, 4]), mask=np.ones((2, 2), dtype=bool))

ids: dict[str, int] = {}
ids["p"] = g.add_node({"t": 0, DEFAULT_ATTR_KEYS.MASK: mask_p})
Expand Down
4 changes: 2 additions & 2 deletions src/tracksdata/functional/_test/test_napari.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@
from tracksdata.constants import DEFAULT_ATTR_KEYS
from tracksdata.functional import to_napari_format
from tracksdata.graph import RustWorkXGraph
from tracksdata.nodes import MaskDiskAttrs
from tracksdata.nodes import MaskDiskAttrs, masks_from_column


@pytest.mark.parametrize("metadata_shape", [True, False])
Expand Down Expand Up @@ -45,7 +45,7 @@ def test_napari_conversion(metadata_shape: bool) -> None:

# Maybe we should update the MaskDiskAttrs to handle bounding boxes
graph.add_node_attr_key(DEFAULT_ATTR_KEYS.BBOX, dtype=pl.Array(pl.Int64, 6))
masks = graph.node_attrs(attr_keys=[DEFAULT_ATTR_KEYS.MASK])[DEFAULT_ATTR_KEYS.MASK]
masks = masks_from_column(graph.node_attrs(attr_keys=[DEFAULT_ATTR_KEYS.MASK])[DEFAULT_ATTR_KEYS.MASK])
graph.update_node_attrs(
attrs={DEFAULT_ATTR_KEYS.BBOX: [mask.bbox for mask in masks]},
node_ids=graph.node_ids(),
Expand Down
16 changes: 11 additions & 5 deletions src/tracksdata/graph/_base_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -1378,16 +1378,19 @@ def compute_overlaps(self, iou_threshold: float = 0.0) -> None:
raise ValueError("iou_threshold must be between 0.0 and 1.0")

def _estimate_overlaps(t: int) -> list[list[int, 2]]:
# Local import: avoids the graph <-> nodes package import cycle.
from tracksdata.nodes._mask import mask_iou, masks_from_column

node_attrs = self.filter(NodeAttr(DEFAULT_ATTR_KEYS.T) == t).node_attrs(
attr_keys=[DEFAULT_ATTR_KEYS.NODE_ID, DEFAULT_ATTR_KEYS.MASK],
)
node_ids = node_attrs[DEFAULT_ATTR_KEYS.NODE_ID].to_list()
masks = node_attrs[DEFAULT_ATTR_KEYS.MASK].to_list()
masks = masks_from_column(node_attrs[DEFAULT_ATTR_KEYS.MASK])
overlaps = []
for i in range(len(masks)):
mask_i = masks[i]
for j in range(i + 1, len(masks)):
if mask_i.iou(masks[j]) > iou_threshold:
if mask_iou(mask_i, masks[j]) > iou_threshold:
overlaps.append([node_ids[i], node_ids[j]])
return overlaps

Expand Down Expand Up @@ -1840,8 +1843,8 @@ def from_geff(
# unsafe operation, changing graph content inplace
for node_attr in indexed_graph.rx_graph.nodes():
node_attr[DEFAULT_ATTR_KEYS.MASK] = Mask(
node_attr[DEFAULT_ATTR_KEYS.MASK].astype(bool),
bbox=node_attr[DEFAULT_ATTR_KEYS.BBOX],
bbox=np.asarray(node_attr[DEFAULT_ATTR_KEYS.BBOX], dtype=np.int64),
mask=node_attr[DEFAULT_ATTR_KEYS.MASK].astype(bool),
)

if cls == IndexedRXGraph:
Expand Down Expand Up @@ -1939,8 +1942,11 @@ def to_geff(
}

if DEFAULT_ATTR_KEYS.MASK in node_attrs.columns:
# Local import: avoids the graph <-> nodes package import cycle.
from tracksdata.nodes._mask import masks_from_column

node_dict[DEFAULT_ATTR_KEYS.MASK] = construct_var_len_props(
[mask.mask.astype(bool) for mask in node_attrs[DEFAULT_ATTR_KEYS.MASK]]
[mask.mask.astype(bool) for mask in masks_from_column(node_attrs[DEFAULT_ATTR_KEYS.MASK])]
)

edge_dict = {k: {"values": column_to_numpy(v), "missing": None} for k, v in edge_attrs.to_dict().items()}
Expand Down
5 changes: 4 additions & 1 deletion src/tracksdata/graph/_rustworkx_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,10 @@ def _pop_time_eq(
def _maybe_fill_null(s: pl.Series, schema: AttrSchema) -> pl.Series:
if s.has_nulls() and schema.default_value is not None:
if schema.dtype == pl.Object:
value = pl.lit(schema.default_value, allow_object=True)
# `pl.lit` infers a struct from a dict default (e.g. a `Mask`), which then
# has no supertype with the object column. Wrapping in an object-typed
# series forces the literal to stay an object.
value = pl.lit(pl.Series([schema.default_value], dtype=pl.Object)).first()

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'm guessing you already tried, but is there an easier way to do this?
polars can be messy somethings when dealing with types :(

elif isinstance(schema.dtype, pl.Array):
if isinstance(schema.default_value, np.ndarray):
value = schema.default_value.tolist()
Expand Down
Loading
Loading