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
7 changes: 7 additions & 0 deletions .github/workflows/python-tests.yml
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,12 @@ on:
- dev
workflow_dispatch:

# A new push to a PR supersedes the run already in flight for that ref, so a
# wedged job cannot sit on a runner while its replacement queues behind it.
concurrency:
group: ${{ github.workflow }}-${{ github.ref }}
cancel-in-progress: true

jobs:
build:
strategy:
Expand All @@ -19,6 +25,7 @@ jobs:
- "windows-latest"
- "macos-latest"
runs-on: ${{matrix.os}}
timeout-minutes: 30

steps:
- uses: actions/checkout@v4
Expand Down
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -9,10 +9,10 @@ readme = "README.md"
requires-python = ">=3.11"
dynamic = ["version"]
dependencies = [
# 3.10.0b2 for AxisArray.chunk_dim, which the shmem sink follows to pick
# 3.10.0b3 for AxisArray.stream_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",
"ezmsg>=3.10.0b3",
"numpy>=1.26.0",
"typer>=0.24.1",
]
Expand Down
8 changes: 4 additions & 4 deletions src/ezmsg/tools/plot/describe.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,7 +112,7 @@ 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
(:attr:`~ezmsg.util.messages.axisarray.AxisArray.stream_dim`) and falls back
to the first of *fallbacks* the message actually has, which is what these
tools did before the field existed.

Expand All @@ -122,10 +122,10 @@ def stream_axis(msg: typing.Any, *fallbacks: str) -> str | None:
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)
stream_dim = getattr(msg, "stream_dim", None)
dims = getattr(msg, "dims", ()) or ()
if chunk_dim is not None and chunk_dim in dims:
return chunk_dim
if stream_dim is not None and stream_dim in dims:
return stream_dim
for name in fallbacks:
if name in dims:
return name
Expand Down
8 changes: 4 additions & 4 deletions src/ezmsg/tools/shmem/aux_meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@
{"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
``stream_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.
Expand Down Expand Up @@ -102,7 +102,7 @@ def encode_aux(
attrs: typing.Mapping[str, typing.Any],
key: str,
buffered_axis: str,
chunk_dim: typing.Optional[str] = None,
stream_dim: typing.Optional[str] = None,
) -> tuple[bytes, list[str]]:
"""Serialize an AxisArray's static metadata.

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

Expand All @@ -148,7 +148,7 @@ def decode_aux(blob: bytes) -> dict:
)
# 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)
payload.setdefault("stream_dim", None)
return payload


Expand Down
8 changes: 4 additions & 4 deletions src/ezmsg/tools/shmem/shmem.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,7 @@ class ShMemCircBuffSettings(ez.Settings):
conn: typing.Optional[multiprocessing.connection.Connection] = None

axis: typing.Optional[str] = None
"""Dimension to buffer along. ``None`` follows the message's ``chunk_dim``.
"""Dimension to buffer along. ``None`` follows the message's ``stream_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
Expand All @@ -171,7 +171,7 @@ class ShMemCircBuffSettings(ez.Settings):
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``.
window. Set explicitly only for a producer that declares no ``stream_dim``.
"""


Expand Down Expand Up @@ -407,7 +407,7 @@ def _update_aux_if_needed(self, msg: AxisArray) -> bool:
# which is knowledge it has no way to arrive at.
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)
blob, dropped = encode_aux(rolled_dims, msg.axes, msg.attrs, msg.key, buff_axis, stream_dim=msg.stream_dim)
if dropped:
dropped_set = frozenset(dropped)
if self.STATE.warned_dropped_attrs != dropped_set:
Expand Down Expand Up @@ -457,7 +457,7 @@ def _resolve_axis(self, msg: AxisArray) -> typing.Optional[str]:
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
axis = self.SETTINGS.axis if self.SETTINGS.axis is not None else msg.stream_dim
if axis is None:
axis = "time" if "time" in msg.dims else None
return axis
Expand Down
4 changes: 2 additions & 2 deletions src/ezmsg/tools/shmem/shmem_mirror.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,14 +121,14 @@ def dims(self) -> typing.Optional[typing.List[str]]:
return None if self._aux is None else self._aux["dims"]

@property
def chunk_dim(self) -> typing.Optional[str]:
def stream_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")
return None if self._aux is None else self._aux.get("stream_dim")

