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
5 changes: 4 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,10 @@ readme = "README.md"
requires-python = ">=3.11"
dynamic = ["version"]
dependencies = [
"ezmsg>=3.6.2",
# 3.10.0b2 for AxisArray.chunk_dim, which the shmem sink follows to pick
# the buffered axis, and CoordinateAxis.fingerprint, which aux_meta uses
# to compare axes without touching their bytes.
"ezmsg>=3.10.0b2",
"numpy>=1.26.0",
"typer>=0.24.1",
]
Expand Down
25 changes: 25 additions & 0 deletions src/ezmsg/tools/plot/describe.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from ..chmeta import channel_names

__all__ = [
"stream_axis",
"METRIC_AXIS_CANDIDATES",
"METRIC_KINDS",
"SWEEP_RENDERABLE_METRICS",
Expand Down Expand Up @@ -107,6 +108,30 @@ def envelope(self) -> bool:
return self.metric is not None and self.metric.kind == "minmax"


def stream_axis(msg: typing.Any, *fallbacks: str) -> str | None:
"""Which dimension of *msg* the stream accumulates along.

Prefers the producer's own declaration
(:attr:`~ezmsg.util.messages.axisarray.AxisArray.chunk_dim`) and falls back
to the first of *fallbacks* the message actually has, which is what these
tools did before the field existed.

The fallback is a guess, and ``"time"`` is the wrong guess downstream of a
windowing stage: a ``(win, time, ch)`` message *has* a ``time`` dimension,
but it is the within-window lag, so a sweep plot keyed on it draws each
window's interior along the x-axis and treats the windows as channels --
and reads an offset that does not advance with the stream.
"""
chunk_dim = getattr(msg, "chunk_dim", None)
dims = getattr(msg, "dims", ()) or ()
if chunk_dim is not None and chunk_dim in dims:
return chunk_dim
for name in fallbacks:
if name in dims:
return name
return None


def metric_axis(dims: typing.Sequence[str], axes: typing.Mapping[str, typing.Any]) -> MetricSpec | None:
"""Describe the trailing per-sample tuple, or None if there is not one.

Expand Down
51 changes: 39 additions & 12 deletions src/ezmsg/tools/shmem/aux_meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,14 +18,24 @@
separate environments with different ezmsg versions installed; pinning the wire
format to ezmsg's dataclass layout would make an upgrade on one side a silent
decode failure on the other. Plain dicts cost one down-conversion and buy
version independence. ``AUX_FORMAT_VERSION`` guards the shape of the dict
itself.
version independence.

``AUX_FORMAT_VERSION`` guards *incompatible* changes to the dict. Adding a key
is not one: a reader that predates it ignores what it does not know, and a
reader that expects it reads a default when an older writer omits it, so both
directions keep working. Bumping the version for an additive change would break
exactly the mixed-version pairing this format exists to support.

Axes decode to::

{"kind": "linear", "unit": str, "gain": float, "offset": float}
{"kind": "coord", "unit": str, "dims": list[str], "data": np.ndarray}

``chunk_dim`` carries the source message's declaration of which dimension it
accumulates along, or ``None`` from a producer that declares nothing. It is
what the sink used to choose ``buffered_axis``, recorded so a consumer can tell
the two apart -- an operator may have overridden the buffered axis.

The buffered axis (normally ``time``) is a deliberate special case: its
``offset`` advances with every message and a coordinate time axis's ``data`` is
wholly new each message, so including either would make the metadata change
Expand Down Expand Up @@ -92,6 +102,7 @@ def encode_aux(
attrs: typing.Mapping[str, typing.Any],
key: str,
buffered_axis: str,
chunk_dim: typing.Optional[str] = None,
) -> tuple[bytes, list[str]]:
"""Serialize an AxisArray's static metadata.

Expand All @@ -113,6 +124,7 @@ def encode_aux(
"attrs": plain_attrs,
"key": key,
"buffered_axis": buffered_axis,
"chunk_dim": chunk_dim,
}
return pickle.dumps(payload, protocol=pickle.HIGHEST_PROTOCOL), dropped

Expand All @@ -134,22 +146,27 @@ def decode_aux(blob: bytes) -> dict:
raise ValueError(
f"shmem metadata blob is format version {version!r}, this build understands {AUX_FORMAT_VERSION}"
)
# Additive keys are defaulted rather than required, so a blob from a writer
# that predates them still decodes. See the module docstring.
payload.setdefault("chunk_dim", None)
return payload


def _axis_equal(a: typing.Any, b: typing.Any) -> bool:
"""Value equality for one axis, compared field by field.

Deliberately does not use ``==``. As of ezmsg 3.6, ``CoordinateAxis.__eq__``
resolves through the MRO to the dataclass-generated ``AxisBase.__eq__``,
which compares ``unit`` and nothing else -- ``ArrayWithNamedDims.__eq__``,
written to compare ``dims`` and ``data``, is shadowed and never runs. Two
coordinate axes with different data, or even different lengths, therefore
compare equal. Relying on that would mean a channel relabelling silently
never reaching the far side of the shmem link, which is the one thing this
module exists to deliver. Comparing explicitly also keeps the check correct
across ezmsg versions, which matters given the two halves of a link need not
share one.
Deliberately does not use ``==``. From ezmsg 3.6 to 3.9,
``CoordinateAxis.__eq__`` resolved through the MRO to the
dataclass-generated ``AxisBase.__eq__``, which compares ``unit`` and nothing
else -- ``ArrayWithNamedDims.__eq__``, written to compare ``dims`` and
``data``, was shadowed and never ran. Two coordinate axes with different
data, or even different lengths, compared equal. Relying on that would mean
a channel relabelling silently never reaching the far side of the shmem
link, which is the one thing this module exists to deliver.

Fixed in ezmsg 3.10, but the explicit comparison stays: the two halves of a
link need not share an ezmsg version, and a writer on 3.9 is still a writer
this has to be correct for.
"""
a_data = getattr(a, "data", None)
b_data = getattr(b, "data", None)
Expand All @@ -163,6 +180,16 @@ def _axis_equal(a: typing.Any, b: typing.Any) -> bool:
return False
if a_data is b_data:
return True
# ezmsg >= 3.10 derives a cached content digest per axis object. Ask before
# reading the bytes: an upstream that computed it -- every ezmsg source does
# now -- makes this comparison O(1) on an axis this process has already
# seen, where `array_equal` is O(bytes) every message. Absent on older
# ezmsg, and None for a dtype it cannot digest, so it is a pure fast path.
a_fp = getattr(a, "fingerprint", None)
if a_fp is not None:
b_fp = getattr(b, "fingerprint", None)
if b_fp is not None:
return bool(a_fp == b_fp)
if a_data.shape != b_data.shape or a_data.dtype != b_data.dtype:
return False
return bool(np.array_equal(a_data, b_data))
Expand Down
51 changes: 44 additions & 7 deletions src/ezmsg/tools/shmem/shmem.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,7 +157,22 @@ class ShMemCircBuffSettings(ez.Settings):
shmem_name: typing.Optional[str]
buf_dur: float
conn: typing.Optional[multiprocessing.connection.Connection] = None
axis: str = "time"

axis: typing.Optional[str] = None
"""Dimension to buffer along. ``None`` follows the message's ``chunk_dim``.

The ring is a history of the stream, so this has to be the dimension
messages accumulate along; buffering a static one would store the same
elements over and over. Only the producer reliably knows which that is --
it is ``time`` on a raw signal and ``win`` downstream of a windowing stage.

The old default of ``"time"`` was silently wrong for the latter. It is
*present* in a ``(win, time, ch)`` message, so nothing rejected it: the
window count ended up inside ``frame_shape`` -- reallocating the buffer
whenever the window count jittered -- and the reported sample rate was the
within-window rate, a 10x error in the viewer's time base for a 10-sample
window. Set explicitly only for a producer that declares no ``chunk_dim``.
"""


class ShMemCircBuffState(ez.State):
Expand All @@ -166,6 +181,8 @@ class ShMemCircBuffState(ez.State):
buffer_shmem: typing.Optional[SharedMemory] = None
buffer_arr: typing.Optional[npt.NDArray] = None
meta_hash: int = -1
# The dimension currently buffered along; see ShMemCircBuffSettings.axis.
buff_axis: typing.Optional[str] = None
# Segment holding the serialized static metadata (see .aux_meta).
aux_shmem: typing.Optional[SharedMemory] = None
# The (dims, axes, attrs, key) we last encoded, held by reference for the
Expand Down Expand Up @@ -388,8 +405,9 @@ def _update_aux_if_needed(self, msg: AxisArray) -> bool:
# meta.shape already describes that order, so dims must too -- a reader
# given the sender's original order would have to know to re-roll it,
# which is knowledge it has no way to arrive at.
rolled_dims = [self.SETTINGS.axis] + [d for d in msg.dims if d != self.SETTINGS.axis]
blob, dropped = encode_aux(rolled_dims, msg.axes, msg.attrs, msg.key, self.SETTINGS.axis)
buff_axis = self.STATE.buff_axis
rolled_dims = [buff_axis] + [d for d in msg.dims if d != buff_axis]
blob, dropped = encode_aux(rolled_dims, msg.axes, msg.attrs, msg.key, buff_axis, chunk_dim=msg.chunk_dim)
if dropped:
dropped_set = frozenset(dropped)
if self.STATE.warned_dropped_attrs != dropped_set:
Expand Down Expand Up @@ -433,6 +451,17 @@ def _cleanup_aux_segment(self) -> None:
del self.STATE.aux_shmem
self.STATE.aux_shmem = None

def _resolve_axis(self, msg: AxisArray) -> typing.Optional[str]:
"""The dimension to buffer along, or None if this message has none.

An explicit setting wins so an operator can still drive a producer that
declares nothing; otherwise the message decides.
"""
axis = self.SETTINGS.axis if self.SETTINGS.axis is not None else msg.chunk_dim
if axis is None:
axis = "time" if "time" in msg.dims else None
return axis

def _n_frames_for_axis(self, axis: AxisBase) -> int:
"""
Utility function to calculate the number of frames to allocate for the buffer.
Expand All @@ -459,8 +488,9 @@ def _get_msg_meta(self, msg: AxisArray) -> typing.Tuple[bytes, float, int, typin
A tuple of metadata extracted from the message.
msg_dtype, msg_srate, n_frames, frame_shape
"""
ax_idx = msg.get_axis_idx(self.SETTINGS.axis)
axis = msg.axes[self.SETTINGS.axis]
buff_axis = self.STATE.buff_axis
ax_idx = msg.get_axis_idx(buff_axis)
axis = msg.axes[buff_axis]
n_frames = self._n_frames_for_axis(axis)
frame_shape = msg.data.shape[:ax_idx] + msg.data.shape[ax_idx + 1 :]
data = np.moveaxis(msg.data, ax_idx, 0)
Expand Down Expand Up @@ -558,10 +588,17 @@ async def on_message(self, msg: AxisArray):
if not isinstance(msg, AxisArray):
return

if self.SETTINGS.axis not in msg.dims:
buff_axis = self._resolve_axis(msg)
if buff_axis is None or buff_axis not in msg.dims:
return
if buff_axis != self.STATE.buff_axis:
# The dimension the stream accumulates along changed under us -- a
# windowing stage inserted upstream, say. The buffer describes the
# old one, so it cannot be appended to.
self.STATE.buff_axis = buff_axis
self._cleanup_buffer()

ax_idx = msg.get_axis_idx(self.SETTINGS.axis)
ax_idx = msg.get_axis_idx(buff_axis)
data = np.moveaxis(msg.data, ax_idx, 0)

# Check if we need to update the metadata, and if so, reset the buffer.
Expand Down
16 changes: 16 additions & 0 deletions src/ezmsg/tools/shmem/shmem_mirror.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,6 +120,22 @@ def dims(self) -> typing.Optional[typing.List[str]]:
self._refresh_aux()
return None if self._aux is None else self._aux["dims"]

@property
def chunk_dim(self) -> typing.Optional[str]:
"""Which dimension the *source* declared it accumulates along.

Distinct from the buffered axis: an operator can override that, and a
producer on ezmsg < 3.10 declares nothing, in which case this is None.
"""
self._refresh_aux()
return None if self._aux is None else self._aux.get("chunk_dim")

@property
def buffered_axis(self) -> typing.Optional[str]:
"""Which dimension the ring is a history along."""
self._refresh_aux()
return None if self._aux is None else self._aux.get("buffered_axis")

@property
def metadata_available(self) -> bool:
"""Whether a decoded metadata blob is currently held."""
Expand Down
9 changes: 6 additions & 3 deletions src/ezmsg/tools/sigmon/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
describe_axisarray,
flatten_for_plot,
require_sweep_renderable,
stream_axis,
)
from ezmsg.tools.sigmon.dag_widget import DAGWidget

