Source code for fbscatnet.generate_bank

from collections.abc import KeysView, ValuesView

import matplotlib.pyplot as plt
import numpy as np

from .logger_config import setup_logger
from .math_utils import (
    _find_neumann_root_muller,
    _generate_fourier_bessel_wavelet,
    _generate_fourier_low_pass_filter,
)

logger = setup_logger(__name__)


[docs] class FourierBesselWaveletBank: """A bank of Fourier-Bessel wavelets indexed by parameters m and k. Attributes: size (int): Image size. m (int): Maximum order. k (int): Maximum angular index. sigma (float): Scale parameter for the wavelets. verbose (bool): Display wavelet diagnostic prints. wavelet_bank (dict[str, np.ndarray]): Stored dictionary of wavelets. mk_to_key (dict[tuple, str]): Mapping from (m, k) tuples to string keys. """ def __init__( self, size: int, m: int, k: int, sigma: float = 0.3, norm: str = "l1", verbose: bool = False ) -> None: """Initialise the FourierBesselWaveletBank. Args: size (int): Image size. m (int): Maximum order. k (int): Maximum angular index. sigma (float, optional): Scale parameter for the wavelets. norm (str, optional): Wavelet normalisation verbose (bool, optional): Display wavelet diagnostic prints. Defaults to False. Raises: ValueError: If angular order k is greater than m. ValueError: If input parameters are negative ValueError: If norm is not 'l1' or 'l2' """ self.size = size self.m = m self.k = k self.sigma = sigma self.sigma2 = sigma**2 self.m_values = np.arange(0, m) self.k_values = np.arange(0, k) self.verbose = verbose self.norm = norm if self.m < self.k: raise ValueError("m <= k condition is not respected") if self.m < 0 or self.k < 0 or self.sigma < 0 or self.size <= 0: raise ValueError("Cannot accept negative parameters") if self.norm not in ("l1", "l2"): raise ValueError(f"Invalid norm: Must be 'l1' or 'l2' not {self.norm}") self.lambda_max = _find_neumann_root_muller( 0, int(self.k_values.max()) if len(self.k_values) > 0 else 0 ) wavelet_bank: dict[str, np.ndarray] = {} mk_to_key: dict[tuple, str] = {} self.freq_limit = int(self.lambda_max + 2 / self.sigma) for k_val in self.k_values: for m_val in self.m_values: if np.abs(m_val) > k_val: continue if k_val == 0: _, Z = _generate_fourier_low_pass_filter( size=self.size, sigma=self.sigma, norm=self.norm, freq_limit=self.freq_limit, verbose=self.verbose, ) else: _, _, Z = _generate_fourier_bessel_wavelet( m_val, k_val, size=self.size, sigma=self.sigma, norm=self.norm, freq_limit=self.freq_limit, verbose=self.verbose, ) key_name = f"m_{m_val}_k_{k_val}_s{self.sigma}" mk_to_key[m_val, k_val] = key_name wavelet_bank[key_name] = Z self.mk_to_key = mk_to_key self.wavelet_bank = wavelet_bank def __getitem__(self, key_or_indices: str | tuple) -> np.ndarray: """Retrieve a specific wavelet by its string key or (m_index, k_index) tuple. Args: key_or_indices (str | tuple): Index to retrieve either as string "m_k_sigma" or by an (m, k) tuple. Returns: np.ndarray: The requested wavelet array. Raises: KeyError: If the key or (m, k) combination does not exist. TypeError: If an invalid index type is provided. """ if isinstance(key_or_indices, str): if key_or_indices not in self.wavelet_bank: raise KeyError(f"Wavelet key '{key_or_indices}' not found.") return self.wavelet_bank[key_or_indices] elif isinstance(key_or_indices, tuple) and len(key_or_indices) == 2: m_index, k_index = key_or_indices key = self.mk_to_key.get((m_index, k_index)) if key is None: raise KeyError(f"Wavelet with parameters m={m_index}, k={k_index} not found.") return self.wavelet_bank[key] raise TypeError("Invalid index type. Use a string key or an (m, k) tuple.") def __len__(self) -> int: """Retrieve the number of wavelets in the bank. Returns: int: Total count of wavelets. """ return len(self.wavelet_bank)
[docs] def get_keys(self) -> KeysView[str]: """Return the dictionary keys representing individual wavelets. Returns: KeysView[str]: View of the string keys. """ return self.wavelet_bank.keys()
[docs] def get_values(self) -> ValuesView[np.ndarray]: """Return the collection of wavelet arrays. Returns: ValuesView[np.ndarray]: View of the wavelet numpy arrays. """ return self.wavelet_bank.values()
[docs] def summary(self, verbose: bool = True) -> tuple[int, int, float]: """Print an optional summary of the wavelet bank parameters and size. Args: verbose (bool, optional): Whether to print summary details. Defaults to True. Returns: tuple[int, int, float]: A tuple containing (m, k, sigma). """ if verbose: print() logger.info("Fourier-Bessel Wavelet bank summary:") logger.info( "Parameters: m = %d, k = %d, sigma = %.2f, norm = %s", self.m, self.k, self.sigma, self.norm, ) logger.info("Total wavelets: %d", len(self.wavelet_bank)) logger.info("Frequency limit: %.2f", self.freq_limit) logger.info("Key naming structure: {m_{m_val}_k_{k_val}_s{sigma}}") return self.m, self.k, self.sigma
[docs] def plot_bank(self) -> None: """Plot the grid of wavelets using matplotlib.""" _, axes = plt.subplots( len(self.k_values), len(self.m_values), figsize=(8, 6), sharex=True, sharey=True ) axes = np.atleast_2d(axes) for row_idx, k_val in enumerate(self.k_values): for col_idx, m_val in enumerate(self.m_values): ax = axes[row_idx, col_idx] ax.set_xlim(-self.freq_limit, self.freq_limit) ax.set_ylim(-self.freq_limit, self.freq_limit) ax.axis("off") if row_idx == 0: ax.set_title(f"m = {m_val}", fontsize=12, fontweight="bold") if col_idx == 0: ax.text( -self.freq_limit * 1.2, 0, f"k = {k_val}", fontsize=12, fontweight="bold", ha="right", va="center", ) if np.abs(m_val) > k_val: continue if f"m_{m_val}_k_{k_val}_s{self.sigma}" not in self.wavelet_bank: continue Z = self.wavelet_bank[f"m_{m_val}_k_{k_val}_s{self.sigma}"] z_max: float = float(np.max(np.abs(Z))) ax.imshow( np.real(Z), extent=[-self.freq_limit, self.freq_limit, -self.freq_limit, self.freq_limit], cmap="inferno", origin="lower", vmin=-z_max, vmax=z_max, ) plt.tight_layout() plt.subplots_adjust(wspace=0.05, hspace=0.05) plt.show()