Source code for ezmsg.sigproc.filterbankdesign

"""Kaiser-window filterbank design with configurable bands and processing mode."""

import typing

import ezmsg.core as ez
import numpy as np
import numpy.typing as npt
from ezmsg.baseproc import (
    BaseStatefulTransformer,
    processor_state,
    resolve_configured_stream_dim,
    suppress_axis_deprecation,
)
from ezmsg.util.messages.axisarray import AxisArray
from ezmsg.util.messages.util import replace

from .filterbank import (
    FilterbankMode,
    FilterbankSettings,
    FilterbankTransformer,
    MinPhaseMode,
)
from .kaiser import KaiserFilterSettings, kaiser_design_fun
from .util.deprecation import warn_axis_deprecated


[docs] class FilterbankDesignSettings(ez.Settings): filters: typing.Iterable[KaiserFilterSettings] mode: FilterbankMode = FilterbankMode.CONV """ "conv", "fft", or "auto". If "auto", the mode is determined by the size of the input data. fft mode is more efficient for long kernels. However, fft mode uses non-overlapping windows and will incur a delay equal to the window length, which is larger than the largest kernel. conv mode is less efficient but will return data for every incoming chunk regardless of how small it is and thus can provide shorter latency updates. """ min_phase: MinPhaseMode = MinPhaseMode.NONE """ If not None, convert the kernels to minimum-phase equivalents. Valid options are 'hilbert', 'homomorphic', and 'homomorphic-full'. Complex filters not supported. See `scipy.signal.minimum_phase` for details. """ 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.stream_dim`; see :mod:`ezmsg.sigproc.util.deprecation`.""" def __post_init__(self) -> None: warn_axis_deprecated(self) new_axis: str = "kernel" """The name of the new axis corresponding to the kernel index."""
[docs] @processor_state class FilterbankDesignState: axis: str = "" """The resolved stream dimension, fixed at reset so every later use agrees.""" filterbank: FilterbankTransformer | None = None needs_redesign: bool = False
[docs] class FilterbankDesignTransformer( BaseStatefulTransformer[FilterbankDesignSettings, AxisArray, AxisArray, FilterbankDesignState], ): """ Transformer that designs and applies a filterbank based on Kaiser windowed FIR filters. """
[docs] @classmethod def get_message_type(cls, dir: str) -> type[AxisArray]: if dir in ("in", "out"): return AxisArray else: raise ValueError(f"Invalid direction: {dir}. Must be 'in' or 'out'.")
[docs] def update_settings(self, new_settings: typing.Optional[FilterbankDesignSettings] = None, **kwargs) -> None: """ Update settings and mark that filter coefficients need to be recalculated. Args: new_settings: Complete new settings object to replace current settings **kwargs: Individual settings to update """ # Update settings if new_settings is not None: self.settings = new_settings else: self.settings = replace(self.settings, **kwargs) # Set flag to trigger recalculation on next message if self.state.filterbank is not None: self.state.needs_redesign = True
def _calculate_kernels(self, fs: float) -> list[npt.NDArray]: kernels = [] for filter in self.settings.filters: output = kaiser_design_fun( fs, cutoff=filter.cutoff, ripple=filter.ripple, width=filter.width, pass_zero=filter.pass_zero, wn_hz=filter.wn_hz, ) kernels.append(np.array([1.0]) if output is None else output[0]) return kernels def __call__(self, message: AxisArray) -> AxisArray: if self.state.filterbank is not None and self.state.needs_redesign: self._reset_state(message) self.state.needs_redesign = False return super().__call__(message) def _hash_message(self, message: AxisArray) -> int: # The only state is a FilterbankTransformer whose kernels are a function # of the sample rate. That inner transformer keeps its own hash and # rebuilds itself when the channels change, so folding the channel # fingerprint in here would only redesign kernels that came out the same. # Runs before `_reset_state`, so the axis is resolved from the message # rather than read back off state that does not exist yet. axis = resolve_configured_stream_dim(self, message, self.settings.axis, legacy_default="time") return hash((message.key, getattr(message.axes.get(axis), "gain", None))) def _reset_state(self, message: AxisArray) -> None: self.state.axis = resolve_configured_stream_dim(self, message, self.settings.axis, legacy_default="time") axis_obj = message.axes[self.state.axis] assert isinstance(axis_obj, AxisArray.LinearAxis) fs = 1 / axis_obj.gain kernels = self._calculate_kernels(fs) # Forwards this stage's own `axis`, which is not the deprecated setting; # warning here would name FilterbankSettings for something the user set on # FilterbankDesignSettings, and would do it on every reset. with suppress_axis_deprecation(): new_settings = FilterbankSettings( kernels=kernels, mode=self.settings.mode, min_phase=self.settings.min_phase, axis=self.state.axis, new_axis=self.settings.new_axis, ) self.state.filterbank = FilterbankTransformer(settings=new_settings) def _process(self, message: AxisArray) -> AxisArray: return self.state.filterbank(message)