From ee37473eac5f4c87bfd6d1ebff163c874d8b2144 Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 4 Aug 2024 11:24:02 -0600 Subject: [PATCH] Power filter fixes --- py/expression_handlers.py | 23 +++++++++++++++++------ py/filtering.py | 27 ++++----------------------- 2 files changed, 21 insertions(+), 29 deletions(-) diff --git a/py/expression_handlers.py b/py/expression_handlers.py index 07c7fb0..3f01b3f 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -560,19 +560,30 @@ if EXT_SONAR: } @classmethod - def make_power_filter(cls, fdict, toplevel=False): + def make_power_filter(cls, fdict, *, toplevel=True): fdict = fdict.copy() compose_with = fdict.pop("compose_with", None) if compose_with: if not isinstance(compose_with, dict): raise TypeError("compose_with must be a dictionary") - fdict["compose_with"] = cls.make_power_filter(compose_with) - power_filter = EXT_SONAR.powernoise.PowerFilter(**fdict) - if not toplevel: - return power_filter + fdict["compose_with"] = cls.make_power_filter( + compose_with, toplevel=False + ) topargs = { k: fdict.pop(k, dv) for k, dv in cls.default_power_filter.items() } + power_filter = EXT_SONAR.powernoise.PowerFilter(**fdict) + if not toplevel: + return power_filter + cc = topargs.get("channel_correlation") + if cc is not None: + if not isinstance(cc, (list, tuple)) or not all( + isinstance(v, (int, float)) for v in cc + ): + raise TypeError( + "Bad channel correlation type: must be comma separated string or numeric sequence" + ) + topargs["channel_correlation"] = ",".join(repr(v) for v in cc) return EXT_SONAR.powernoise.PowerNoiseItem( 1, power_filter=power_filter, time_brownian=True, **topargs ) @@ -581,7 +592,7 @@ if EXT_SONAR: tensor, filter_def = self.safe_get_all(obj, getter) if not isinstance(filter_def, dict): raise TypeError("filter argument must be a dictionary") - power_filter = self.make_power_filter(filter_def, toplevel=True) + power_filter = self.make_power_filter(filter_def) filter_rfft = power_filter.make_filter(tensor.shape).to( tensor.device, non_blocking=True ) diff --git a/py/filtering.py b/py/filtering.py index 76e8fbf..77a8c81 100644 --- a/py/filtering.py +++ b/py/filtering.py @@ -482,12 +482,6 @@ if EXT_SONAR: class SonarPowerFilter(Filter): name = "sonar_power_filter" - default_power_filter = { - "mix": 1.0, - "normalization_factor": 1.0, - "common_mode": 0.0, - "channel_correlation": "1,1,1,1,1,1", - } default_options = Filter.default_options def __init__(self, **kwargs): @@ -498,23 +492,10 @@ if EXT_SONAR: return if not isinstance(power_filter, dict): raise ValueError("power_filter key must be dict or null") - self.power_filter = self.make_power_filter(power_filter, toplevel=True) - - @classmethod - def make_power_filter(cls, fdict, toplevel=False): - fdict = fdict.copy() - compose_with = fdict.pop("compose_with", None) - if compose_with: - fdict["compose_with"] = cls.make_power_filter(compose_with) - if toplevel: - topargs = { - k: fdict.pop(k, dv) for k, dv in cls.default_power_filter.items() - } - power_filter = EXT_SONAR.powernoise.PowerFilter(**fdict) - if not toplevel: - return power_filter - return EXT_SONAR.powernoise.PowerNoiseItem( - 1, power_filter=power_filter, time_brownian=True, **topargs + self.power_filter = ( + expression_handlers.SonarPowerFilterHandler.make_power_filter( + power_filter + ) ) def filter(self, latent, ref_latent, *args, refs=None, **kwargs):