Skip to content
Merged
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
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
11 changes: 9 additions & 2 deletions src/ezmsg/tools/plot/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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__ = [
Expand All @@ -38,7 +39,9 @@
"MetricSpec",
"ShmemSweepWidget",
"StreamShape",
"ChannelLayoutCache",
"UnsupportedMetricError",
"channel_layout",
"describe_axisarray",
"describe_mirror",
"flatten_for_plot",
Expand All @@ -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}")
116 changes: 116 additions & 0 deletions src/ezmsg/tools/plot/layout.py
Original file line number Diff line number Diff line change
@@ -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
207 changes: 207 additions & 0 deletions tests/test_plot_layout.py
Original file line number Diff line number Diff line change
@@ -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)
Loading