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)