From 0b479c359c4e2cc6878778635059eebc31ec9b78 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Fri, 4 Sep 2026 02:08:47 -0400 Subject: [PATCH] Follow chunk_dim rather than assuming "time" `ShMemCircBuffSettings.axis` defaulted to `"time"`, and the ring is a history of the stream, so it has to be the dimension messages accumulate along. That is `time` on a raw signal and `win` downstream of a windowing stage, and only the producer reliably knows which. Nothing rejected the old default on a windowed stream, because `time` *is* present in a `(win, time, ch)` message. It just buffered the wrong thing: buffered axis = 'time' n_win=4: frame_shape=(4, 3) frames written=10 srate=100 n_win=7: frame_shape=(7, 3) frames written=10 srate=100 buffered axis = 'win' n_win=4: frame_shape=(10, 3) frames written=4 srate=10 n_win=7: frame_shape=(10, 3) frames written=7 srate=10 The window *count* lands inside `frame_shape`, so the buffer is reallocated whenever it jitters, and the rate published to the viewer is the within-window rate -- a 10x error in its time base for a 10-sample window. `axis` now defaults to None, meaning "follow the message". An explicit setting still wins, so an operator can drive a producer that declares nothing; without either, the old `"time"` fallback stands. The buffered dimension is resolved per message and held in state, and the buffer is torn down if it changes -- a windowing stage inserted upstream leaves the ring describing the old layout. The blob gains `chunk_dim`, recorded separately from `buffered_axis` so a consumer can tell the source's declaration from an operator's override, and the mirror exposes both. Adding a key deliberately does not bump `AUX_FORMAT_VERSION`: a reader that predates it ignores what it does not know and a newer reader defaults it, so a mixed-version link -- the pairing this plain-dict format exists to support -- keeps working. Bumping would break exactly that. There is a test for a blob written without the key. `_axis_equal` asks for `CoordinateAxis.fingerprint` before reading bytes. It is on the publisher's per-message path, and its identity shortcut misses precisely when a producer rebuilds its axes -- where the cached digest, which every ezmsg source now primes, turns an O(bytes) comparison into a comparison of two tuples. Absent on older ezmsg and None for an undigestable dtype, so it stays a pure fast path. Its docstring is also corrected: the `CoordinateAxis.__eq__` MRO bug it describes was fixed in ezmsg 3.10, but the explicit comparison stays, because the two halves of a link need not share a version. The same wrong assumption was in the viewer and sigmon plot paths, where a sweep keyed on `time` draws each window's interior along the x-axis and treats the windows as channels. Both now go through `describe.stream_axis`, which prefers the declaration and falls back to the old guess. A declaration naming a dimension the message no longer has is ignored rather than trusted. Requires ezmsg 3.10.0b2 for both fields. 72 passed, up from 52. Mutation-checked: ignoring `chunk_dim` in the sink fails 1, dropping it from the blob fails 3, making the fingerprint path answer unconditionally fails 3, and both failure modes of `stream_axis` fail 1 each. --- pyproject.toml | 5 +- src/ezmsg/tools/plot/describe.py | 25 +++++++ src/ezmsg/tools/shmem/aux_meta.py | 51 ++++++++++--- src/ezmsg/tools/shmem/shmem.py | 51 +++++++++++-- src/ezmsg/tools/shmem/shmem_mirror.py | 16 ++++ src/ezmsg/tools/sigmon/cli.py | 9 ++- src/ezmsg/tools/viewer/cli.py | 17 +++-- tests/test_plot_describe.py | 65 ++++++++++++++++ tests/test_shmem_aux_meta.py | 104 +++++++++++++++++++++++++- tests/test_shmem_sink.py | 103 ++++++++++++++++++++++++- 10 files changed, 413 insertions(+), 33 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 1607409..bb452d3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", ] diff --git a/src/ezmsg/tools/plot/describe.py b/src/ezmsg/tools/plot/describe.py index 5210e38..e957fb6 100644 --- a/src/ezmsg/tools/plot/describe.py +++ b/src/ezmsg/tools/plot/describe.py @@ -19,6 +19,7 @@ from ..chmeta import channel_names __all__ = [ + "stream_axis", "METRIC_AXIS_CANDIDATES", "METRIC_KINDS", "SWEEP_RENDERABLE_METRICS", @@ -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. diff --git a/src/ezmsg/tools/shmem/aux_meta.py b/src/ezmsg/tools/shmem/aux_meta.py index 61ec934..e6af6dd 100644 --- a/src/ezmsg/tools/shmem/aux_meta.py +++ b/src/ezmsg/tools/shmem/aux_meta.py @@ -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 @@ -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. @@ -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 @@ -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) @@ -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)) diff --git a/src/ezmsg/tools/shmem/shmem.py b/src/ezmsg/tools/shmem/shmem.py index d34d52e..853a8d5 100644 --- a/src/ezmsg/tools/shmem/shmem.py +++ b/src/ezmsg/tools/shmem/shmem.py @@ -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): @@ -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 @@ -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: @@ -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. @@ -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) @@ -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. diff --git a/src/ezmsg/tools/shmem/shmem_mirror.py b/src/ezmsg/tools/shmem/shmem_mirror.py index c38af94..7f269bf 100644 --- a/src/ezmsg/tools/shmem/shmem_mirror.py +++ b/src/ezmsg/tools/shmem/shmem_mirror.py @@ -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.""" diff --git a/src/ezmsg/tools/sigmon/cli.py b/src/ezmsg/tools/sigmon/cli.py index 900564d..318c242 100644 --- a/src/ezmsg/tools/sigmon/cli.py +++ b/src/ezmsg/tools/sigmon/cli.py @@ -22,6 +22,7 @@ describe_axisarray, flatten_for_plot, require_sweep_renderable, + stream_axis, ) from ezmsg.tools.sigmon.dag_widget import DAGWidget @@ -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)) @@ -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) diff --git a/src/ezmsg/tools/viewer/cli.py b/src/ezmsg/tools/viewer/cli.py index e4edac1..18015b1 100644 --- a/src/ezmsg/tools/viewer/cli.py +++ b/src/ezmsg/tools/viewer/cli.py @@ -23,6 +23,7 @@ describe_axisarray, flatten_for_plot, require_sweep_renderable, + stream_axis, ) logger = logging.getLogger(__name__) @@ -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): @@ -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) diff --git a/tests/test_plot_describe.py b/tests/test_plot_describe.py index cc71dd6..5bbbce1 100644 --- a/tests/test_plot_describe.py +++ b/tests/test_plot_describe.py @@ -17,6 +17,7 @@ flatten_for_plot, metric_axis, require_sweep_renderable, + stream_axis, ) CHANNEL_DTYPE = np.dtype([("bank", "U2"), ("elec", " AxisArray: + """`(win, time, ch)` -- what a windowing stage emits. + + `time` here is the *within-window* lag dimension. Both are LinearAxes, so + nothing about the message distinguishes them except `chunk_dim`. + """ + kwargs = {"chunk_dim": chunk_dim} if chunk_dim else {} + return AxisArray( + np.zeros((n_win, n_lag, n_ch), np.float32), + dims=["win", "time", "ch"], + axes={ + "win": AxisArray.TimeAxis(fs=10.0), + "time": AxisArray.TimeAxis(fs=100.0), + "ch": CoordinateAxis(data=np.array(["a", "b", "c"]), dims=["ch"]), + }, + key="dev", + **kwargs, + ) + + +class TestTheBlobCarriesChunkDim: + def test_it_round_trips(self): + msg = _windowed() + blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "win", chunk_dim=msg.chunk_dim) + assert decode_aux(blob)["chunk_dim"] == "win" + + def test_none_from_a_producer_that_declares_nothing(self): + msg = _windowed(chunk_dim=None) + blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "win", chunk_dim=msg.chunk_dim) + assert decode_aux(blob)["chunk_dim"] is None + + def test_it_is_distinct_from_the_buffered_axis(self): + """An operator can override which axis the ring buffers; the source's + own declaration is recorded separately so a consumer can tell.""" + msg = _windowed() + blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "time", chunk_dim=msg.chunk_dim) + payload = decode_aux(blob) + assert payload["buffered_axis"] == "time" + assert payload["chunk_dim"] == "win" + + def test_a_blob_written_before_the_key_existed_still_decodes(self): + """Adding a key must not break a mixed-version link -- that pairing is + the whole reason this format is plain dicts. See the module docstring.""" + import pickle + + msg = _windowed() + blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "win", chunk_dim="win") + payload = pickle.loads(blob) + del payload["chunk_dim"] # what an older writer emits + old_blob = pickle.dumps(payload, protocol=pickle.HIGHEST_PROTOCOL) + + decoded = decode_aux(old_blob) + assert decoded["chunk_dim"] is None + assert decoded["buffered_axis"] == "win" + + +class TestTheFingerprintFastPath: + """`_axis_equal` asks for the cached digest before reading the bytes. It has + to stay exactly as discriminating as the byte comparison it replaces.""" + + @staticmethod + def _coord(labels): + return CoordinateAxis(data=np.array(labels), dims=["ch"]) + + def test_equal_content_in_distinct_objects_compares_equal(self): + a, b = self._coord(["a", "b", "c"]), self._coord(["a", "b", "c"]) + assert a is not b + assert _axis_equal(a, b) + + def test_a_relabel_at_fixed_length_compares_unequal(self): + """The case the module exists to deliver: same key, same channel count, + different channels.""" + assert not _axis_equal(self._coord(["a", "b", "c"]), self._coord(["x", "y", "z"])) + + def test_a_length_change_compares_unequal(self): + assert not _axis_equal(self._coord(["a", "b"]), self._coord(["a", "b", "c"])) + + def test_it_uses_the_digest_when_one_is_available(self): + a, b = self._coord(["a", "b", "c"]), self._coord(["a", "b", "c"]) + a.fingerprint # prime, as every ezmsg source now does + b.fingerprint + called = [] + real = np.array_equal + + def spy(*args, **kwargs): + called.append(1) + return real(*args, **kwargs) + + np.array_equal = spy + try: + assert _axis_equal(a, b) + finally: + np.array_equal = real + assert not called, "should have settled on the digest without reading the bytes" diff --git a/tests/test_shmem_sink.py b/tests/test_shmem_sink.py index 1b5192b..47fd0be 100644 --- a/tests/test_shmem_sink.py +++ b/tests/test_shmem_sink.py @@ -9,10 +9,10 @@ from ezmsg.simbiophys.eeg import EEGSynth from ezmsg.util.messagecodec import message_log from ezmsg.util.messagelogger import MessageLogger -from ezmsg.util.messages.axisarray import AxisArray +from ezmsg.util.messages.axisarray import AxisArray, CoordinateAxis from ezmsg.util.terminate import TerminateOnTotal -from ezmsg.tools.shmem.shmem import ShMemCircBuff +from ezmsg.tools.shmem.shmem import ShMemCircBuff, ShMemCircBuffSettings class CrazyUnitSettings(ez.Settings): @@ -113,3 +113,102 @@ def test_shmem_change(change_type: str): assert all(msg.data.dtype == float for msg in messages) file_path.unlink(missing_ok=True) + + +# --------------------------------------------------------------------------- +# Which dimension the ring is a history along +# --------------------------------------------------------------------------- + + +def _windowed_msg(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, chunk_dim: str | None = "win") -> AxisArray: + """`(win, time, ch)` -- what a windowing stage emits. + + `time` is the *within-window* lag dimension. Both it and `win` are + LinearAxes, so nothing distinguishes them but `chunk_dim`. + """ + kwargs = {"chunk_dim": chunk_dim} if chunk_dim else {} + return AxisArray( + np.zeros((n_win, n_lag, n_ch), np.float32), + dims=["win", "time", "ch"], + axes={ + "win": AxisArray.TimeAxis(fs=10.0), + "time": AxisArray.TimeAxis(fs=100.0), + "ch": CoordinateAxis(data=np.array(["a", "b", "c"]), dims=["ch"]), + }, + key="dev", + **kwargs, + ) + + +def _plain_msg(n_time: int = 20, n_ch: int = 3, chunk_dim: str | None = "time") -> AxisArray: + kwargs = {"chunk_dim": chunk_dim} if chunk_dim else {} + return AxisArray( + np.zeros((n_time, n_ch), np.float32), + dims=["time", "ch"], + axes={ + "time": AxisArray.TimeAxis(fs=100.0), + "ch": CoordinateAxis(data=np.array(["a", "b", "c"]), dims=["ch"]), + }, + key="dev", + **kwargs, + ) + + +def _sink(axis=None): + unit = ShMemCircBuff(ShMemCircBuffSettings(shmem_name=None, buf_dur=1.0, axis=axis)) + unit.STATE = ShMemCircBuff.STATE() + return unit + + +class TestTheBufferedAxisFollowsTheMessage: + """The ring is a history of the stream, so it has to be the dimension + messages accumulate along. `"time"` is present in a windowed message but is + the wrong one, so nothing rejected the old default -- it just buffered the + within-window samples and reported a 10x wrong sample rate. + """ + + def test_a_windowed_stream_resolves_to_win(self): + assert _sink()._resolve_axis(_windowed_msg()) == "win" + + def test_a_plain_stream_resolves_to_time(self): + assert _sink()._resolve_axis(_plain_msg()) == "time" + + def test_an_explicit_setting_still_wins(self): + """For a producer that declares nothing, an operator must still be able + to say which dimension to buffer.""" + assert _sink(axis="time")._resolve_axis(_windowed_msg()) == "time" + + def test_an_undeclared_producer_falls_back_to_time(self): + """Nothing better is available. A windowed producer that declares no + `chunk_dim` still gets the old, wrong answer -- the fix is for it to + declare one, which every ezmsg source now does.""" + assert _sink()._resolve_axis(_plain_msg(chunk_dim=None)) == "time" + assert _sink()._resolve_axis(_windowed_msg(chunk_dim=None)) == "time" + + def test_a_message_with_neither_is_skipped(self): + msg = AxisArray( + np.zeros((4, 3), np.float32), + dims=["freq", "ch"], + axes={"freq": AxisArray.LinearAxis(gain=1.0)}, + key="dev", + ) + assert _sink()._resolve_axis(msg) is None + + def test_what_the_old_default_did_to_a_windowed_stream(self): + """Buffering `time` puts the window *count* inside the frame, so the + buffer is reallocated whenever the window count jitters, and the rate + reported to the viewer is the within-window rate.""" + for axis, want_frame, want_rate in (("time", (4, 3), 100.0), ("win", (10, 3), 10.0)): + msg = _windowed_msg(n_win=4) + ax_idx = msg.get_axis_idx(axis) + frame_shape = msg.data.shape[:ax_idx] + msg.data.shape[ax_idx + 1 :] + assert frame_shape == want_frame + assert 1 / msg.axes[axis].gain == want_rate + + # ...and the window count is not stable, so `time` reshapes the frame. + shapes = set() + for n_win in (4, 7, 5): + msg = _windowed_msg(n_win=n_win) + ax_idx = msg.get_axis_idx("time") + shapes.add(msg.data.shape[:ax_idx] + msg.data.shape[ax_idx + 1 :]) + assert len(shapes) == 3, "frame shape must vary with window count when buffering `time`"