Expand Down Expand Up @@ -222,7 +223,8 @@ def _push_message(self, msg) -> None:
widget = self._plot_widget

if isinstance(widget, SweepWidget):
time_idx = msg.get_axis_idx("time") if "time" in msg.dims else 0
sweep_dim = stream_axis(msg, "time")
time_idx = msg.get_axis_idx(sweep_dim) if sweep_dim else 0
shape = self._shape or describe_axisarray(msg)
data = flatten_for_plot(np.moveaxis(msg.data, time_idx, 0), shape)
widget.push_data(data.astype(np.float32))
Expand All @@ -238,8 +240,9 @@ def _push_message(self, msg) -> None:
# Scatter expects (n_channels,) or (n_samples, n_channels).
if len(msg.shape) > 1:
targ_idx = 0
if "time" in msg.dims or "freq" in msg.dims:
targ_idx = msg.get_axis_idx("time") if "time" in msg.dims else msg.get_axis_idx("freq")
targ_dim = stream_axis(msg, "time", "freq")
if targ_dim is not None:
targ_idx = msg.get_axis_idx(targ_dim)
n_items = msg.shape[targ_idx]
n_channels = msg.data.size // n_items if n_items > 0 else 1
data_2d = np.moveaxis(msg.data, targ_idx, 0).reshape(n_items, n_channels)
Expand Down
17 changes: 10 additions & 7 deletions src/ezmsg/tools/viewer/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
describe_axisarray,
flatten_for_plot,
require_sweep_renderable,
stream_axis,
)

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -172,12 +173,15 @@ def _push_message(self, msg) -> None:
widget = self._plot_widget