@property
def buffered_axis(self) -> typing.Optional[str]:
Expand Down
14 changes: 7 additions & 7 deletions tests/test_plot_describe.py
Original file line number Diff line number Diff line change
Expand Up @@ -233,8 +233,8 @@ class TestStreamAxis:
"""

@staticmethod
def _windowed(chunk_dim="win"):
kwargs = {"chunk_dim": chunk_dim} if chunk_dim else {}
def _windowed(stream_dim="win"):
kwargs = {"stream_dim": stream_dim} if stream_dim else {}
return AxisArray(
np.zeros((4, 10, 3), np.float32),
dims=["win", "time", "ch"],
Expand All @@ -247,7 +247,7 @@ def test_it_prefers_the_declaration(self):
assert stream_axis(self._windowed(), "time") == "win"

def test_it_falls_back_when_nothing_is_declared(self):
assert stream_axis(self._windowed(chunk_dim=None), "time") == "time"
assert stream_axis(self._windowed(stream_dim=None), "time") == "time"

def test_the_fallback_order_is_honoured(self):
msg = AxisArray(
Expand All @@ -259,17 +259,17 @@ def test_the_fallback_order_is_honoured(self):
assert stream_axis(msg, "time", "freq") == "freq"

def test_a_declaration_naming_an_absent_dim_is_ignored(self):
"""`chunk_dim` is validated at construction, but a message can reach a
"""`stream_dim` is validated at construction, but a message can reach a
viewer after a transform that dropped the dimension without updating
it. Falling back beats indexing on a name that is not there."""
msg = AxisArray(
np.zeros((8, 3), np.float32),
dims=["time", "ch"],
axes={"time": AxisArray.TimeAxis(fs=100.0)},
key="dev",
chunk_dim="time",
stream_dim="time",
)
object.__setattr__(msg, "chunk_dim", "win")
object.__setattr__(msg, "stream_dim", "win")
assert stream_axis(msg, "time") == "time"

def test_none_when_nothing_matches(self):
Expand All @@ -282,6 +282,6 @@ def test_a_plain_stream_is_unaffected(self):
dims=["time", "ch"],
axes={"time": AxisArray.TimeAxis(fs=100.0)},
key="dev",
chunk_dim="time",
stream_dim="time",
)
assert stream_axis(msg, "time") == "time"
30 changes: 15 additions & 15 deletions tests/test_shmem_aux_meta.py
Original file line number Diff line number Diff line change
Expand Up @@ -438,17 +438,17 @@ def test_reader_rejects_a_foreign_header_loudly():


# ---------------------------------------------------------------------------
# chunk_dim: which dimension the stream accumulates along
# stream_dim: which dimension the stream accumulates along
# ---------------------------------------------------------------------------


def _windowed(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, chunk_dim: str | None = "win") -> AxisArray:
def _windowed(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, stream_dim: str | None = "win") -> 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`.
nothing about the message distinguishes them except `stream_dim`.
"""
kwargs = {"chunk_dim": chunk_dim} if chunk_dim else {}
kwargs = {"stream_dim": stream_dim} if stream_dim else {}
return AxisArray(
np.zeros((n_win, n_lag, n_ch), np.float32),
dims=["win", "time", "ch"],
Expand All @@ -462,39 +462,39 @@ def _windowed(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, chunk_dim: str | N
)


class TestTheBlobCarriesChunkDim:
class TestTheBlobCarriesStreamDim:
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"
blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "win", stream_dim=msg.stream_dim)
assert decode_aux(blob)["stream_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
msg = _windowed(stream_dim=None)
blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "win", stream_dim=msg.stream_dim)
assert decode_aux(blob)["stream_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)
blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "time", stream_dim=msg.stream_dim)
payload = decode_aux(blob)
assert payload["buffered_axis"] == "time"
assert payload["chunk_dim"] == "win"
assert payload["stream_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")
blob, _ = encode_aux(list(msg.dims), msg.axes, msg.attrs, msg.key, "win", stream_dim="win")
payload = pickle.loads(blob)
del payload["chunk_dim"] # what an older writer emits
del payload["stream_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["stream_dim"] is None
assert decoded["buffered_axis"] == "win"


Expand Down
16 changes: 8 additions & 8 deletions tests/test_shmem_sink.py
Original file line number Diff line number Diff line change
Expand Up @@ -120,13 +120,13 @@ def test_shmem_change(change_type: str):
# ---------------------------------------------------------------------------


def _windowed_msg(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, chunk_dim: str | None = "win") -> AxisArray:
def _windowed_msg(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, stream_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`.
LinearAxes, so nothing distinguishes them but `stream_dim`.
"""
kwargs = {"chunk_dim": chunk_dim} if chunk_dim else {}
kwargs = {"stream_dim": stream_dim} if stream_dim else {}
return AxisArray(
np.zeros((n_win, n_lag, n_ch), np.float32),
dims=["win", "time", "ch"],
Expand All @@ -140,8 +140,8 @@ def _windowed_msg(n_win: int = 4, n_lag: int = 10, n_ch: int = 3, chunk_dim: str
)


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 {}
def _plain_msg(n_time: int = 20, n_ch: int = 3, stream_dim: str | None = "time") -> AxisArray:
kwargs = {"stream_dim": stream_dim} if stream_dim else {}
return AxisArray(
np.zeros((n_time, n_ch), np.float32),
dims=["time", "ch"],
Expand Down Expand Up @@ -180,10 +180,10 @@ def test_an_explicit_setting_still_wins(self):

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
`stream_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"
assert _sink()._resolve_axis(_plain_msg(stream_dim=None)) == "time"
assert _sink()._resolve_axis(_windowed_msg(stream_dim=None)) == "time"

def test_a_message_with_neither_is_skipped(self):
msg = AxisArray(
Expand Down
Loading