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()