if isinstance(widget, SweepWidget):
time_idx = msg.get_axis_idx("time") if "time" in msg.dims else 0
# The dimension the stream accumulates along, which is `win` rather
# than `time` downstream of a windowing stage. See `stream_axis`.
sweep_dim = stream_axis(msg, "time")
time_idx = msg.get_axis_idx(sweep_dim) if sweep_dim else 0
shape = self._shape or describe_axisarray(msg)
data = flatten_for_plot(np.moveaxis(msg.data, time_idx, 0), shape)
# Pass the AxisArray time-axis offset so the sweep buffer
# tracks the same clock as the event timestamps.
ts = msg.get_axis("time").offset if "time" in msg.dims else None
# Pass that axis's offset so the sweep buffer tracks the same clock
# as the event timestamps.
ts = msg.get_axis(sweep_dim).offset if sweep_dim else None
widget.push_data(data.astype(np.float32), timestamps=ts)

elif isinstance(widget, SpectrumWidget):
Expand All @@ -189,9 +193,8 @@ def _push_message(self, msg) -> None:

elif isinstance(widget, ScatterWidget):
if len(msg.shape) > 1:
targ_idx = 0
if "time" in msg.dims or "freq" in msg.dims:
targ_idx = msg.get_axis_idx("time") if "time" in msg.dims else msg.get_axis_idx("freq")
targ_dim = stream_axis(msg, "time", "freq")
targ_idx = msg.get_axis_idx(targ_dim) if targ_dim else 0
n_items = msg.shape[targ_idx]
n_channels = msg.data.size // n_items if n_items > 0 else 1
data_2d = np.moveaxis(msg.data, targ_idx, 0).reshape(n_items, n_channels)
Expand Down
Loading
Loading