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
32 changes: 25 additions & 7 deletions src/ezmsg/sigproc/adaptive_lattice_notch.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,9 @@
from ezmsg.util.messages.axisarray import AxisArray, CoordinateAxis
from ezmsg.util.messages.util import replace

from .util.deprecation import warn_axis_deprecated
from .util.message import resolve_configured_chunk_dim


class AdaptiveLatticeNotchFilterSettings(ez.Settings):
"""Settings for the Adaptive Lattice Notch Filter."""
Expand All @@ -18,7 +21,15 @@ class AdaptiveLatticeNotchFilterSettings(ez.Settings):
"""Smoothing factor"""
eta: float = 0.99
"""Forgetting factor"""
axis: str = "time"
axis: str | None = None
""".. deprecated:: 3.8
Scheduled for removal in 4.0. The dimension messages accumulate along
now comes from :attr:`~ezmsg.util.messages.axisarray.AxisArray.chunk_dim`;
see :mod:`ezmsg.sigproc.util.deprecation`."""

def __post_init__(self) -> None:
warn_axis_deprecated(self)

"""Axis to apply filter to"""
init_notch_freq: float | None = None
"""Initial notch frequency. Should be < nyquist."""
Expand All @@ -28,6 +39,9 @@ class AdaptiveLatticeNotchFilterSettings(ez.Settings):

@processor_state
class AdaptiveLatticeNotchFilterState:
axis: str = ""
"""The resolved chunk dimension, fixed at reset so every later use agrees."""

"""State for the Adaptive Lattice Notch Filter."""

