diff --git a/.github/workflows/python-tests.yml b/.github/workflows/python-tests.yml index 22602eb..57abf79 100644 --- a/.github/workflows/python-tests.yml +++ b/.github/workflows/python-tests.yml @@ -1,10 +1,12 @@ name: Test package on: -# push: -# branches: [main] -# pull_request: -# branches: [main] + push: + branches: [main] + pull_request: + branches: + - main + - dev workflow_dispatch: jobs: @@ -25,7 +27,9 @@ jobs: uses: astral-sh/setup-uv@v3 with: enable-cache: true - cache-dependency-glob: "uv.lock" + # No uv.lock is committed here (nor in any sibling ezmsg package), + # so keying the cache on it matches nothing and setup-uv fails outright. + cache-dependency-glob: "pyproject.toml" - name: Set up Python ${{ matrix.python-version }} run: uv python install ${{ matrix.python-version }} diff --git a/pyproject.toml b/pyproject.toml index 150815a..45243a7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,9 @@ authors = [ requires-python = ">=3.10.15" dynamic = ["version"] dependencies = [ - "ezmsg>=3.6.0", + # 3.10.0b2 for AxisArray.chunk_dim and CoordinateAxis.fingerprint. + "ezmsg>=3.10.0b2", + "ezmsg-baseproc>=1.11.0", "numpy>=2.0.2", "pyxdf>=1.16.8", ] @@ -62,7 +64,7 @@ select = ["E", "F", "I", "W"] [tool.ruff.lint.isort] known-first-party = ["ezmsg.xdf"] -known-third-party = ["ezmsg"] +known-third-party = ["ezmsg", "ezmsg.baseproc"] [tool.uv.sources] # Uncomment to use development version of ezmsg from git diff --git a/src/ezmsg/xdf/iter.py b/src/ezmsg/xdf/iter.py index d683698..e204351 100644 --- a/src/ezmsg/xdf/iter.py +++ b/src/ezmsg/xdf/iter.py @@ -1,9 +1,16 @@ -from pathlib import Path +import asyncio +import os import queue +import typing +from dataclasses import field +from pathlib import Path +import ezmsg.core as ez import numpy as np import numpy.typing as npt import pyxdf +from ezmsg.baseproc.protocols import processor_state +from ezmsg.baseproc.stateful import BaseStatefulProducer from ezmsg.util.messages.axisarray import AxisArray from ezmsg.util.messages.util import replace @@ -12,8 +19,7 @@ class XDFIterator: def __init__( self, filepath: Path | str, - select: set[str] - | None = None, # If set, then the iterator yields only AxisArray of selected stream(s). + select: set[str] | None = None, # If set, then the iterator yields only AxisArray of selected stream(s). # If None (default), then the iterator yields dicts with keys for each stream chunk_dur: float = 1.0, # Attempt to chunk data into chunks of this duration. start_time: float | None = None, @@ -60,9 +66,7 @@ def __init__( self._chunk_ix = 0 self._last_time = 0.0 self._metadata = {} - self._prev_file_read_s: float = ( - 0 # File read header in seconds for previous iteration - ) + self._prev_file_read_s: float = 0 # File read header in seconds for previous iteration self._time_range: tuple[float | None, float | None] = (start_time, stop_time) self._scan_file() @@ -79,9 +83,7 @@ def _scan_file(self): # Load xdf self._streams, fileheader = pyxdf.load_xdf( self._filepath, - select_streams=None - if (self._select is None or self._rezero) - else [{"name": _} for _ in self._select], + select_streams=None if (self._select is None or self._rezero) else [{"name": _} for _ in self._select], ) self._metadata = {} self._file_read_s = 0 @@ -148,7 +150,7 @@ def _scan_file(self): self._streams = [self._streams[stream_names.index(_)] for _ in self._select] self._metadata = {k: self._metadata[k] for k in self._select} - print( + ez.logger.info( f"Imported {len(self._streams)} streams from {self._filepath} " f"spanning {xdf_dur:.2f} s beginning at t={xdf_t0:.2f}." ) @@ -161,12 +163,17 @@ def stream_meta(self) -> list[dict] | dict: def n_chunks(self) -> int: return self._n_chunks + @property + def exhausted(self) -> bool: + """True once every chunk boundary has been handed out.""" + return self._chunk_ix >= self._n_chunks + def __iter__(self): self._chunk_ix = 0 return self def __next__(self) -> dict[str, tuple[npt.NDArray, npt.NDArray]]: - if self._chunk_ix >= self.n_chunks: + if self.exhausted: raise StopIteration else: out_dict = {} @@ -175,9 +182,7 @@ def __next__(self) -> dict[str, tuple[npt.NDArray, npt.NDArray]]: (self._chunk_ix + 1) * self._chunk_dur + self._t0, ) for strm in self._streams: - b_chunk = np.logical_and( - strm["time_stamps"] >= t_start, strm["time_stamps"] < t_stop - ) + b_chunk = np.logical_and(strm["time_stamps"] >= t_start, strm["time_stamps"] < t_stop) out_tvec = strm["time_stamps"][b_chunk] out_data = strm["time_series"][b_chunk] out_dict[strm["info"]["name"][0]] = (out_data, out_tvec) @@ -197,156 +202,236 @@ def labels_from_strm(strm: dict) -> list[str]: return labels -class XDFAxisArrayIterator(XDFIterator): - def __init__(self, *args, select: str, **kwargs): - """ - This Iterator loads only a single stream and yields a single :obj:`AxisArray` object per chunk. +def _build_template(stream: dict, name: str, n_ch: int, fs: float) -> AxisArray: + """The message every chunk of *stream* is a `replace` of. - Args: - *args: - select: Unlike :obj:`XDFIterator`, this must be a single string, the name of the stream to select. - **kwargs: + Built once per stream so the `ch` axis object -- and the fingerprint cached + on it -- is shared by every message, which is what makes priming cheap. + """ + labels = labels_from_strm(stream) + time_ax = ( + AxisArray.TimeAxis(fs=fs, offset=0.0) + if fs + else AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s") + ) + ch_ax = AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]) + # Compute the channel fingerprint once, now. It is cached on the axis and + # pickled with it, and every message from this stream reuses this same axis + # object, so one checksum covers the whole file. Left cold it would be + # computed by the first stateful consumer in this process -- and, since + # unpickling builds a new axis object per message, by the first consumer in + # every other process, on every message. + ch_ax.fingerprint + return AxisArray( + data=np.zeros((0, n_ch), dtype=stream["time_series"].dtype), + dims=["time", "ch"], + axes={"time": time_ax, "ch": ch_ax}, + key=name, + # Messages accumulate along `time`, whether the stream is regular or + # carries per-sample timestamps; `ch` describes the stream itself. + chunk_dim="time", + ) + + +def _with_time(template: AxisArray, data: npt.NDArray, tvec: npt.NDArray, fallback_t: float) -> AxisArray: + """A chunk message: the template's data replaced, and its time axis advanced. + + An irregular stream carries every timestamp; a regular one carries only where + the chunk starts, since its gain says the rest. + """ + time_ax = template.axes["time"] + if isinstance(time_ax, AxisArray.CoordinateAxis): + t_kwargs = {"data": tvec if len(tvec) else np.array([])} + else: + t_kwargs = {"offset": tvec[0] if len(tvec) else fallback_t} + return replace( + template, + data=data, + axes={**template.axes, "time": replace(time_ax, **t_kwargs)}, + ) + + +class XDFIteratorSettings(ez.Settings): + """Settings shared by both AxisArray iterators. + + ``playback_rate`` and ``self_terminating`` belong to the unit rather than to + the reader, and are listed in :attr:`NONRESET_SETTINGS_FIELDS` so changing + either does not reopen the file. + """ + + filepath: typing.Union[os.PathLike, str] + select: str = "" + chunk_dur: float = 1.0 + start_time: float | None = None + stop_time: float | None = None + rezero: bool = True + playback_rate: float | None = None + self_terminating: bool = False + """ + If True, the unit will raise a :obj:`ez.NormalTermination` exception when the file is exhausted. + Note, however, that this will terminate the pipeline even if the data published by this unit are still in transit, + which will lead to the pipeline output being truncated before it has finished processing the stream. + `self_terminating` should only be used when it is not important that the pipeline finish processing data, such + as during prototyping and testing. + """ + + +class XDFMultiIteratorSettings(XDFIteratorSettings): + select: set[str] | None = None + force_single_sample: set = field(default_factory=set) + + +@processor_state +class XDFIteratorState: + reader: XDFIterator | None = None + template: AxisArray | None = None + + +@processor_state +class XDFMultiIteratorState: + reader: XDFIterator | None = None + templates: dict | None = None + pubqueue: queue.SimpleQueue | None = None + + +class _XDFProducerBase: + """Shared plumbing for the two AxisArray producers. + + The file load is deliberately *not* run from ``__init__``. ezmsg-neo and + ezmsg-nwb both reset eagerly there and so pay the whole open on the event + loop during ``initialize``; here the first ``__acall__`` triggers + ``_areset_state``, which puts it on a worker thread. Nothing in this + package's public surface reads stream metadata before the first chunk, so + there is nothing to lose by waiting. + """ + + NONRESET_SETTINGS_FIELDS = frozenset({"playback_rate", "self_terminating"}) + + async def _areset_state(self) -> None: + """Offload the sync open onto a worker thread. + + ``pyxdf.load_xdf`` reads and decodes the entire file, which for a + several-hundred-megabyte recording is seconds of pure CPU and I/O. On the + event loop that stalls every other unit in the process. """ - kwargs["select"] = set((select,)) - super().__init__(*args, **kwargs) - _sel = [_ for _ in self._select][0] - labels = labels_from_strm(self._streams[0]) - if self._metadata[_sel].get("nominal_srate", None): - time_ax = AxisArray.TimeAxis( - fs=self._metadata[_sel]["nominal_srate"], offset=0 - ) - else: - time_ax = AxisArray.CoordinateAxis( - data=np.array([]), - dims=["time"], - unit="s" - ) - self._template = AxisArray( - data=np.zeros( - (0, len(labels)), dtype=self._streams[0]["time_series"].dtype - ), - dims=["time", "ch"], - axes={ - "time": time_ax, - "ch": AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]), - }, - key=self._streams[0]["info"]["name"][0], + await asyncio.to_thread(self._reset_state) + + def _build_reader(self, select: set[str] | None) -> XDFIterator: + return XDFIterator( + filepath=self.settings.filepath, + select=select, + chunk_dur=self.settings.chunk_dur, + start_time=self.settings.start_time, + stop_time=self.settings.stop_time, + rezero=self.settings.rezero, + ) + + +class XDFAxisArrayIterator( + _XDFProducerBase, + BaseStatefulProducer[XDFIteratorSettings, AxisArray, XDFIteratorState], +): + """Loads a single stream and produces one :obj:`AxisArray` per chunk. + + ``select`` must be a single stream name, unlike :obj:`XDFIterator`. + """ + + @property + def exhausted(self) -> bool: + reader = self._state.reader + return reader is not None and reader.exhausted + + def _reset_state(self) -> None: + reader = self._build_reader({self.settings.select}) + meta = reader.stream_meta[self.settings.select] + self._state.reader = reader + self._state.template = _build_template( + reader._streams[0], + name=reader._streams[0]["info"]["name"][0], + n_ch=meta["channel_count"], + fs=meta["nominal_srate"], ) + async def _produce(self) -> AxisArray | None: + reader = self._state.reader + try: + chunk_dict = next(reader) + except StopIteration: + return None + data, tvec = chunk_dict.get(self.settings.select, (None, None)) + if data is None: + return None + return _with_time(self._state.template, data, tvec, reader._last_time) + def __next__(self) -> AxisArray: - result: AxisArray | None = None - chunk_dict = super().__next__() - # Should only be 1 in self._select. If there are more then we overwrite with the last. - for strm_name in self._select: - if strm_name in chunk_dict: - data, tvec = chunk_dict[strm_name] - if isinstance(self._template.axes["time"], AxisArray.CoordinateAxis): - t_kwargs = {"data": tvec} - else: - t_kwargs = {"offset": tvec[0] if len(tvec) else self._last_time} - result = replace( - self._template, - data=data, - axes={ - **self._template.axes, - "time": replace( - self._template.axes["time"], - **t_kwargs, - ), - }, - ) + result = self() + if result is None: + raise StopIteration return result -class XDFMultiAxArrIterator(XDFIterator): - def __init__(self, *args, force_single_sample: set = set(), **kwargs): - """ - This Iterator loads multiple streams and yields a :obj:`AxisArray` object per iteration, - but the stream source might different between chunks. +class XDFMultiAxArrIterator( + _XDFProducerBase, + BaseStatefulProducer[XDFMultiIteratorSettings, AxisArray, XDFMultiIteratorState], +): + """Loads multiple streams and produces one :obj:`AxisArray` per iteration. - Args: - *args: - force_single_sample: Use this to identify irregular-rate streams that might conceivably have more than one - event within the defined chunk_dur, for which :obj:`AxisArray` cannot represent timestamps properly. - **kwargs: - """ - super().__init__(*args, **kwargs) - self._force_single_sample = force_single_sample - stream_names = [_["info"]["name"][0] for _ in self._streams] - - # Create template messages for each stream - self._templates = {} - for stream_name, stream_meta in self._metadata.items(): - stream = self._streams[stream_names.index(stream_name)] - labels = labels_from_strm(stream) - fs = stream_meta["nominal_srate"] - time_ax = ( - AxisArray.TimeAxis(fs=fs, offset=0.0) - if fs - else AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s") - ) - self._templates[stream_name] = AxisArray( - data=np.zeros( - (0, stream_meta["channel_count"]), dtype=stream["time_series"].dtype - ), - dims=["time", "ch"], - axes={ - "time": time_ax, - "ch": AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"]), - }, - key=stream_name, + Which stream a given message came from varies; read ``.key``. Returns + ``None`` when a chunk held nothing for any stream, and raises + ``StopIteration`` only once the file is done. + + ``force_single_sample`` names irregular-rate streams that may carry more than + one event within ``chunk_dur``, which :obj:`AxisArray` cannot represent as a + single message with correct timestamps; those are split one event per + message. + """ + + @property + def exhausted(self) -> bool: + reader = self._state.reader + if reader is None: + return False + return reader.exhausted and self._state.pubqueue.empty() + + def _reset_state(self) -> None: + reader = self._build_reader(self.settings.select) + stream_names = [_["info"]["name"][0] for _ in reader._streams] + self._state.reader = reader + self._state.pubqueue = queue.SimpleQueue() + self._state.templates = { + name: _build_template( + reader._streams[stream_names.index(name)], + name=name, + n_ch=meta["channel_count"], + fs=meta["nominal_srate"], ) - self._pubqueue: queue.SimpleQueue[AxisArray] = queue.SimpleQueue() + for name, meta in reader.stream_meta.items() + } - def __next__(self) -> AxisArray | None: - if self._pubqueue.empty(): - chunk_dict = super().__next__() - for k, template in self._templates.items(): - if k in chunk_dict and len(chunk_dict[k][1]) > 0: - data, tvec = chunk_dict[k] - if k in self._force_single_sample: - if isinstance(template.axes["time"], AxisArray.CoordinateAxis): - t_kwargs = {"data": np.array([])} - else: - t_kwargs = {"offset": 0.0} - for ix, _t in enumerate(tvec): - if "data" in t_kwargs: - t_kwargs["data"] = np.array([_t]) - else: - t_kwargs["offset"] = _t - self._pubqueue.put_nowait( - replace( - template, - data=data[ix : ix + 1], - axes={ - **template.axes, - "time": replace( - template.axes["time"], **t_kwargs - ), - }, - ) - ) - else: - if isinstance(template.axes["time"], AxisArray.CoordinateAxis): - t_kwargs = {"data": tvec if len(tvec) else np.array([])} - else: - t_kwargs = { - "offset": tvec[0] if len(tvec) else self._last_time - } - self._pubqueue.put_nowait( - replace( - template, - data=data, - axes={ - **template.axes, - "time": replace( - template.axes["time"], - **t_kwargs, - ), - }, - ) - ) + def _enqueue_chunk(self, chunk_dict: dict) -> None: + reader = self._state.reader + for name, template in self._state.templates.items(): + if name not in chunk_dict or len(chunk_dict[name][1]) == 0: + continue + data, tvec = chunk_dict[name] + if name in self.settings.force_single_sample: + for ix, stamp in enumerate(tvec): + self._state.pubqueue.put_nowait(_with_time(template, data[ix : ix + 1], np.array([stamp]), stamp)) + else: + self._state.pubqueue.put_nowait(_with_time(template, data, tvec, reader._last_time)) + + async def _produce(self) -> AxisArray | None: + if self._state.pubqueue.empty(): + try: + self._enqueue_chunk(next(self._state.reader)) + except StopIteration: + return None try: - return self._pubqueue.get_nowait() + return self._state.pubqueue.get_nowait() except queue.Empty: return None + + def __next__(self) -> AxisArray | None: + if self.exhausted: + raise StopIteration + return self() diff --git a/src/ezmsg/xdf/source.py b/src/ezmsg/xdf/source.py index 0228384..bc4a12e 100644 --- a/src/ezmsg/xdf/source.py +++ b/src/ezmsg/xdf/source.py @@ -1,13 +1,30 @@ import asyncio -import os import time import typing -from dataclasses import field import ezmsg.core as ez +from ezmsg.baseproc.units import BaseProducerUnit from ezmsg.util.messages.axisarray import AxisArray -from .iter import XDFAxisArrayIterator, XDFMultiAxArrIterator +from .iter import ( + XDFAxisArrayIterator, + XDFIteratorSettings, + XDFMultiAxArrIterator, + XDFMultiIteratorSettings, +) + +# The settings types moved to `iter.py` so the producers can own them, but they +# were importable from here first. +XDFMultiIteratorUnitSettings = XDFMultiIteratorSettings + +__all__ = [ + "PlaybackClock", + "XDFIteratorSettings", + "XDFIteratorUnit", + "XDFMultiIteratorSettings", + "XDFMultiIteratorUnit", + "XDFMultiIteratorUnitSettings", +] class PlaybackClock: @@ -47,113 +64,75 @@ def step(self) -> None: time.sleep(self._get_duration()) -class XDFIteratorSettings(ez.Settings): - filepath: typing.Union[os.PathLike, str] - select: str - chunk_dur: float = 1.0 - start_time: float | None = None - stop_time: float | None = None - rezero: bool = True - playback_rate: float | None = None - self_terminating: bool = False - """ - If True, the unit will raise a :obj:`ez.NormalTermination` exception when the file is exhausted. - Note, however, that this will terminate the pipeline even if the data published by this unit are still in transit, - which will lead to the pipeline output being truncated before it has finished processing the stream. - `self_terminating` should only be used when it is not important that the pipeline finish processing data, such - as during prototyping and testing. +class _XDFUnitBase: + """Playback pacing and end-of-file handling, shared by both units. + + Note both subclasses name their publisher ``produce``: that is the name + ``BaseProducerUnit`` uses, and ezmsg collects publishers per attribute, so a + differently named one would run *alongside* the base class\'s rather than + replacing it -- two publishers draining one producer, neither stopping. + + The producer supplies chunks as fast as they can be sliced out of memory; + ``playback_rate`` is what turns that into a paced stream, and it is a + property of the unit rather than of the reader. """ + OUTPUT_TERM = ez.OutputStream(typing.Any) + + async def initialize(self) -> None: + await super().initialize() + self._clock = ( + PlaybackClock(rate=self.SETTINGS.playback_rate, step_dur=self.SETTINGS.chunk_dur) + if self.SETTINGS.playback_rate is not None + else None + ) -class XDFIteratorState(ez.State): - gen: typing.Any = None + async def _finish(self) -> typing.AsyncGenerator: + ez.logger.debug(f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted.") + if self.SETTINGS.self_terminating: + raise ez.NormalTermination + yield self.OUTPUT_TERM, True -class XDFIteratorUnit(ez.Unit): - STATE = XDFIteratorState +class XDFIteratorUnit( + _XDFUnitBase, + BaseProducerUnit[XDFIteratorSettings, AxisArray, XDFAxisArrayIterator], +): SETTINGS = XDFIteratorSettings OUTPUT_SIGNAL = ez.OutputStream(AxisArray) - OUTPUT_TERM = ez.OutputStream(typing.Any) - - def initialize(self) -> None: - self.construct_generator() - - def construct_generator(self): - self.STATE.gen = XDFAxisArrayIterator( - filepath=self.SETTINGS.filepath, - select=self.SETTINGS.select, - chunk_dur=self.SETTINGS.chunk_dur, - start_time=self.SETTINGS.start_time, - stop_time=self.SETTINGS.stop_time, - rezero=self.SETTINGS.rezero, - ) - if self.SETTINGS.playback_rate is not None: - self._clock = PlaybackClock(rate=self.SETTINGS.playback_rate, step_dur=self.SETTINGS.chunk_dur) - else: - self._clock = None @ez.publisher(OUTPUT_SIGNAL) - async def pub_chunk(self) -> typing.AsyncGenerator: - try: - while True: - if self._clock is not None: - await self._clock.astep() - msg = next(self.STATE.gen) - if msg.data.size > 0: - yield self.OUTPUT_SIGNAL, msg - else: - await asyncio.sleep(0) - except StopIteration: - ez.logger.debug(f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted.") - if self.SETTINGS.self_terminating: - raise ez.NormalTermination - yield self.OUTPUT_TERM, True - - -class XDFMultiIteratorUnitSettings(XDFIteratorSettings): - select: set[str] | None = None # Override with a default - force_single_sample: set = field(default_factory=set) - - -class XDFMultiIteratorUnit(ez.Unit): - STATE = XDFIteratorState - SETTINGS = XDFMultiIteratorUnitSettings + async def produce(self) -> typing.AsyncGenerator: + while not self.producer.exhausted: + if self._clock is not None: + await self._clock.astep() + msg = await self.producer.__acall__() + if msg is not None and msg.data.size > 0: + yield self.OUTPUT_SIGNAL, msg + else: + await asyncio.sleep(0) + async for out in self._finish(): + yield out + + +class XDFMultiIteratorUnit( + _XDFUnitBase, + BaseProducerUnit[XDFMultiIteratorSettings, AxisArray, XDFMultiAxArrIterator], +): + SETTINGS = XDFMultiIteratorSettings OUTPUT_SIGNAL = ez.OutputStream(AxisArray) - OUTPUT_TERM = ez.OutputStream(typing.Any) - - def initialize(self) -> None: - self.construct_generator() - - def construct_generator(self): - self.STATE.gen = XDFMultiAxArrIterator( - filepath=self.SETTINGS.filepath, - select=self.SETTINGS.select, - chunk_dur=self.SETTINGS.chunk_dur, - start_time=self.SETTINGS.start_time, - stop_time=self.SETTINGS.stop_time, - rezero=self.SETTINGS.rezero, - force_single_sample=self.SETTINGS.force_single_sample, - ) - if self.SETTINGS.playback_rate is not None: - self._clock = PlaybackClock(rate=self.SETTINGS.playback_rate, step_dur=self.SETTINGS.chunk_dur) - else: - self._clock = None @ez.publisher(OUTPUT_SIGNAL) - async def pub_multi(self) -> typing.AsyncGenerator: - try: - while True: - if self._clock is not None: - await self._clock.astep() - msg = next(self.STATE.gen) - if msg is not None: - yield self.OUTPUT_SIGNAL, msg - else: - await asyncio.sleep(0) - except StopIteration: - ez.logger.debug(f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted.") - if self.SETTINGS.self_terminating: - raise ez.NormalTermination - yield self.OUTPUT_TERM, True + async def produce(self) -> typing.AsyncGenerator: + while not self.producer.exhausted: + if self._clock is not None: + await self._clock.astep() + msg = await self.producer.__acall__() + if msg is not None: + yield self.OUTPUT_SIGNAL, msg + else: + await asyncio.sleep(0) + async for out in self._finish(): + yield out diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..7e8d496 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,26 @@ +"""Shared fixtures for the ezmsg-xdf tests. + +The XDF file is generated rather than checked in as a binary -- see +``create_test_xdf.py`` for why, and for what it contains. +""" + +import sys +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).parent)) + +from create_test_xdf import DEFAULT_STREAMS, StreamSpec, write_test_xdf # noqa: E402 + +# Kept alongside the fixture so tests can assert against them by name rather +# than by index into DEFAULT_STREAMS. +EEG_STREAM: StreamSpec = DEFAULT_STREAMS[0] +MARKER_STREAM: StreamSpec = DEFAULT_STREAMS[1] + + +@pytest.fixture(scope="session") +def test_xdf_path(tmp_path_factory) -> Path: + """A two-stream XDF: a 100 Hz 4-channel float32 stream and an irregular + string marker stream, both starting at t=10 s.""" + return write_test_xdf(tmp_path_factory.mktemp("xdf") / "test.xdf") diff --git a/tests/create_test_xdf.py b/tests/create_test_xdf.py new file mode 100644 index 0000000..cf9cd54 --- /dev/null +++ b/tests/create_test_xdf.py @@ -0,0 +1,281 @@ +"""Write a small, self-describing XDF file for the tests. + +Checked in as a generator rather than a binary blob so the fixture is readable +and adjustable: a test that needs an irregular stream, a string stream or a +particular channel layout changes an argument here instead of asking someone to +produce a new recording. + +The format is the XDF specification's chunked layout, and the field-by-field +choices below are pinned against what ``pyxdf``'s reader actually consumes -- +:func:`pyxdf.pyxdf._read_varlen_int`, :func:`~pyxdf.pyxdf._read_chunk3` and the +tag dispatch in :func:`~pyxdf.load_xdf` -- since that is the only reader these +files ever meet. + +Run standalone to drop a file next to this script:: + + python tests/create_test_xdf.py /tmp/test.xdf +""" + +from __future__ import annotations + +import struct +import typing +from pathlib import Path + +import numpy as np + +# XDF chunk tags. Boundary chunks (5) are omitted -- they only help a reader +# resynchronise after corruption, and these files are not corrupt. ClockOffset +# chunks (4) are written even though there is no clock skew to model, because +# pyxdf warns ("Segments and clock-segments differ") on a stream that has sample +# segments but no clock segments, and a fixture that logs a warning on every load +# trains readers to ignore warnings. +TAG_FILE_HEADER = 1 +TAG_STREAM_HEADER = 2 +TAG_SAMPLES = 3 +TAG_CLOCK_OFFSET = 4 +TAG_STREAM_FOOTER = 6 + +FORMAT_DTYPES: dict[str, np.dtype] = { + "int8": np.dtype(" bytes: + """XDF's length prefix: one byte saying how wide the count is, then the count. + + The reader accepts widths of 1, 4 and 8 only, so this picks the narrowest of + those rather than the narrowest possible. + """ + if value < 256: + return b"\x01" + bytes([value]) + if value < 2**32: + return b"\x04" + struct.pack(" bytes: + """One length-prefixed chunk. + + The declared length covers the tag and, where the tag has one, the stream id + -- the reader subtracts both back off before reading the payload. + """ + body = struct.pack(" str: + rows = "".join( + f"{unit}{ch_type}" for label in labels + ) + return f"{rows}" + + +def _stream_header_xml( + name: str, + stream_type: str, + labels: typing.Sequence[str], + srate: float, + channel_format: str, + unit: str, + ch_type: str, +) -> bytes: + return ( + '' + f"{name}" + f"{stream_type}" + f"{len(labels)}" + f"{srate:g}" + f"{channel_format}" + f"0.0" + f"{name}-uid" + f"default" + f"test" + f"{_xml_channels(labels, unit, ch_type)}" + "" + ).encode("utf-8") + + +def _stream_footer_xml(first: float, last: float, n_samples: int, srate: float) -> bytes: + return ( + '' + f"{first!r}" + f"{last!r}" + f"{n_samples}" + f"{srate:g}" + "" + ).encode("utf-8") + + +def _clock_offset_chunk(stream_id: int, collection_time: float, offset: float) -> bytes: + return _chunk( + TAG_CLOCK_OFFSET, + struct.pack(" bytes: + """A [Samples] chunk with an explicit timestamp on every sample. + + Every sample carries its stamp rather than relying on delta decompression. + That is more bytes than a real recorder would write, and deliberately so: + the alternative encodes "same as last plus 1/srate", which would make a + fixture silently agree with any reader that got the nominal rate wrong. + """ + out = [_varlen_int(len(stamps))] + if channel_format == "string": + for row, stamp in zip(values, stamps): + out.append(b"\x08" + struct.pack(" np.ndarray: + if spec.channel_format == "string": + return np.array( + [[f"{spec.name}-{i}-{ch}" for ch in range(len(spec.labels))] for i in range(spec.n_samples)], + dtype=object, + ) + # A per-channel ramp offset by channel index: every sample is identifiable, + # so a test can assert on ordering and chunk boundaries, not just on shape. + ramp = np.arange(spec.n_samples, dtype=np.float64)[:, None] + offsets = np.arange(len(spec.labels), dtype=np.float64)[None, :] * 1000.0 + return (ramp + offsets).astype(FORMAT_DTYPES[spec.channel_format]) + + +def _stamps_for(spec: StreamSpec, rng: np.random.Generator) -> np.ndarray: + if spec.srate > 0: + return spec.t0 + np.arange(spec.n_samples) / spec.srate + # Irregular: monotonic but unevenly spaced, which is the whole point of the + # format's per-sample timestamps. + gaps = rng.uniform(0.05, 0.25, size=spec.n_samples) + return spec.t0 + np.cumsum(gaps) + + +DEFAULT_STREAMS: tuple[StreamSpec, ...] = ( + StreamSpec(name="EEGSignal", labels=("Fz", "Cz", "Pz", "Oz"), srate=100.0, n_samples=250), + StreamSpec( + name="Markers", + labels=("marker",), + srate=0.0, + n_samples=12, + channel_format="string", + stream_type="Markers", + unit="none", + ch_type="Marker", + ), +) + + +def write_test_xdf( + path: Path | str, + streams: typing.Sequence[StreamSpec] = DEFAULT_STREAMS, + samples_per_chunk: int = 32, + seed: int = 0, +) -> Path: + """Write *streams* to *path* and return it. + + Samples are split across several [Samples] chunks so the file exercises the + reader's chunk stitching rather than arriving as one block. + """ + rng = np.random.default_rng(seed) + path = Path(path) + + out = [b"XDF:"] + out.append( + _chunk( + TAG_FILE_HEADER, + b'1.0', + ) + ) + + prepared = [] + for stream_id, spec in enumerate(streams, start=1): + values = _values_for(spec, rng) + stamps = _stamps_for(spec, rng) + prepared.append((stream_id, spec, values, stamps)) + out.append( + _chunk( + TAG_STREAM_HEADER, + _stream_header_xml( + spec.name, spec.stream_type, spec.labels, spec.srate, spec.channel_format, spec.unit, spec.ch_type + ), + stream_id=stream_id, + ) + ) + + # Interleave the streams' sample chunks, as a real recording would have them. + offset = 0 + while any(offset < spec.n_samples for _, spec, _, _ in prepared): + for stream_id, spec, values, stamps in prepared: + if offset >= spec.n_samples: + continue + stop = min(offset + samples_per_chunk, spec.n_samples) + out.append( + _chunk( + TAG_SAMPLES, + _samples_chunk(values[offset:stop], stamps[offset:stop], spec.channel_format), + stream_id=stream_id, + ) + ) + offset += samples_per_chunk + + # A pair of zero-offset clock measurements per stream, bracketing its + # samples. Zero because these timestamps are already on one clock; the + # chunks exist so the reader sees a clock segment covering the data. + for stream_id, spec, _, stamps in prepared: + out.append(_clock_offset_chunk(stream_id, float(stamps[0]) - 1.0, 0.0)) + out.append(_clock_offset_chunk(stream_id, float(stamps[-1]) + 1.0, 0.0)) + + for stream_id, spec, _, stamps in prepared: + out.append( + _chunk( + TAG_STREAM_FOOTER, + _stream_footer_xml(float(stamps[0]), float(stamps[-1]), spec.n_samples, spec.srate), + stream_id=stream_id, + ) + ) + + path.write_bytes(b"".join(out)) + return path + + +if __name__ == "__main__": + import sys + + target = Path(sys.argv[1] if len(sys.argv) > 1 else "test.xdf") + write_test_xdf(target) + print(f"wrote {target} ({target.stat().st_size} bytes)") diff --git a/tests/test_iter.py b/tests/test_iter.py index f1a04f4..deb67c3 100644 --- a/tests/test_iter.py +++ b/tests/test_iter.py @@ -1,3 +1,306 @@ -def test_dummy(): - # TODO: Add tests. Requires a small XDF source. - pass +"""Behaviour of the XDF iterators, against a generated fixture. + +These pin what the iterators produce -- sample values, ordering, chunking, and +the two message fields consumers key their cached state on -- so that the +producer's contract is checked rather than merely exercised. +""" + +from __future__ import annotations + +import math +import pickle + +import ezmsg.core as ez +import numpy as np +import pytest +from conftest import EEG_STREAM, MARKER_STREAM +from ezmsg.util.messages.axisarray import AxisArray +from ezmsg.util.messages.util import replace as replace_settings + +from ezmsg.xdf.iter import ( + XDFAxisArrayIterator, + XDFIterator, + XDFIteratorSettings, + XDFMultiAxArrIterator, + XDFMultiIteratorSettings, +) +from ezmsg.xdf.source import XDFIteratorUnit, XDFMultiIteratorUnit + + +def eeg_messages(path, **kwargs) -> list[AxisArray]: + return list(XDFAxisArrayIterator(filepath=path, select=EEG_STREAM.name, **kwargs)) + + +class TestTheRawIterator: + def test_it_finds_both_streams(self, test_xdf_path): + it = XDFIterator(filepath=test_xdf_path, chunk_dur=1.0) + seen: set[str] = set() + for chunk in it: + seen |= set(chunk) + assert seen == {EEG_STREAM.name, MARKER_STREAM.name} + + +class TestTheSingleStreamIterator: + def test_it_yields_every_sample_in_order(self, test_xdf_path): + msgs = eeg_messages(test_xdf_path, chunk_dur=0.5) + data = np.concatenate([m.data for m in msgs], axis=0) + assert data.shape == (EEG_STREAM.n_samples, len(EEG_STREAM.labels)) + # create_test_xdf writes a per-channel ramp: sample i, channel c == i + 1000c. + expected = np.arange(EEG_STREAM.n_samples)[:, None] + np.arange(len(EEG_STREAM.labels))[None, :] * 1000.0 + np.testing.assert_allclose(data, expected) + + def test_chunk_dur_controls_how_much_arrives_at_once(self, test_xdf_path): + long = eeg_messages(test_xdf_path, chunk_dur=1.0) + short = eeg_messages(test_xdf_path, chunk_dur=0.25) + assert len(short) > len(long) + assert max(m.data.shape[0] for m in long) > max(m.data.shape[0] for m in short) + + def test_it_carries_the_channel_labels(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + assert list(msg.axes["ch"].data) == list(EEG_STREAM.labels) + + def test_rezero_moves_the_first_sample_to_zero(self, test_xdf_path): + """The fixture starts at t=10 s, so this distinguishes honouring the + setting from ignoring it.""" + rezeroed = eeg_messages(test_xdf_path, rezero=True)[0] + assert rezeroed.axes["time"].offset == pytest.approx(0.0, abs=1e-6) + + def test_without_rezero_the_file_clock_is_preserved(self, test_xdf_path): + raw = eeg_messages(test_xdf_path, rezero=False)[0] + assert raw.axes["time"].offset == pytest.approx(EEG_STREAM.t0, abs=1e-6) + + def test_the_nominal_rate_becomes_the_axis_gain(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + assert msg.axes["time"].gain == pytest.approx(1.0 / EEG_STREAM.srate) + + +class TestTheMultiStreamIterator: + @staticmethod + def _messages(path, **kwargs) -> list[AxisArray]: + it = XDFMultiAxArrIterator(filepath=path, chunk_dur=1.0, **kwargs) + return [msg for msg in it if msg is not None] + + def test_both_streams_come_through(self, test_xdf_path): + keys = {m.key for m in self._messages(test_xdf_path)} + assert keys == {EEG_STREAM.name, MARKER_STREAM.name} + + def test_the_irregular_stream_keeps_per_sample_timestamps(self, test_xdf_path): + markers = [m for m in self._messages(test_xdf_path) if m.key == MARKER_STREAM.name] + assert markers, "no marker messages" + time_ax = markers[0].axes["time"] + assert isinstance(time_ax, AxisArray.CoordinateAxis) + stamps = np.concatenate([m.axes["time"].data for m in markers]) + assert len(stamps) == MARKER_STREAM.n_samples + assert np.all(np.diff(stamps) > 0), "timestamps must stay monotonic" + + def test_force_single_sample_splits_an_irregular_stream(self, test_xdf_path): + """Without it, several events inside one ``chunk_dur`` arrive together.""" + batched = [m for m in self._messages(test_xdf_path) if m.key == MARKER_STREAM.name] + split = [ + m + for m in self._messages(test_xdf_path, force_single_sample={MARKER_STREAM.name}) + if m.key == MARKER_STREAM.name + ] + assert len(split) == MARKER_STREAM.n_samples + assert len(split) > len(batched) + assert all(m.data.shape[0] == 1 for m in split) + + +class TestMessagesArriveReadyForConsumers: + """Two things only the source can supply, both set once per stream. + + ``chunk_dim`` names the dimension messages accumulate along -- the one whose + length is just however much of the file this chunk covered, and which a + consumer must leave out of the state it caches against the stream's + configuration. ``fingerprint`` is the channel axis's content digest, cached + on the axis and pickled with it; priming it at construction spares the first + consumer in every process from recomputing it on every message. + """ + + def test_the_single_stream_iterator_declares_its_chunk_dim(self, test_xdf_path): + assert all(m.chunk_dim == "time" for m in eeg_messages(test_xdf_path)) + + def test_the_multi_stream_iterator_declares_it_for_every_stream(self, test_xdf_path): + it = XDFMultiAxArrIterator(filepath=test_xdf_path, chunk_dur=1.0) + undeclared = sorted({m.key for m in it if m is not None and m.chunk_dim != "time"}) + assert not undeclared, f"streams not declaring chunk_dim='time': {undeclared}" + + def test_the_channel_axis_is_primed(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + assert "_fingerprint" in msg.axes["ch"].__dict__ + assert msg.axes["ch"].fingerprint is not None + + def test_every_stream_of_the_multi_iterator_is_primed(self, test_xdf_path): + it = XDFMultiAxArrIterator(filepath=test_xdf_path, chunk_dur=1.0) + cold = sorted({m.key for m in it if m is not None and "_fingerprint" not in m.axes["ch"].__dict__}) + assert not cold, f"streams handing over a cold ch axis: {cold}" + + def test_one_axis_object_serves_the_whole_stream(self, test_xdf_path): + """What makes priming cheap: the checksum is paid once, not per message.""" + msgs = eeg_messages(test_xdf_path, chunk_dur=0.25) + assert len(msgs) > 1 + assert len({id(m.axes["ch"]) for m in msgs}) == 1 + + def test_the_chunk_axis_is_left_cold(self, test_xdf_path): + """Digesting per-message timestamps would be pure cost: no consumer reads + the chunk axis's fingerprint.""" + it = XDFMultiAxArrIterator(filepath=test_xdf_path, chunk_dur=1.0) + markers = [m for m in it if m is not None and m.key == MARKER_STREAM.name] + assert markers, "no marker messages" + assert all("_fingerprint" not in m.axes["time"].__dict__ for m in markers) + + def test_it_all_survives_the_transport(self, test_xdf_path): + msg = eeg_messages(test_xdf_path)[0] + landed = pickle.loads(pickle.dumps(msg)) + assert landed.chunk_dim == "time" + assert "_fingerprint" in landed.axes["ch"].__dict__ + assert landed.axes["ch"].__dict__["_fingerprint"] == msg.axes["ch"].fingerprint + + +class TestTheFixtureItself: + """The generator is test code, and a wrong fixture would make every + assertion above agree with the wrong thing.""" + + def test_pyxdf_reads_back_what_was_written(self, test_xdf_path): + import pyxdf + + streams, header = pyxdf.load_xdf(str(test_xdf_path)) + assert header["info"]["version"] == ["1.0"] + by_name = {s["info"]["name"][0]: s for s in streams} + assert set(by_name) == {EEG_STREAM.name, MARKER_STREAM.name} + + eeg = by_name[EEG_STREAM.name] + assert np.asarray(eeg["time_series"]).shape == (EEG_STREAM.n_samples, len(EEG_STREAM.labels)) + assert float(eeg["info"]["nominal_srate"][0]) == EEG_STREAM.srate + assert eeg["time_stamps"][0] == pytest.approx(EEG_STREAM.t0) + labels = [c["label"][0] for c in eeg["info"]["desc"][0]["channels"][0]["channel"]] + assert labels == list(EEG_STREAM.labels) + + markers = by_name[MARKER_STREAM.name] + assert float(markers["info"]["nominal_srate"][0]) == 0.0 + assert np.asarray(markers["time_series"]).shape == (MARKER_STREAM.n_samples, 1) + + def test_it_loads_without_pyxdf_warnings(self, test_xdf_path, caplog): + """A fixture that logs on every load trains readers to ignore warnings.""" + import logging + + import pyxdf + + with caplog.at_level(logging.WARNING, logger="pyxdf"): + pyxdf.load_xdf(str(test_xdf_path)) + assert not caplog.records, [r.getMessage() for r in caplog.records] + + def test_the_samples_span_more_than_one_chunk(self, test_xdf_path): + """Otherwise the reader's chunk stitching is never exercised.""" + assert EEG_STREAM.n_samples > 32 + assert math.ceil(EEG_STREAM.n_samples / 32) > 1 + + +class TestTheProducerContract: + """The producers are `BaseStatefulProducer`s, which means the file open is a + state reset rather than construction work.""" + + def test_construction_does_not_read_the_file(self, test_xdf_path): + """Deliberately unlike ezmsg-neo and ezmsg-nwb, which reset eagerly in + ``__init__`` and so pay the whole open on the event loop during + ``initialize``. Nothing here reads stream metadata before the first + chunk, so there is nothing to lose by waiting.""" + it = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name) + assert it._state.reader is None + assert next(it) is not None + assert it._state.reader is not None + + def test_the_file_load_runs_off_the_event_loop(self, test_xdf_path): + """``pyxdf.load_xdf`` reads and decodes the whole file. On the event loop + that stalls every other unit in the process. + + Driven with ``asyncio.run`` rather than an async test, because this repo + has no pytest-asyncio -- the pytest config names ``asyncio_mode`` but the + plugin is not installed, so an ``async def`` test is silently never + awaited and passes without running. + """ + import asyncio + import threading + + seen: list[int] = [] + + class Spy(XDFAxisArrayIterator): + def _reset_state(self): + seen.append(threading.get_ident()) + super()._reset_state() + + producer = Spy(filepath=test_xdf_path, select=EEG_STREAM.name) + assert not seen, "construction should not have opened anything" + + loop_tid: list[int] = [] + + async def drive(): + loop_tid.append(threading.get_ident()) + await producer.__acall__() + + asyncio.run(drive()) + + assert len(seen) == 1 + assert seen[0] != loop_tid[0], "_reset_state ran on the event-loop thread" + + def test_settings_arrive_as_keywords_or_as_a_settings_object(self, test_xdf_path): + by_kwargs = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name, chunk_dur=0.5) + by_settings = XDFAxisArrayIterator( + settings=XDFIteratorSettings(filepath=test_xdf_path, select=EEG_STREAM.name, chunk_dur=0.5) + ) + assert by_kwargs.settings == by_settings.settings + + def test_pacing_settings_do_not_reopen_the_file(self, test_xdf_path): + """``playback_rate`` and ``self_terminating`` belong to the unit, so + changing either must not throw away a loaded file.""" + it = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name) + next(it) + reader = it._state.reader + it.update_settings(replace_settings(it.settings, playback_rate=2.0, self_terminating=True)) + next(it) + assert it._state.reader is reader + + def test_changing_the_file_does_reopen(self, test_xdf_path): + it = XDFAxisArrayIterator(filepath=test_xdf_path, select=EEG_STREAM.name) + next(it) + reader = it._state.reader + it.update_settings(replace_settings(it.settings, chunk_dur=0.25)) + next(it) + assert it._state.reader is not reader + + +class TestTheUnitsInAGraph: + @staticmethod + def _run(unit_cls, settings) -> list[AxisArray]: + collected: list[AxisArray] = [] + + class Collector(ez.Unit): + INPUT_SIGNAL = ez.InputStream(AxisArray) + + @ez.subscriber(INPUT_SIGNAL) + async def on_msg(self, msg: AxisArray) -> None: + collected.append(msg) + + src, sink = unit_cls(settings), Collector() + ez.run(SRC=src, SINK=sink, connections=((src.OUTPUT_SIGNAL, sink.INPUT_SIGNAL),)) + return collected + + def test_the_single_stream_unit_publishes_the_whole_file(self, test_xdf_path): + msgs = self._run( + XDFIteratorUnit, + XDFIteratorSettings(filepath=test_xdf_path, select=EEG_STREAM.name, self_terminating=True), + ) + assert msgs, "no messages published" + assert sum(m.data.shape[0] for m in msgs) == EEG_STREAM.n_samples + assert all(m.chunk_dim == "time" for m in msgs) + assert all("_fingerprint" in m.axes["ch"].__dict__ for m in msgs) + + def test_the_multi_stream_unit_publishes_both_streams(self, test_xdf_path): + msgs = self._run( + XDFMultiIteratorUnit, + XDFMultiIteratorSettings(filepath=test_xdf_path, self_terminating=True), + ) + assert {m.key for m in msgs} == {EEG_STREAM.name, MARKER_STREAM.name} + eeg = [m for m in msgs if m.key == EEG_STREAM.name] + assert sum(m.data.shape[0] for m in eeg) == EEG_STREAM.n_samples + assert all(m.chunk_dim == "time" for m in msgs)