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] diff --git a/src/ezmsg/tools/plot/__init__.py b/src/ezmsg/tools/plot/__init__.py index 8b6ff5c..ef7143e 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 ChannelLayoutCache, channel_layout from .shmem_sweep import ShmemSweepWidget __all__ = [ @@ -38,7 +39,9 @@ "MetricSpec", "ShmemSweepWidget", "StreamShape", + "ChannelLayoutCache", "UnsupportedMetricError", + "channel_layout", "describe_axisarray", "describe_mirror", "flatten_for_plot", @@ -48,9 +51,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 in ("channel_layout", "ChannelLayoutCache"): + from . import 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 new file mode 100644 index 0000000..103d92b --- /dev/null +++ b/src/ezmsg/tools/plot/layout.py @@ -0,0 +1,116 @@ +"""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__ = ["ChannelLayoutCache", "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) + + +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 new file mode 100644 index 0000000..1d0ef0d --- /dev/null +++ b/tests/test_plot_layout.py @@ -0,0 +1,207 @@ +"""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 + +# 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")]) + + +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 + + +# ---- 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)