s_history: npt.NDArray | None = None
Expand Down Expand Up @@ -71,10 +85,11 @@ class AdaptiveLatticeNotchFilterTransformer(
NONRESET_SETTINGS_FIELDS = frozenset({"gamma", "mu", "eta", "chunkwise"})

def _reset_state(self, message: AxisArray) -> None:
ax_idx = message.get_axis_idx(self.settings.axis)
axis = resolve_configured_chunk_dim(self, message, self.settings.axis, legacy_default="time")
ax_idx = message.get_axis_idx(axis)
sample_shape = message.data.shape[:ax_idx] + message.data.shape[ax_idx + 1 :]

fs = 1 / message.axes[self.settings.axis].gain
fs = 1 / message.axes[axis].gain
init_f = (
self.settings.init_notch_freq if self.settings.init_notch_freq is not None else 0.07178314656435313 * fs
)
Expand All @@ -83,13 +98,15 @@ def _reset_state(self, message: AxisArray) -> None:

"""Reset filter state to initial values."""
self._state = AdaptiveLatticeNotchFilterState()
# Set after the wholesale replacement above, which would otherwise drop it.
self._state.axis = axis
self._state.s_history = np.zeros((2,) + sample_shape, dtype=float)
self._state.p = np.zeros(sample_shape, dtype=float)
self._state.q = np.zeros(sample_shape, dtype=float)
self._state.k1 = init_k1 + np.zeros(sample_shape, dtype=float)
self._state.freq_template = CoordinateAxis(
data=np.zeros((0,) + sample_shape, dtype=float),
dims=[self.settings.axis] + message.dims[:ax_idx] + message.dims[ax_idx + 1 :],
dims=[axis] + message.dims[:ax_idx] + message.dims[ax_idx + 1 :],
unit="Hz",
)

Expand All @@ -105,17 +122,18 @@ def _reset_state(self, message: AxisArray) -> None:

def _process(self, message: AxisArray) -> AxisArray:
x_data = message.data
ax_idx = message.get_axis_idx(self.settings.axis)
axis = self._state.axis
ax_idx = message.get_axis_idx(axis)

# TODO: Time should be moved to -1th axis, not the 0th axis
if message.dims[0] != self.settings.axis:
if message.dims[0] != axis:
x_data = np.moveaxis(x_data, ax_idx, 0)

# Access settings once
gamma = self.settings.gamma
eta = self.settings.eta
mu = self.settings.mu
fs = 1 / message.axes[self.settings.axis].gain
fs = 1 / message.axes[axis].gain

# Pre-compute constants
one_minus_eta = 1 - eta
Expand Down
25 changes: 19 additions & 6 deletions src/ezmsg/sigproc/adaptive_lnc.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,9 @@
from ezmsg.util.messages.axisarray import AxisArray
from ezmsg.util.messages.util import replace

from .util.deprecation import warn_axis_deprecated
from .util.message import resolve_configured_chunk_dim

# Optional Apple-Silicon GPU backend. The canceller is an LTI SOS notch
# cascade (see `design_lnc_sos`), so on MLX arrays we dispatch to the Metal
# `sosfilt` kernel; everything else runs through scipy on the array's own
Expand Down Expand Up @@ -165,14 +168,23 @@ class AdaptiveLNCSettings(ez.Settings):
unchanged. (Per-channel sampling-delay alignment is handled separately,
upstream, by ``SamplingDelayAlignmentTransformer``.)"""

axis: str = "time"
"""Name of the axis to filter along."""
axis: str | None = None
""".. deprecated:: 3.8
Scheduled for removal in 4.0. The dimension messages accumulate along
now comes from :attr:`~ezmsg.util.messages.axisarray.AxisArray.chunk_dim`;
see :mod:`ezmsg.sigproc.util.deprecation`."""

def __post_init__(self) -> None:
warn_axis_deprecated(self)


@processor_state
class AdaptiveLNCState:
"""State for :class:`AdaptiveLNCTransformer`."""

axis: str = ""
"""The resolved chunk dimension, fixed at reset so every later use agrees."""

omega: float = 0.0
"""Current NCO angular frequency in rad/sample (tracked by the FLL)."""

Expand Down Expand Up @@ -288,11 +300,12 @@ class AdaptiveLNCTransformer(
NONRESET_SETTINGS_FIELDS = frozenset({"adapt_time_constant", "freq_time_constant"})

def _reset_state(self, message: AxisArray) -> None:
ax_idx = message.get_axis_idx(self.settings.axis)
self._state.axis = resolve_configured_chunk_dim(self, message, self.settings.axis, legacy_default="time")
ax_idx = message.get_axis_idx(self._state.axis)
sample_shape = message.data.shape[:ax_idx] + message.data.shape[ax_idx + 1 :]
xp, is_mlx = _namespace(message.data)

fs = 1.0 / message.axes[self.settings.axis].gain
fs = 1.0 / message.axes[self._state.axis].gain
# Seed the NCO at the nominal normalised frequency; the FLL refines it.
self._state.omega = 2.0 * np.pi * self.settings.line_freq / fs
max_deviation = self.settings.max_freq_deviation
Expand Down Expand Up @@ -472,7 +485,7 @@ def _process(self, message: AxisArray) -> AxisArray:
# No cancellation and no frequency tracking; emit the input as-is.
return message

ax_idx = message.get_axis_idx(self.settings.axis)
ax_idx = message.get_axis_idx(self._state.axis)
x_data = message.data
xp, is_mlx = _namespace(x_data)
moved = ax_idx != 0
Expand All @@ -482,7 +495,7 @@ def _process(self, message: AxisArray) -> AxisArray:
n = x_data.shape[0]
st = self._state
dtype = x_data.dtype
fs = 1.0 / message.axes[self.settings.axis].gain
fs = 1.0 / message.axes[self._state.axis].gain

# Time constants -> gains (independent of chunk size and fs).
# mu = 2 / (tau_adapt * fs); beta = 1 - exp(-window_dt / tau_freq).
Expand Down
10 changes: 5 additions & 5 deletions src/ezmsg/sigproc/affinetransform.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@
from ezmsg.sigproc.util.array import array_device, is_float_dtype, xp_asarray, xp_copy, xp_create, xp_empty
from ezmsg.sigproc.util.blockdiag import plan_block_matmul
from ezmsg.sigproc.util.channels import ChannelGroupSpec, resolve_channel_groups
from ezmsg.sigproc.util.message import with_fingerprint
from ezmsg.sigproc.util.message import resolve_feature_dim, with_fingerprint
from ezmsg.sigproc.util.rereference import RereferenceKind, rereference_matrix

KERNELS = ("auto", "dense", "blocks")
Expand Down Expand Up @@ -242,7 +242,7 @@ def _reset_state(self, message: AxisArray) -> None:
if self.settings.kernel not in KERNELS:
raise ValueError(f"kernel must be one of {KERNELS}, got {self.settings.kernel!r}")

axis = self.settings.axis or message.dims[-1]
axis = self.settings.axis or resolve_feature_dim(message)
axis_idx = message.get_axis_idx(axis)
n_in = message.data.shape[axis_idx]
xp = get_namespace(message.data)
Expand Down Expand Up @@ -447,7 +447,7 @@ def _stacked_split(self, xp):

def _process(self, message: AxisArray) -> AxisArray:
xp = get_namespace(message.data)
axis = self.settings.axis or message.dims[-1]
axis = self.settings.axis or resolve_feature_dim(message)
axis_idx = message.get_axis_idx(axis)
data = message.data

Expand Down Expand Up @@ -588,7 +588,7 @@ class CommonRereferenceTransformer(
def _reset_state(self, message: AxisArray) -> None:
xp = get_namespace(message.data)
dev = array_device(message.data)
axis = self.settings.axis or message.dims[-1]
axis = self.settings.axis or resolve_feature_dim(message)
axis_idx = message.get_axis_idx(axis)
n_ch = message.data.shape[axis_idx]
include_current = self.settings.include_current
Expand Down Expand Up @@ -638,7 +638,7 @@ def _process(self, message: AxisArray) -> AxisArray:
return message

xp = get_namespace(message.data)
axis = self.settings.axis or message.dims[-1]
axis = self.settings.axis or resolve_feature_dim(message)
axis_idx = message.get_axis_idx(axis)
state = self._state
data = message.data
Expand Down
6 changes: 3 additions & 3 deletions src/ezmsg/sigproc/aggregate.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
)

from .spectral import OptionsEnum
from .util.message import with_fingerprint
from .util.message import resolve_feature_dim, with_fingerprint


class AggregationFunction(OptionsEnum):
Expand Down Expand Up @@ -260,7 +260,7 @@ async def __acall__(self, message: AxisArray) -> AxisArray:
return await super().__acall__(message)

def _reset_state(self, message: AxisArray) -> None:
axis = self.settings.axis or message.dims[0]
axis = self.settings.axis or resolve_feature_dim(message, 0)
target_axis = message.get_axis(axis)
ax_idx = message.get_axis_idx(axis)

Expand Down Expand Up @@ -293,7 +293,7 @@ def _reset_state(self, message: AxisArray) -> None:
)

def _process(self, message: AxisArray) -> AxisArray:
axis = self.settings.axis or message.dims[0]
axis = self.settings.axis or resolve_feature_dim(message, 0)
ax_idx = message.get_axis_idx(axis)

# The bands are already resolved to slices in _reset_state, and ax_vec
Expand Down
18 changes: 13 additions & 5 deletions src/ezmsg/sigproc/align.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,11 +12,19 @@
from ezmsg.util.messages.axisarray import AxisArray

from .util.axisarray_buffer import HybridAxisArrayBuffer
from .util.deprecation import warn_axis_deprecated
from .util.message import resolve_configured_chunk_dim


class AlignAlongAxisSettings(ez.Settings):
axis: str = "time"
"""Axis used for alignment (typically the time axis)."""
axis: str | None = None
""".. deprecated:: 3.8
Scheduled for removal in 4.0. The dimension messages accumulate along
now comes from :attr:`~ezmsg.util.messages.axisarray.AxisArray.chunk_dim`;
see :mod:`ezmsg.sigproc.util.deprecation`."""

def __post_init__(self) -> None:
warn_axis_deprecated(self)

buffer_dur: float = 10.0
"""Buffer duration in seconds for each input stream."""
Expand Down Expand Up @@ -79,7 +87,7 @@ def _hash_message(self, message: AxisArray) -> int:
return hash(self._extract_gain(message))

def _extract_gain(self, message: AxisArray) -> float | None:
align_name = self.settings.axis or message.dims[0]
align_name = resolve_configured_chunk_dim(self, message, self.settings.axis)
ax = message.axes.get(align_name)
if ax is not None and hasattr(ax, "gain"):
return ax.gain
Expand Down Expand Up @@ -131,7 +139,7 @@ def _request_reset(self) -> None:
super()._request_reset()

def _reset_state(self, message: AxisArray) -> None:
align_axis = self.settings.axis or message.dims[0]
align_axis = resolve_configured_chunk_dim(self, message, self.settings.axis)
if self._hash == -1 and not getattr(self, "_force_full_reset", False):
self._state.align_axis = align_axis
if self._state.buf_a is None:
Expand Down Expand Up @@ -159,7 +167,7 @@ def _process(self, message: AxisArray) -> _AlignPair | None:

def push_b(self, message: AxisArray) -> _AlignPair | None:
"""Process input B: check gain, detect shape changes, buffer, try align."""
align_axis = self.settings.axis or message.dims[0]
align_axis = resolve_configured_chunk_dim(self, message, self.settings.axis)

# Gain compatibility check. Skipped when B's gain can't be estimated
# (e.g. a single-sample CoordinateAxis yields None) — there is nothing
Expand Down
27 changes: 19 additions & 8 deletions src/ezmsg/sigproc/binned_aggregate.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,14 +56,21 @@
from .aggregate import AggregationFunction, aggregate_slices, needs_coordinates
from .util.array import xp_copy
from .util.binning import BinSchedule, BinStep
from .util.message import is_empty_along, with_fingerprint
from .util.deprecation import warn_axis_deprecated
from .util.message import is_empty_along, resolve_configured_chunk_dim, with_fingerprint


class BinnedAggregateSettings(ez.Settings):
"""Settings for :obj:`BinnedAggregate`."""

axis: str = "time"
"""The name of the axis to bin and aggregate along."""
axis: str | None = None
""".. deprecated:: 3.8
Scheduled for removal in 4.0. The dimension messages accumulate along
now comes from :attr:`~ezmsg.util.messages.axisarray.AxisArray.chunk_dim`;
see :mod:`ezmsg.sigproc.util.deprecation`."""

def __post_init__(self) -> None:
warn_axis_deprecated(self)

bin_duration: float = 0.02
"""Output bin duration in seconds."""
Expand Down Expand Up @@ -113,6 +120,9 @@ class BinnedAggregateSettings(ez.Settings):

@processor_state
class BinnedAggregateState:
axis: str = ""
"""The resolved chunk dimension, fixed at reset so every later use agrees."""

schedule: BinSchedule | None = None
"""Shared bin-boundary schedule (see :obj:`ezmsg.sigproc.util.binning`). Owns
the sample rate, samples-per-bin, output gain, global bin index, and carried
Expand Down Expand Up @@ -167,7 +177,8 @@ async def __acall__(self, message: AxisArray) -> AxisArray:
return await super().__acall__(message)

def _reset_state(self, message: AxisArray) -> None:
axis_info = message.get_axis(self.settings.axis)
self._state.axis = resolve_configured_chunk_dim(self, message, self.settings.axis, legacy_default="time")
axis_info = message.get_axis(self._state.axis)
schedule = BinSchedule(
bin_duration=self.settings.bin_duration,
fractional=self.settings.fractional,
Expand Down Expand Up @@ -227,10 +238,10 @@ def _out_dims(self, message: AxisArray) -> list[str]:
return dims + [self.settings.newaxis] if self._multi else dims

def _out_axes(self, message: AxisArray, step: BinStep) -> dict:
axis_info = message.get_axis(self.settings.axis)
axis_info = message.get_axis(self._state.axis)
axes = {
**message.axes,
self.settings.axis: replace(axis_info, gain=step.output_gain, offset=step.output_offset),
self._state.axis: replace(axis_info, gain=step.output_gain, offset=step.output_offset),
}
if self._multi:
axes[self.settings.newaxis] = self._state.metric_axis
Expand All @@ -252,7 +263,7 @@ def _empty_like(self, message: AxisArray, axis_idx: int, step: BinStep) -> AxisA
)

def _process(self, message: AxisArray) -> AxisArray:
axis = self.settings.axis
axis = self._state.axis
axis_info = message.get_axis(axis)
axis_idx = message.get_axis_idx(axis)
xp = get_namespace(message.data)
Expand Down Expand Up @@ -321,5 +332,5 @@ async def on_signal(self, message: AxisArray) -> typing.AsyncGenerator:
cadence.
"""
result = await self.processor.__acall__(message)
if result is not None and not is_empty_along(result, (self.SETTINGS.axis,)):
if result is not None and not is_empty_along(result, (self.processor.state.axis,)):
yield self.OUTPUT_SIGNAL, result
5 changes: 3 additions & 2 deletions src/ezmsg/sigproc/butterworthzerophase.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
_sosfilt_mlx_metal_xp,
)
from .util.array import xp_asarray, xp_copy, xp_empty, xp_flip
from .util.message import resolve_configured_chunk_dim

if _HAS_MLX_METAL:
import mlx.core as _mx
Expand Down Expand Up @@ -188,7 +189,7 @@ def _reset_state(self, message: AxisArray) -> None:
self._tail = None
self._tail_offset = 0.0
# Compute pad_length based on the message's sampling rate
axis = message.dims[0] if self.settings.axis is None else self.settings.axis
axis = resolve_configured_chunk_dim(self, message, self.settings.axis)
fs = 1 / message.axes[axis].gain
self._pad_length = self._compute_pad_length(fs)
self.state.needs_redesign = True
Expand Down Expand Up @@ -230,7 +231,7 @@ def _initialize_zi(self, data, ax_idx: int, xp):
return self._zi_tiled * first_sample

def _process(self, message: AxisArray) -> AxisArray:
axis = message.dims[0] if self.settings.axis is None else self.settings.axis
axis = resolve_configured_chunk_dim(self, message, self.settings.axis)
ax_idx = message.get_axis_idx(axis)
fs = 1 / message.axes[axis].gain

Expand Down
4 changes: 2 additions & 2 deletions src/ezmsg/sigproc/coordinatespaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,7 @@
)
from ezmsg.util.messages.axisarray import AxisArray, replace

from .util.message import with_fingerprint
from .util.message import resolve_feature_dim, with_fingerprint

# -- Utility functions for coordinate transformations --

Expand Down Expand Up @@ -109,7 +109,7 @@ class CoordinateSpacesTransformer(BaseTransformer[CoordinateSpacesSettings, Axis

def _process(self, message: AxisArray) -> AxisArray:
xp = get_namespace(message.data)
axis = self.settings.axis or message.dims[-1]
axis = self.settings.axis or resolve_feature_dim(message)
axis_idx = message.get_axis_idx(axis)

if message.data.shape[axis_idx] != 2:
Expand Down
Loading
Loading