"""Which dimension a processor should operate on.
An ``AxisArray`` names its dimensions but nothing about a *position* says what a
dimension means. ``dims[0]`` is not "the streaming axis" and ``dims[-1]`` is not
"the channel axis"; both guesses break under
:meth:`~ezmsg.util.messages.axisarray.AxisArray.transpose` and downstream of any
windowing stage, where a ``(time, ch)`` stream becomes ``(win, time, ch)``.
:attr:`~ezmsg.util.messages.axisarray.AxisArray.stream_dim` is the producer's
declaration of which dimension messages accumulate along -- the one party that
reliably knows. These helpers turn that declaration into the axis a given kind of
processor should use, and they live here because
:meth:`~ezmsg.baseproc.BaseStatefulTransformer._message_hash` already resolves
the same thing for its own purposes: a processor whose arithmetic disagreed with
its state-reset logic about which dimension is which would reset on the wrong
changes and cache state along the wrong axis.
Three rules, because one does not fit every case:
* :func:`resolve_stream_dim` -- for state carried *between* messages.
* :func:`resolve_feature_dim` -- for a static axis (channels, components).
* :func:`resolve_transform_dim` -- for a transform that consumes a regularly
sampled axis, which downstream of a windowing stage is *not* the stream one.
"""
import typing
import ezmsg.core as ez
from ezmsg.util.messages.axisarray import AxisArray
__all__ = [
"STREAMING_DIMS",
"resolve_configured_stream_dim",
"resolve_feature_dim",
"resolve_stream_dim",
"resolve_transform_dim",
]
STREAMING_DIMS: tuple[str, ...] = ("time",)
"""Default fallback stream dimension, matching ``BaseStatefulTransformer``."""
[docs]
def resolve_stream_dim(message: AxisArray, streaming_dims: typing.Iterable[str] = STREAMING_DIMS) -> str:
"""The dimension successive messages accumulate along.
This is the axis a processor that carries state *between* messages must
operate on -- filter initial conditions, a running mean, a sample buffer,
a previous-sample cache. Carrying such state along any other dimension is
not a smaller error but a different operation: a static axis has the same
length every message, so state carried across it applies message N's tail
to message N+1's head at the same coordinate, forever.
The producer renamed the dims and so is the only party that reliably knows
which one grows; ``message.stream_dim`` is that declaration. When a producer
is silent, *streaming_dims* supplies the guess -- ``("time",)`` is right for
a raw signal and wrong downstream of a windowing stage, where the message is
``(win, time, ch)`` and ``win`` is what grows.
``dims[0]`` is the last resort only. It is a position, not a meaning, and it
breaks under :meth:`~ezmsg.util.messages.axisarray.AxisArray.transpose`.
"""
if message.stream_dim is not None:
return message.stream_dim
for name in streaming_dims:
if name in message.dims:
return name
return message.dims[0]
[docs]
def resolve_feature_dim(message: AxisArray, position: int = -1) -> str:
"""The dimension at *position*, skipping the stream dimension.
For processors whose axis is a *static* one -- channels, coordinate
components, feature labels. ``stream_dim`` is emphatically not the answer
here, but the naive ``dims[position]`` can silently *be* the stream
dimension: a ``(ch, time)`` stream makes ``dims[-1]`` the accumulating axis,
and an affine transform would then matmul across time while a slicer would
discard samples.
Falls back to ``dims[position]`` when the stream dimension is all there is,
which keeps 1-D messages working rather than raising on them.
"""
candidates = [d for d in message.dims if d != message.stream_dim]
if not candidates:
return message.dims[position]
return candidates[position]
[docs]
def with_fingerprint(axis: AxisArray.CoordinateAxis) -> AxisArray.CoordinateAxis:
"""Compute *axis*'s fingerprint now, and return the axis.
Every stateful consumer reads the fingerprint of the coordinate axes that
describe a stream's configuration, and the value is cached on the instance
and pickled with it. Computing it at the point of construction therefore
pays the checksum once, for everybody:
* In this process, the axis object is reused for the life of the stream, so
one call covers every message and every consumer downstream of it.
* Across a process boundary it is better than that. Unpickling hands out a
*new* axis object per message, so a cold axis is re-checksummed by the
first consumer in every receiving process, on every message, forever.
A primed one arrives with the answer already attached.
Apply it to axes that describe the stream -- channel labels, frequency
labels, feature labels -- not to per-message coordinates along the stream
dimension, whose fingerprint no consumer reads and whose data is new every
message anyway.
"""
axis.fingerprint # noqa: B018 -- evaluated for the caching side effect
return axis
[docs]
def is_empty_along(message: AxisArray, dims: typing.Iterable[str]) -> bool:
"""True iff any of the named dims is present in ``message`` with zero length.
Publish gates use this instead of ``data.size == 0`` so a message that is
empty only along *other* axes — e.g. an upstream selection removed every
channel while time samples remain — still flows downstream, preserving the
stream's cadence for consumers that align or merge multiple sources.
Dims not present in the message are ignored.
"""
return any(d in message.dims and message.data.shape[message.get_axis_idx(d)] == 0 for d in dims)
[docs]
def has_samples_along(message: AxisArray, dim: str) -> bool:
"""True iff ``dim`` is present in ``message`` with nonzero length.
Stricter than ``not is_empty_along(...)``: the dim must exist. Drain loops
use this to decide whether a chunk is real output, so that a placeholder
lacking the axis entirely (e.g. ResampleProcessor's pre-init null template,
``dims=[""]``) counts as "nothing ready" rather than a publishable chunk.
"""
return dim in message.dims and message.data.shape[message.get_axis_idx(dim)] > 0