From 26c3df9bd76334302ef3a0c6d114c31d9e045e38 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 6 Aug 2026 12:19:21 -0400 Subject: [PATCH 1/4] Derive a grid layout from the ch axis The geometry counterpart to chmeta, which already does this for names: given a ch coordinate axis, work out where each channel's cell goes and how big it is. It comes from intent-tools, and deliberately did not move into phosphor with the grids it feeds. Reading x/y/size/headstage off a structured axis is decoding one acquisition system's convention, and a renderer that takes plain arrays should not have to know it. This is the layer that already owns that translation. Labels delegate to chmeta.channel_names rather than repeating its fallback, so a grid can name channels by bank and elec exactly as the sweep does -- the intent-tools version could only read a label field. Everything degrades to a square-ish tiling: no ch axis, no coordinate fields, or only one of the two. A plot that draws nothing is less useful than one that draws the right number of cells in the wrong places, and geometry is missing often enough that this is the common path, not the edge case. Exported lazily, like ShmemSweepWidget: it needs phosphor's geometry helpers, and importing ezmsg.tools.plot must not pull in a GPU stack. 12 tests. --- src/ezmsg/tools/plot/__init__.py | 10 ++- src/ezmsg/tools/plot/layout.py | 73 +++++++++++++++++ tests/test_plot_layout.py | 131 +++++++++++++++++++++++++++++++ 3 files changed, 212 insertions(+), 2 deletions(-) create mode 100644 src/ezmsg/tools/plot/layout.py create mode 100644 tests/test_plot_layout.py diff --git a/src/ezmsg/tools/plot/__init__.py b/src/ezmsg/tools/plot/__init__.py index 8b6ff5c..01bb83f 100644 --- a/src/ezmsg/tools/plot/__init__.py +++ b/src/ezmsg/tools/plot/__init__.py @@ -6,7 +6,7 @@ :mod:`.shmem_sweep` is the Qt widget built on it, and needs the ``viewer`` or ``sigmon`` extra. -``ShmemSweepWidget`` is resolved lazily so that importing this package, or +``ShmemSweepWidget`` and :mod:`.layout` are resolved lazily so that importing this package, or anything under it, does not pull in Qt. Eagerly importing it here would make ``from ezmsg.tools.plot.describe import ...`` fail without phosphor installed, since importing a submodule runs its parent's ``__init__`` first -- which would @@ -30,6 +30,7 @@ ) if typing.TYPE_CHECKING: # pragma: no cover - import for type checkers only + from .layout import channel_layout from .shmem_sweep import ShmemSweepWidget __all__ = [ @@ -39,6 +40,7 @@ "ShmemSweepWidget", "StreamShape", "UnsupportedMetricError", + "channel_layout", "describe_axisarray", "describe_mirror", "flatten_for_plot", @@ -48,9 +50,13 @@ def __getattr__(name: str) -> typing.Any: - """Resolve the Qt widget on first use (PEP 562).""" + """Resolve the phosphor-backed names on first use (PEP 562).""" if name == "ShmemSweepWidget": from .shmem_sweep import ShmemSweepWidget return ShmemSweepWidget + if name == "channel_layout": + from .layout import channel_layout + + return channel_layout raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/src/ezmsg/tools/plot/layout.py b/src/ezmsg/tools/plot/layout.py new file mode 100644 index 0000000..4bc66ac --- /dev/null +++ b/src/ezmsg/tools/plot/layout.py @@ -0,0 +1,73 @@ +"""Turning a structured ``ch`` coordinate axis into a per-channel grid layout. + +The geometry counterpart to :mod:`ezmsg.tools.chmeta`, which does the same job +for names. An AxisArray's ``ch`` axis may carry electrode coordinates; a grid +plot wants positions, sizes and labels. Which fields those live in is a +property of the acquisition system, so the decoding belongs here rather than in +phosphor, whose grids take plain arrays and have no opinion about where they +came from. + +Everything degrades: a source with no ``ch`` axis at all, or one carrying names +but no coordinates, still gets a layout -- a square-ish tiling -- because a plot +that draws nothing is less useful than one that draws the right number of cells +in the wrong places. +""" + +import typing + +import numpy as np + +from ..chmeta import channel_names + +__all__ = ["channel_layout"] + +DEFAULT_POSITION_FIELDS = ("x", "y") +DEFAULT_SIZE_FIELD = "size" +DEFAULT_GROUP_FIELD = "headstage" + + +def channel_layout( + ch_axis_data: typing.Optional[np.ndarray], + n_ch: int, + *, + position_fields: typing.Tuple[str, str] = DEFAULT_POSITION_FIELDS, + size_field: str = DEFAULT_SIZE_FIELD, + group_field: typing.Optional[str] = DEFAULT_GROUP_FIELD, + label_fields: typing.Sequence[str] = ("label",), +) -> typing.Tuple[np.ndarray, typing.Optional[np.ndarray], typing.List[str]]: + """Per-channel ``(positions, sizes, labels)`` for a grid plot. + + :param ch_axis_data: The ``ch`` axis' structured data, or ``None`` when the + stream carries no channel metadata. + :param n_ch: Channels in the data, used when the axis cannot say. + :param position_fields: Fields holding each channel's coordinates. Both must + be present, or the layout falls back to a tiling. + :param size_field: Field holding each channel's extent. Absent gives + ``None``, which lets the renderer size cells by the inferred pitch. + :param group_field: Field identifying which device a channel belongs to. + Devices commonly number their electrodes from a shared origin, so + without this two of them draw on top of each other. ``None`` skips the + check. + :param label_fields: Passed to :func:`~ezmsg.tools.chmeta.channel_names`. + + :returns: ``positions`` as ``(n, 2)`` float32, ``sizes`` as ``(n,)`` float32 + or ``None``, and one label per channel. + """ + from phosphor.grid_layout import tile_by_group, tiled_grid_positions + + if ch_axis_data is None: + return tiled_grid_positions(n_ch), None, channel_names(None, n_ch, fields=label_fields) + + fields = ch_axis_data.dtype.fields or {} + actual_n = ch_axis_data.shape[0] + x_field, y_field = position_fields + + if {x_field, y_field} <= set(fields): + positions = np.column_stack([ch_axis_data[x_field], ch_axis_data[y_field]]).astype(np.float32) + if group_field and group_field in fields: + positions = tile_by_group(positions, ch_axis_data[group_field]) + else: + positions = tiled_grid_positions(actual_n) + + sizes = ch_axis_data[size_field].astype(np.float32) if size_field in fields else None + return positions, sizes, channel_names(ch_axis_data, actual_n, fields=label_fields) diff --git a/tests/test_plot_layout.py b/tests/test_plot_layout.py new file mode 100644 index 0000000..2c46d3a --- /dev/null +++ b/tests/test_plot_layout.py @@ -0,0 +1,131 @@ +"""Deriving a grid layout from a ``ch`` coordinate axis. + +The decoding half of the grid: which fields hold coordinates, and what to do +when they are missing. Every fallback here exists because a plot that draws +nothing is less useful than one that draws the right number of cells in the +wrong places -- and because in practice the geometry is often absent, partial, +or shared between two devices that each numbered from their own origin. +""" + +import numpy as np +import pytest + +from ezmsg.tools.plot import channel_layout + +GEOMETRY = np.dtype([("x", "f4"), ("y", "f4"), ("size", "f4"), ("label", "U8"), ("headstage", "i4")]) + + +def make_axis(n=4, *, fields=("x", "y", "size", "label", "headstage")): + """A ``ch`` axis carrying only the named fields.""" + dt = np.dtype([(name, GEOMETRY[name]) for name in fields]) + ch = np.zeros(n, dtype=dt) + if "x" in fields: + ch["x"] = np.arange(n) % 2 + if "y" in fields: + ch["y"] = np.arange(n) // 2 + if "size" in fields: + ch["size"] = 0.5 + if "label" in fields: + ch["label"] = [f"E{i}" for i in range(n)] + return ch + + +def test_coordinates_are_used_verbatim(): + ch = make_axis(4, fields=("x", "y")) + positions, sizes, labels = channel_layout(ch, 4) + + np.testing.assert_allclose(positions, [[0, 0], [1, 0], [0, 1], [1, 1]]) + assert sizes is None, "no size field means the renderer picks from the pitch" + assert labels == ["ch0", "ch1", "ch2", "ch3"] + + +def test_sizes_and_labels_come_along_when_present(): + positions, sizes, labels = channel_layout(make_axis(4), 4) + assert positions.shape == (4, 2) + np.testing.assert_allclose(sizes, 0.5) + assert labels == ["E0", "E1", "E2", "E3"] + + +def test_devices_sharing_a_coordinate_range_are_separated(): + """Two headstages numbering electrodes from the same origin would otherwise + render one on top of the other.""" + ch = make_axis(4) + ch["x"] = [0, 1, 0, 1] + ch["y"] = [0, 0, 0, 0] + ch["headstage"] = [0, 0, 1, 1] + + positions, _, _ = channel_layout(ch, 4) + assert positions[2, 0] > positions[1, 0], "the second device must start clear of the first" + + +def test_the_grouping_field_can_be_ignored(): + ch = make_axis(4) + ch["x"] = [0, 1, 0, 1] + ch["headstage"] = [0, 0, 1, 1] + + positions, _, _ = channel_layout(ch, 4, group_field=None) + np.testing.assert_allclose(positions[:, 0], [0, 1, 0, 1]) + + +def test_an_axis_without_coordinates_still_gets_a_layout(): + """Channel names but no geometry: common for LSL and plain NWB streams.""" + ch = make_axis(4, fields=("label",)) + positions, sizes, labels = channel_layout(ch, 4) + + assert positions.shape == (4, 2) + assert len({tuple(p) for p in positions}) == 4, "cells must not stack" + assert sizes is None + assert labels == ["E0", "E1", "E2", "E3"], "names survive even with no geometry" + + +def test_no_channel_axis_at_all_still_gets_a_layout(): + positions, sizes, labels = channel_layout(None, 3) + assert positions.shape == (3, 2) + assert sizes is None + assert labels == ["ch0", "ch1", "ch2"] + + +def test_only_one_coordinate_is_not_half_a_layout(): + """An x with no y cannot place anything; falling back is the honest move.""" + ch = make_axis(4, fields=("x", "label")) + positions, _, _ = channel_layout(ch, 4) + assert len({tuple(p) for p in positions}) == 4 + + +def test_the_coordinate_fields_are_configurable(): + """Nothing says a source must call them x and y.""" + dt = np.dtype([("col", "f4"), ("row", "f4")]) + ch = np.zeros(3, dtype=dt) + ch["col"] = [5.0, 6.0, 7.0] + ch["row"] = 1.0 + + positions, _, _ = channel_layout(ch, 3, position_fields=("col", "row")) + np.testing.assert_allclose(positions[:, 0], [5.0, 6.0, 7.0]) + + +def test_labels_follow_the_requested_fields(): + """The same choice the sweep offers: a Blackrock user reads bank and elec + off the front panel, not a label field.""" + dt = np.dtype([("x", "f4"), ("y", "f4"), ("bank", "U4"), ("elec", "i4")]) + ch = np.zeros(2, dtype=dt) + ch["bank"] = ["A", "A"] + ch["elec"] = [1, 2] + + _, _, labels = channel_layout(ch, 2, label_fields=("bank", "elec")) + assert labels == ["A-1", "A-2"] + + +def test_the_axis_decides_the_count_when_it_disagrees(): + """The axis is the thing that knows how many rows it has; trusting a stale + n_ch would index off the end of it.""" + ch = make_axis(4) + positions, sizes, labels = channel_layout(ch, 99) + assert positions.shape[0] == 4 + assert len(labels) == 4 + + +@pytest.mark.parametrize("n", [0, 1]) +def test_degenerate_channel_counts_do_not_raise(n): + positions, _, labels = channel_layout(None, n) + assert positions.shape == (n, 2) + assert len(labels) == n From 8b7097484da171c7cc02181652fc765ac5095350 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 6 Aug 2026 19:20:50 -0400 Subject: [PATCH 2/4] Derive a channel layout once per change, not once per message A grid deriving its layout from every message pays for it on every message, while the answer changes about once a session. ChannelLayoutCache wraps channel_layout and recomputes only when the axis actually differs: 102 us down to 6 us at 128 channels. Fingerprinted on the axis' bytes rather than its identity. It arrives deserialized from another process, so it is a new object every message describing the same electrodes -- keying on identity would make the cache miss exactly where it is wanted. A cache rather than a shared instance, so two consumers watching one stream each keep their own and neither has to know the other exists. The derivation stays here, where any application can call it. --- src/ezmsg/tools/plot/__init__.py | 9 +++-- src/ezmsg/tools/plot/layout.py | 45 +++++++++++++++++++++- tests/test_plot_layout.py | 66 +++++++++++++++++++++++++++++++- 3 files changed, 114 insertions(+), 6 deletions(-) diff --git a/src/ezmsg/tools/plot/__init__.py b/src/ezmsg/tools/plot/__init__.py index 01bb83f..ef7143e 100644 --- a/src/ezmsg/tools/plot/__init__.py +++ b/src/ezmsg/tools/plot/__init__.py @@ -30,7 +30,7 @@ ) if typing.TYPE_CHECKING: # pragma: no cover - import for type checkers only - from .layout import channel_layout + from .layout import ChannelLayoutCache, channel_layout from .shmem_sweep import ShmemSweepWidget __all__ = [ @@ -39,6 +39,7 @@ "MetricSpec", "ShmemSweepWidget", "StreamShape", + "ChannelLayoutCache", "UnsupportedMetricError", "channel_layout", "describe_axisarray", @@ -55,8 +56,8 @@ def __getattr__(name: str) -> typing.Any: from .shmem_sweep import ShmemSweepWidget return ShmemSweepWidget - if name == "channel_layout": - from .layout import channel_layout + if name in ("channel_layout", "ChannelLayoutCache"): + from . import layout - return channel_layout + return getattr(layout, name) raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/src/ezmsg/tools/plot/layout.py b/src/ezmsg/tools/plot/layout.py index 4bc66ac..103d92b 100644 --- a/src/ezmsg/tools/plot/layout.py +++ b/src/ezmsg/tools/plot/layout.py @@ -19,7 +19,7 @@ from ..chmeta import channel_names -__all__ = ["channel_layout"] +__all__ = ["ChannelLayoutCache", "channel_layout"] DEFAULT_POSITION_FIELDS = ("x", "y") DEFAULT_SIZE_FIELD = "size" @@ -71,3 +71,46 @@ def channel_layout( sizes = ch_axis_data[size_field].astype(np.float32) if size_field in fields else None return positions, sizes, channel_names(ch_axis_data, actual_n, fields=label_fields) + + +class ChannelLayoutCache: + """:func:`channel_layout`, recomputed only when the axis actually changes. + + A grid that derives its layout per message pays for it per message, while + the answer changes about once a session. Fingerprinting the axis costs + around 1 us against the 75 us the derivation takes, so this is worth having + wherever messages arrive faster than the geometry does. + + Deliberately a cache rather than a shared instance: two widgets watching one + stream each keep their own, so neither has to know the other exists, and the + derivation stays where any application can call it. + """ + + def __init__(self) -> None: + self._key: typing.Optional[tuple] = None + self._value: typing.Optional[tuple] = None + + @staticmethod + def _fingerprint(ch_axis_data: typing.Optional[np.ndarray], n_ch: int, kwargs: dict) -> tuple: + """Enough to tell one layout apart from another, cheaply. + + The axis' bytes rather than its identity: it arrives deserialized from + another process, so a new object every message describes the same + electrodes. + """ + options = tuple(sorted((k, tuple(v) if isinstance(v, (list, tuple)) else v) for k, v in kwargs.items())) + if ch_axis_data is None: + return (n_ch, options, None) + return (n_ch, options, str(ch_axis_data.dtype), ch_axis_data.shape, ch_axis_data.tobytes()) + + def __call__( + self, + ch_axis_data: typing.Optional[np.ndarray], + n_ch: int, + **kwargs: typing.Any, + ) -> typing.Tuple[np.ndarray, typing.Optional[np.ndarray], typing.List[str]]: + key = self._fingerprint(ch_axis_data, n_ch, kwargs) + if key != self._key: + self._value = channel_layout(ch_axis_data, n_ch, **kwargs) + self._key = key + return self._value diff --git a/tests/test_plot_layout.py b/tests/test_plot_layout.py index 2c46d3a..45a4d54 100644 --- a/tests/test_plot_layout.py +++ b/tests/test_plot_layout.py @@ -10,7 +10,7 @@ import numpy as np import pytest -from ezmsg.tools.plot import channel_layout +from ezmsg.tools.plot import ChannelLayoutCache, channel_layout GEOMETRY = np.dtype([("x", "f4"), ("y", "f4"), ("size", "f4"), ("label", "U8"), ("headstage", "i4")]) @@ -129,3 +129,67 @@ def test_degenerate_channel_counts_do_not_raise(n): positions, _, labels = channel_layout(None, n) assert positions.shape == (n, 2) assert len(labels) == n + + +# ---- caching ---------------------------------------------------------------- +# +# A grid deriving its layout per message pays per message, while the answer +# changes about once a session. + + +def test_an_unchanged_axis_is_not_derived_twice(): + cache = ChannelLayoutCache() + ch = make_axis(4) + assert cache(ch, 4) is cache(ch, 4) + + +def test_an_equal_but_distinct_axis_still_hits(): + """The axis arrives deserialized from another process, so it is a new object + every message while describing the same electrodes. Keying on identity would + make the cache useless exactly where it is needed.""" + cache = ChannelLayoutCache() + ch = make_axis(4) + first = cache(ch, 4) + assert cache(ch.copy(), 4) is first + + +def test_a_changed_axis_is_derived_again(): + cache = ChannelLayoutCache() + ch = make_axis(4) + first = cache(ch, 4) + + moved = ch.copy() + moved["x"][0] = 99.0 + second = cache(moved, 4) + + assert second is not first + assert second[0][0, 0] == pytest.approx(99.0) + + +def test_a_changed_channel_count_is_derived_again(): + cache = ChannelLayoutCache() + assert cache(None, 4) is not cache(None, 9) + + +def test_changed_options_are_derived_again(): + """Same axis, different question -- the answer is not the cached one.""" + dt = np.dtype([("x", "f4"), ("y", "f4"), ("bank", "U4"), ("elec", "i4")]) + ch = np.zeros(2, dtype=dt) + ch["bank"] = ["A", "A"] + ch["elec"] = [1, 2] + + cache = ChannelLayoutCache() + assert cache(ch, 2)[2] == ["ch0", "ch1"] + assert cache(ch, 2, label_fields=("bank", "elec"))[2] == ["A-1", "A-2"] + + +def test_no_axis_at_all_caches_too(): + cache = ChannelLayoutCache() + assert cache(None, 3) is cache(None, 3) + + +def test_two_caches_do_not_share(): + """Each consumer keeps its own, so neither has to know the other exists.""" + a, b = ChannelLayoutCache(), ChannelLayoutCache() + ch = make_axis(4) + assert a(ch, 4) is not b(ch, 4) From ae49200802467871ae8978e56a14c2bf8ca42778 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 6 Aug 2026 21:20:02 -0400 Subject: [PATCH 3/4] Require phosphor 0.9.1 --- pyproject.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index e8f2d9d..1607409 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,14 +54,14 @@ sigmon = [ "pygraphviz>=1.14", "typer>=0.15.1", "ezmsg-qt>=0.2.1", - "phosphor>=0.8.0", + "phosphor>=0.9.1", "pandas", ] viewer = [ "PySide6>=6.7", "typer>=0.15.1", "ezmsg-qt>=0.2.1", - "phosphor>=0.8.0", + "phosphor>=0.9.1", ] [project.scripts] From 570c8acc26b9b5b75df1b3c9ff17232380e1b9e6 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 6 Aug 2026 21:23:11 -0400 Subject: [PATCH 4/4] Skip the layout tests where phosphor is not installed CI installs the test group, not the viewer or sigmon extras, so phosphor is absent -- which is the point of it being optional, and which these tests did not allow for. Nine jobs failed on ModuleNotFoundError. Guarded the way test_shmem_sweep already guards for the same reason, but keyed on phosphor itself rather than the module under test: layout.py imports it inside the function, so the module imports cleanly without it and only fails when called. Keying on the module would never have skipped. 19 pass where phosphor is installed, 1 skipped where it is not. --- tests/test_plot_layout.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/tests/test_plot_layout.py b/tests/test_plot_layout.py index 45a4d54..1d0ef0d 100644 --- a/tests/test_plot_layout.py +++ b/tests/test_plot_layout.py @@ -10,7 +10,19 @@ import numpy as np import pytest -from ezmsg.tools.plot import ChannelLayoutCache, channel_layout +# The geometry helpers this builds on live in phosphor, which is an optional +# extra -- so a runner that installs only the test group cannot reach them. +# Same guard as test_shmem_sweep, and skipped for the same reason. +# Keyed on phosphor itself, not on the module under test: layout.py imports it +# inside the function, so the module imports fine without it and only fails when +# called. +pytest.importorskip( + "phosphor.grid_layout", + reason="needs phosphor (the 'viewer' or 'sigmon' extra)", + exc_type=ImportError, +) + +from ezmsg.tools.plot import ChannelLayoutCache, channel_layout # noqa: E402 GEOMETRY = np.dtype([("x", "f4"), ("y", "f4"), ("size", "f4"), ("label", "U8"), ("headstage", "i4")])