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
14 changes: 9 additions & 5 deletions .github/workflows/python-tests.yml
Original file line number Diff line number Diff line change
@@ -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:
Expand All @@ -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 }}
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,8 @@ 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",
"numpy>=2.0.2",
"pyxdf>=1.16.8",
]
Expand Down
61 changes: 27 additions & 34 deletions src/ezmsg/xdf/iter.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from pathlib import Path
import queue
from pathlib import Path

import numpy as np
import numpy.typing as npt
Expand All @@ -12,8 +12,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,
Expand Down Expand Up @@ -60,9 +59,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()

Expand All @@ -79,9 +76,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
Expand Down Expand Up @@ -175,9 +170,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)
Expand Down Expand Up @@ -212,25 +205,28 @@ def __init__(self, *args, select: str, **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
)
time_ax = AxisArray.TimeAxis(fs=self._metadata[_sel]["nominal_srate"], offset=0)
else:
time_ax = AxisArray.CoordinateAxis(
data=np.array([]),
dims=["time"],
unit="s"
)
time_ax = 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
self._template = AxisArray(
data=np.zeros(
(0, len(labels)), dtype=self._streams[0]["time_series"].dtype
),
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"]),
"ch": ch_ax,
},
key=self._streams[0]["info"]["name"][0],
# Messages accumulate along `time`, whether the stream is regular or
# carries per-sample timestamps; `ch` describes the stream itself.
chunk_dim="time",
)

def __next__(self) -> AxisArray:
Expand Down Expand Up @@ -285,16 +281,17 @@ def __init__(self, *args, force_single_sample: set = set(), **kwargs):
if fs
else AxisArray.CoordinateAxis(data=np.array([]), dims=["time"], unit="s")
)
ch_ax = AxisArray.CoordinateAxis(data=np.array(labels), dims=["ch"])
ch_ax.fingerprint # primed once per stream -- see the single-stream iterator above
self._templates[stream_name] = AxisArray(
data=np.zeros(
(0, stream_meta["channel_count"]), dtype=stream["time_series"].dtype
),
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"]),
"ch": ch_ax,
},
key=stream_name,
chunk_dim="time",
)
self._pubqueue: queue.SimpleQueue[AxisArray] = queue.SimpleQueue()

Expand All @@ -320,19 +317,15 @@ def __next__(self) -> AxisArray | None:
data=data[ix : ix + 1],
axes={
**template.axes,
"time": replace(
template.axes["time"], **t_kwargs
),
"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
}
t_kwargs = {"offset": tvec[0] if len(tvec) else self._last_time}
self._pubqueue.put_nowait(
replace(
template,
Expand Down
10 changes: 3 additions & 7 deletions src/ezmsg/xdf/source.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ class XDFIteratorSettings(ez.Settings):
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.
as during prototyping and testing.
"""


Expand Down Expand Up @@ -102,9 +102,7 @@ async def pub_chunk(self) -> typing.AsyncGenerator:
else:
await asyncio.sleep(0)
except StopIteration:
ez.logger.debug(
f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted."
)
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
Expand Down Expand Up @@ -152,9 +150,7 @@ async def pub_multi(self) -> typing.AsyncGenerator:
else:
await asyncio.sleep(0)
except StopIteration:
ez.logger.debug(
f"File ({self.SETTINGS.filepath} :: {self.SETTINGS.select}) exhausted."
)
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
26 changes: 26 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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")
Loading
Loading