fbscatnet 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- fbscatnet/__init__.py +14 -0
- fbscatnet/generate_bank.py +231 -0
- fbscatnet/logger_config.py +47 -0
- fbscatnet/math_utils.py +236 -0
- fbscatnet/py.typed +0 -0
- fbscatnet/scatnet.py +468 -0
- fbscatnet-1.0.0.dist-info/METADATA +120 -0
- fbscatnet-1.0.0.dist-info/RECORD +10 -0
- fbscatnet-1.0.0.dist-info/WHEEL +4 -0
- fbscatnet-1.0.0.dist-info/licenses/LICENSE +21 -0
fbscatnet/__init__.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
|
|
3
|
+
from .generate_bank import FourierBesselWaveletBank
|
|
4
|
+
from .scatnet import FourierBesselScatNet
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def setup_logger(name: str) -> logging.Logger:
|
|
8
|
+
logger = logging.getLogger(name)
|
|
9
|
+
logger.addHandler(logging.NullHandler())
|
|
10
|
+
|
|
11
|
+
return logger
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
__all__ = ["FourierBesselWaveletBank", "FourierBesselScatNet"]
|
|
@@ -0,0 +1,231 @@
|
|
|
1
|
+
from collections.abc import KeysView, ValuesView
|
|
2
|
+
|
|
3
|
+
import matplotlib.pyplot as plt
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from .logger_config import setup_logger
|
|
7
|
+
from .math_utils import (
|
|
8
|
+
_find_neumann_root_muller,
|
|
9
|
+
_generate_fourier_bessel_wavelet,
|
|
10
|
+
_generate_fourier_low_pass_filter,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
logger = setup_logger(__name__)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class FourierBesselWaveletBank:
|
|
17
|
+
"""A bank of Fourier-Bessel wavelets indexed by parameters m and k.
|
|
18
|
+
|
|
19
|
+
Attributes:
|
|
20
|
+
size (int): Image size.
|
|
21
|
+
m (int): Maximum order.
|
|
22
|
+
k (int): Maximum angular index.
|
|
23
|
+
sigma (float): Scale parameter for the wavelets.
|
|
24
|
+
verbose (bool): Display wavelet diagnostic prints.
|
|
25
|
+
wavelet_bank (dict[str, np.ndarray]): Stored dictionary of wavelets.
|
|
26
|
+
mk_to_key (dict[tuple, str]): Mapping from (m, k) tuples to string keys.
|
|
27
|
+
"""
|
|
28
|
+
|
|
29
|
+
def __init__(
|
|
30
|
+
self, size: int, m: int, k: int, sigma: float = 0.3, norm: str = "l1", verbose: bool = False
|
|
31
|
+
) -> None:
|
|
32
|
+
"""Initialise the FourierBesselWaveletBank.
|
|
33
|
+
|
|
34
|
+
Args:
|
|
35
|
+
size (int): Image size.
|
|
36
|
+
m (int): Maximum order.
|
|
37
|
+
k (int): Maximum angular index.
|
|
38
|
+
sigma (float, optional): Scale parameter for the wavelets.
|
|
39
|
+
norm (str, optional): Wavelet normalisation
|
|
40
|
+
verbose (bool, optional): Display wavelet diagnostic prints. Defaults to False.
|
|
41
|
+
|
|
42
|
+
Raises:
|
|
43
|
+
ValueError: If angular order k is greater than m.
|
|
44
|
+
ValueError: If input parameters are negative
|
|
45
|
+
ValueError: If norm is not 'l1' or 'l2'
|
|
46
|
+
"""
|
|
47
|
+
self.size = size
|
|
48
|
+
self.m = m
|
|
49
|
+
self.k = k
|
|
50
|
+
self.sigma = sigma
|
|
51
|
+
self.sigma2 = sigma**2
|
|
52
|
+
self.m_values = np.arange(0, m)
|
|
53
|
+
self.k_values = np.arange(0, k)
|
|
54
|
+
self.verbose = verbose
|
|
55
|
+
self.norm = norm
|
|
56
|
+
|
|
57
|
+
if self.m < self.k:
|
|
58
|
+
raise ValueError("m <= k condition is not respected")
|
|
59
|
+
|
|
60
|
+
if self.m < 0 or self.k < 0 or self.sigma < 0 or self.size <= 0:
|
|
61
|
+
raise ValueError("Cannot accept negative parameters")
|
|
62
|
+
|
|
63
|
+
if self.norm not in ("l1", "l2"):
|
|
64
|
+
raise ValueError(f"Invalid norm: Must be 'l1' or 'l2' not {self.norm}")
|
|
65
|
+
|
|
66
|
+
self.lambda_max = _find_neumann_root_muller(
|
|
67
|
+
0, int(self.k_values.max()) if len(self.k_values) > 0 else 0
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
wavelet_bank: dict[str, np.ndarray] = {}
|
|
71
|
+
mk_to_key: dict[tuple, str] = {}
|
|
72
|
+
|
|
73
|
+
self.freq_limit = int(self.lambda_max + 2 / self.sigma)
|
|
74
|
+
|
|
75
|
+
for k_val in self.k_values:
|
|
76
|
+
for m_val in self.m_values:
|
|
77
|
+
if np.abs(m_val) > k_val:
|
|
78
|
+
continue
|
|
79
|
+
|
|
80
|
+
if k_val == 0:
|
|
81
|
+
_, Z = _generate_fourier_low_pass_filter(
|
|
82
|
+
size=self.size,
|
|
83
|
+
sigma=self.sigma,
|
|
84
|
+
norm=self.norm,
|
|
85
|
+
freq_limit=self.freq_limit,
|
|
86
|
+
verbose=self.verbose,
|
|
87
|
+
)
|
|
88
|
+
else:
|
|
89
|
+
_, _, Z = _generate_fourier_bessel_wavelet(
|
|
90
|
+
m_val,
|
|
91
|
+
k_val,
|
|
92
|
+
size=self.size,
|
|
93
|
+
sigma=self.sigma,
|
|
94
|
+
norm=self.norm,
|
|
95
|
+
freq_limit=self.freq_limit,
|
|
96
|
+
verbose=self.verbose,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
key_name = f"m_{m_val}_k_{k_val}_s{self.sigma}"
|
|
100
|
+
mk_to_key[m_val, k_val] = key_name
|
|
101
|
+
wavelet_bank[key_name] = Z
|
|
102
|
+
|
|
103
|
+
self.mk_to_key = mk_to_key
|
|
104
|
+
self.wavelet_bank = wavelet_bank
|
|
105
|
+
|
|
106
|
+
def __getitem__(self, key_or_indices: str | tuple) -> np.ndarray:
|
|
107
|
+
"""Retrieve a specific wavelet by its string key or (m_index, k_index) tuple.
|
|
108
|
+
|
|
109
|
+
Args:
|
|
110
|
+
key_or_indices (str | tuple): Index to retrieve either as string "m_k_sigma"
|
|
111
|
+
or by an (m, k) tuple.
|
|
112
|
+
|
|
113
|
+
Returns:
|
|
114
|
+
np.ndarray: The requested wavelet array.
|
|
115
|
+
|
|
116
|
+
Raises:
|
|
117
|
+
KeyError: If the key or (m, k) combination does not exist.
|
|
118
|
+
TypeError: If an invalid index type is provided.
|
|
119
|
+
"""
|
|
120
|
+
if isinstance(key_or_indices, str):
|
|
121
|
+
if key_or_indices not in self.wavelet_bank:
|
|
122
|
+
raise KeyError(f"Wavelet key '{key_or_indices}' not found.")
|
|
123
|
+
return self.wavelet_bank[key_or_indices]
|
|
124
|
+
|
|
125
|
+
elif isinstance(key_or_indices, tuple) and len(key_or_indices) == 2:
|
|
126
|
+
m_index, k_index = key_or_indices
|
|
127
|
+
|
|
128
|
+
key = self.mk_to_key.get((m_index, k_index))
|
|
129
|
+
if key is None:
|
|
130
|
+
raise KeyError(f"Wavelet with parameters m={m_index}, k={k_index} not found.")
|
|
131
|
+
|
|
132
|
+
return self.wavelet_bank[key]
|
|
133
|
+
|
|
134
|
+
raise TypeError("Invalid index type. Use a string key or an (m, k) tuple.")
|
|
135
|
+
|
|
136
|
+
def __len__(self) -> int:
|
|
137
|
+
"""Retrieve the number of wavelets in the bank.
|
|
138
|
+
|
|
139
|
+
Returns:
|
|
140
|
+
int: Total count of wavelets.
|
|
141
|
+
"""
|
|
142
|
+
return len(self.wavelet_bank)
|
|
143
|
+
|
|
144
|
+
def get_keys(self) -> KeysView[str]:
|
|
145
|
+
"""Return the dictionary keys representing individual wavelets.
|
|
146
|
+
|
|
147
|
+
Returns:
|
|
148
|
+
KeysView[str]: View of the string keys.
|
|
149
|
+
"""
|
|
150
|
+
return self.wavelet_bank.keys()
|
|
151
|
+
|
|
152
|
+
def get_values(self) -> ValuesView[np.ndarray]:
|
|
153
|
+
"""Return the collection of wavelet arrays.
|
|
154
|
+
|
|
155
|
+
Returns:
|
|
156
|
+
ValuesView[np.ndarray]: View of the wavelet numpy arrays.
|
|
157
|
+
"""
|
|
158
|
+
return self.wavelet_bank.values()
|
|
159
|
+
|
|
160
|
+
def summary(self, verbose: bool = True) -> tuple[int, int, float]:
|
|
161
|
+
"""Print an optional summary of the wavelet bank parameters and size.
|
|
162
|
+
|
|
163
|
+
Args:
|
|
164
|
+
verbose (bool, optional): Whether to print summary details. Defaults to True.
|
|
165
|
+
|
|
166
|
+
Returns:
|
|
167
|
+
tuple[int, int, float]: A tuple containing (m, k, sigma).
|
|
168
|
+
"""
|
|
169
|
+
if verbose:
|
|
170
|
+
print()
|
|
171
|
+
logger.info("Fourier-Bessel Wavelet bank summary:")
|
|
172
|
+
logger.info(
|
|
173
|
+
"Parameters: m = %d, k = %d, sigma = %.2f, norm = %s",
|
|
174
|
+
self.m,
|
|
175
|
+
self.k,
|
|
176
|
+
self.sigma,
|
|
177
|
+
self.norm,
|
|
178
|
+
)
|
|
179
|
+
logger.info("Total wavelets: %d", len(self.wavelet_bank))
|
|
180
|
+
logger.info("Frequency limit: %.2f", self.freq_limit)
|
|
181
|
+
logger.info("Key naming structure: {m_{m_val}_k_{k_val}_s{sigma}}")
|
|
182
|
+
|
|
183
|
+
return self.m, self.k, self.sigma
|
|
184
|
+
|
|
185
|
+
def plot_bank(self) -> None:
|
|
186
|
+
"""Plot the grid of wavelets using matplotlib."""
|
|
187
|
+
_, axes = plt.subplots(
|
|
188
|
+
len(self.k_values), len(self.m_values), figsize=(8, 6), sharex=True, sharey=True
|
|
189
|
+
)
|
|
190
|
+
|
|
191
|
+
axes = np.atleast_2d(axes)
|
|
192
|
+
for row_idx, k_val in enumerate(self.k_values):
|
|
193
|
+
for col_idx, m_val in enumerate(self.m_values):
|
|
194
|
+
ax = axes[row_idx, col_idx]
|
|
195
|
+
ax.set_xlim(-self.freq_limit, self.freq_limit)
|
|
196
|
+
ax.set_ylim(-self.freq_limit, self.freq_limit)
|
|
197
|
+
ax.axis("off")
|
|
198
|
+
|
|
199
|
+
if row_idx == 0:
|
|
200
|
+
ax.set_title(f"m = {m_val}", fontsize=12, fontweight="bold")
|
|
201
|
+
if col_idx == 0:
|
|
202
|
+
ax.text(
|
|
203
|
+
-self.freq_limit * 1.2,
|
|
204
|
+
0,
|
|
205
|
+
f"k = {k_val}",
|
|
206
|
+
fontsize=12,
|
|
207
|
+
fontweight="bold",
|
|
208
|
+
ha="right",
|
|
209
|
+
va="center",
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
if np.abs(m_val) > k_val:
|
|
213
|
+
continue
|
|
214
|
+
|
|
215
|
+
if f"m_{m_val}_k_{k_val}_s{self.sigma}" not in self.wavelet_bank:
|
|
216
|
+
continue
|
|
217
|
+
|
|
218
|
+
Z = self.wavelet_bank[f"m_{m_val}_k_{k_val}_s{self.sigma}"]
|
|
219
|
+
z_max: float = float(np.max(np.abs(Z)))
|
|
220
|
+
ax.imshow(
|
|
221
|
+
np.real(Z),
|
|
222
|
+
extent=[-self.freq_limit, self.freq_limit, -self.freq_limit, self.freq_limit],
|
|
223
|
+
cmap="inferno",
|
|
224
|
+
origin="lower",
|
|
225
|
+
vmin=-z_max,
|
|
226
|
+
vmax=z_max,
|
|
227
|
+
)
|
|
228
|
+
|
|
229
|
+
plt.tight_layout()
|
|
230
|
+
plt.subplots_adjust(wspace=0.05, hspace=0.05)
|
|
231
|
+
plt.show()
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
import sys
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class LogColors:
|
|
6
|
+
RESET = "\033[0m"
|
|
7
|
+
DEBUG = "\033[36m" # Cyan
|
|
8
|
+
INFO = "\033[32m" # Green
|
|
9
|
+
WARNING = "\033[33m" # Yellow
|
|
10
|
+
ERROR = "\033[31m" # Red
|
|
11
|
+
CRITICAL = "\033[35m" # Magenta
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class ColoredFormatter(logging.Formatter):
|
|
15
|
+
def format(self, record: logging.LogRecord) -> str:
|
|
16
|
+
color = getattr(LogColors, record.levelname, LogColors.RESET)
|
|
17
|
+
colored_level = f"{color}%(levelname)s{LogColors.RESET}"
|
|
18
|
+
log_fmt = f"%(asctime)s [{colored_level}] %(message)s"
|
|
19
|
+
formatter = logging.Formatter(log_fmt, datefmt="%Y-%m-%d %H:%M:%S")
|
|
20
|
+
return formatter.format(record)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def setup_logger(name: str) -> logging.Logger:
|
|
24
|
+
"""Internal function to set up the library logger with a NullHandler."""
|
|
25
|
+
logger = logging.getLogger(name)
|
|
26
|
+
if not logger.handlers:
|
|
27
|
+
logger.addHandler(logging.NullHandler())
|
|
28
|
+
return logger
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def enable_colored_logs(level: int = logging.INFO) -> None:
|
|
32
|
+
"""
|
|
33
|
+
Helper function for users to easily enable colored console logging for fbscatnet.
|
|
34
|
+
|
|
35
|
+
Usage:
|
|
36
|
+
import fbscatnet
|
|
37
|
+
fbscatnet.enable_colored_logs()
|
|
38
|
+
"""
|
|
39
|
+
# Get the top-level logger for the package
|
|
40
|
+
logger = logging.getLogger("fbscatnet")
|
|
41
|
+
logger.setLevel(level)
|
|
42
|
+
|
|
43
|
+
# Check if we already added a StreamHandler to avoid duplicates
|
|
44
|
+
if not any(isinstance(h, logging.StreamHandler) for h in logger.handlers):
|
|
45
|
+
handler = logging.StreamHandler(sys.stdout)
|
|
46
|
+
handler.setFormatter(ColoredFormatter())
|
|
47
|
+
logger.addHandler(handler)
|
fbscatnet/math_utils.py
ADDED
|
@@ -0,0 +1,236 @@
|
|
|
1
|
+
from typing import cast
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import numpy.typing as npt
|
|
5
|
+
import scipy.special
|
|
6
|
+
|
|
7
|
+
from .logger_config import setup_logger
|
|
8
|
+
|
|
9
|
+
logger = setup_logger(__name__)
|
|
10
|
+
|
|
11
|
+
### BESSEL WRAPPERS VIA SCIPY ###
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _first_kind_bessel(X: npt.ArrayLike, order: int | np.integer) -> float | np.ndarray:
|
|
15
|
+
"""Vectorised Bessel function supporting both scalar points and arrays."""
|
|
16
|
+
X_arr = np.asarray(X)
|
|
17
|
+
abs_order = int(abs(int(order)))
|
|
18
|
+
|
|
19
|
+
if order < 0:
|
|
20
|
+
val = (-1) ** abs_order * scipy.special.jv(abs_order, X_arr)
|
|
21
|
+
else:
|
|
22
|
+
val = scipy.special.jv(int(order), X_arr)
|
|
23
|
+
|
|
24
|
+
# If the input was a scalar, return a Python float; otherwise the array
|
|
25
|
+
if np.isscalar(X):
|
|
26
|
+
return float(np.asarray(val).item())
|
|
27
|
+
return cast(np.ndarray, np.asarray(val))
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _first_kind_bessel_deriv(X: npt.ArrayLike, order: int | np.integer) -> float | np.ndarray:
|
|
31
|
+
"""Derivative supporting both scalar points and arrays using recurrence relations."""
|
|
32
|
+
val = 0.5 * (_first_kind_bessel(X, order - 1) - _first_kind_bessel(X, order + 1))
|
|
33
|
+
return val
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _first_modified_bessel(X: npt.ArrayLike, order: int | np.integer) -> float | np.ndarray:
|
|
37
|
+
"""Modified Bessel function supporting both scalar points and arrays."""
|
|
38
|
+
|
|
39
|
+
X_arr = np.asarray(X)
|
|
40
|
+
val = scipy.special.iv(order, X_arr)
|
|
41
|
+
|
|
42
|
+
# If the input was a scalar, return a Python float; otherwise the array
|
|
43
|
+
if np.isscalar(X):
|
|
44
|
+
return float(np.asarray(val).item())
|
|
45
|
+
return cast(np.ndarray, np.asarray(val))
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
### MULLER'S METHOD TO FIND EIGENVALUE ###
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _mcmahon_seed(m: int | np.integer, k: int | np.integer) -> float:
|
|
52
|
+
"""Approximate the k-th positive root using McMahon's asymptotic expansion."""
|
|
53
|
+
m = int(m)
|
|
54
|
+
k = int(k)
|
|
55
|
+
beta: float
|
|
56
|
+
if m == 0:
|
|
57
|
+
s = k
|
|
58
|
+
nu = 1
|
|
59
|
+
beta = (s + nu / 2.0 - 0.25) * np.pi
|
|
60
|
+
return beta - (4.0 * nu**2 - 1.0) / (8.0 * beta)
|
|
61
|
+
else:
|
|
62
|
+
s = k - m + 1
|
|
63
|
+
beta = (s + m / 2.0 - 0.75) * np.pi
|
|
64
|
+
if beta <= 0:
|
|
65
|
+
return float(m) + 0.5
|
|
66
|
+
return beta - (4.0 * m**2 + 3.0) / (8.0 * beta)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _find_neumann_root_muller(
|
|
70
|
+
m: int | np.integer, k: int | np.integer, thresh: float = 1e-12, max_iter: int = 200
|
|
71
|
+
) -> float:
|
|
72
|
+
"""Find the k-th positive root of using Muller's method."""
|
|
73
|
+
|
|
74
|
+
x3 = 0.0
|
|
75
|
+
|
|
76
|
+
if k < m or (m == 0 and k == 0):
|
|
77
|
+
return x3
|
|
78
|
+
|
|
79
|
+
x_seed = _mcmahon_seed(m, k)
|
|
80
|
+
x0 = max(0.01, x_seed - 0.15)
|
|
81
|
+
x1 = x_seed
|
|
82
|
+
x2 = x_seed + 0.15
|
|
83
|
+
|
|
84
|
+
def f(x: npt.ArrayLike) -> float | np.ndarray:
|
|
85
|
+
return _first_kind_bessel_deriv(x, m)
|
|
86
|
+
|
|
87
|
+
d0: float = float(f(x0))
|
|
88
|
+
d1: float = float(f(x1))
|
|
89
|
+
d2: float = float(f(x2))
|
|
90
|
+
|
|
91
|
+
converged = False
|
|
92
|
+
|
|
93
|
+
for _ in range(max_iter):
|
|
94
|
+
h1 = x1 - x0
|
|
95
|
+
h2 = x2 - x1
|
|
96
|
+
if h1 == 0 or h2 == 0:
|
|
97
|
+
break
|
|
98
|
+
delta1 = (d1 - d0) / h1
|
|
99
|
+
delta2 = (d2 - d1) / h2
|
|
100
|
+
if (h2 + h1) == 0:
|
|
101
|
+
break
|
|
102
|
+
d_coef = (delta2 - delta1) / (h2 + h1)
|
|
103
|
+
a = d_coef
|
|
104
|
+
b = delta2 + h2 * d_coef
|
|
105
|
+
c = d2
|
|
106
|
+
disc = np.lib.scimath.sqrt(b**2 - 4 * a * c)
|
|
107
|
+
|
|
108
|
+
dx = -2 * c / (b + disc) if np.real(b) >= 0 else -2 * c / (b - disc)
|
|
109
|
+
|
|
110
|
+
x3 = np.real(x2 + dx).item()
|
|
111
|
+
d3 = float(f(x3))
|
|
112
|
+
if abs(d3) < thresh or abs(dx) < thresh:
|
|
113
|
+
converged = True
|
|
114
|
+
break
|
|
115
|
+
x0, x1, x2 = x1, x2, x3
|
|
116
|
+
d0, d1, d2 = d1, d2, d3
|
|
117
|
+
|
|
118
|
+
if not converged:
|
|
119
|
+
logger.warning(
|
|
120
|
+
"Root finder did not converge for m=%d, k=%d after %d iterations "
|
|
121
|
+
"(residual=%.2e). Consider increasing max_iter.",
|
|
122
|
+
int(m),
|
|
123
|
+
int(k),
|
|
124
|
+
max_iter,
|
|
125
|
+
abs(d3) if "d3" in locals() else float("nan"),
|
|
126
|
+
)
|
|
127
|
+
return float(x3)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
### FOURIER WAVELET ###
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def _generate_fourier_bessel_wavelet(
|
|
134
|
+
m: int | np.integer,
|
|
135
|
+
k: int | np.integer,
|
|
136
|
+
size: int = 50,
|
|
137
|
+
sigma: float = 0.1,
|
|
138
|
+
norm: str = "l1",
|
|
139
|
+
freq_limit: int = 20,
|
|
140
|
+
verbose: bool = False,
|
|
141
|
+
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
|
142
|
+
|
|
143
|
+
abs_m = np.abs(m)
|
|
144
|
+
freq = np.linspace(-freq_limit, freq_limit, size)
|
|
145
|
+
Kx: np.ndarray
|
|
146
|
+
Ky: np.ndarray
|
|
147
|
+
Kx, Ky = np.meshgrid(freq, freq, indexing="ij")
|
|
148
|
+
|
|
149
|
+
Q = np.sqrt(Kx**2 + Ky**2)
|
|
150
|
+
Psi = np.arctan2(Ky, Kx)
|
|
151
|
+
|
|
152
|
+
eigenvalue = _find_neumann_root_muller(abs_m, k)
|
|
153
|
+
sigma2 = sigma**2
|
|
154
|
+
eig2 = eigenvalue**2
|
|
155
|
+
angular_profile = np.exp(1j * m * Psi)
|
|
156
|
+
K = np.exp(-(eig2 * sigma2) / 2)
|
|
157
|
+
|
|
158
|
+
mod_bessel = _first_modified_bessel((eig2 * sigma2) / 2, m)
|
|
159
|
+
mod_bessel_freq = _first_modified_bessel(eigenvalue * sigma2 * Q, m)
|
|
160
|
+
|
|
161
|
+
start = (1j**m) * angular_profile
|
|
162
|
+
left_hand = sigma2 * np.exp(-(sigma2 * (eig2 + Q**2)) / 2) * mod_bessel_freq
|
|
163
|
+
|
|
164
|
+
if m == 0:
|
|
165
|
+
bracket = K * mod_bessel - 2 * np.exp(-(3 * sigma2 * eig2) / 4) + np.exp(-sigma2 * eig2)
|
|
166
|
+
norm_term = 1 / (np.sqrt((np.pi * sigma2) * bracket))
|
|
167
|
+
|
|
168
|
+
else:
|
|
169
|
+
norm_term = 1 / (np.sqrt(np.pi * sigma2 * K * mod_bessel))
|
|
170
|
+
|
|
171
|
+
if m == 0:
|
|
172
|
+
Z = start * norm_term * (left_hand - K * sigma2 * np.exp(-(sigma2 * Q**2) / 2))
|
|
173
|
+
else:
|
|
174
|
+
Z = start * norm_term * left_hand
|
|
175
|
+
|
|
176
|
+
if norm == "l1":
|
|
177
|
+
z_max = np.max(np.abs(Z))
|
|
178
|
+
Z /= z_max
|
|
179
|
+
|
|
180
|
+
# --- DIAGNOSTIC PRINTS ---
|
|
181
|
+
if verbose:
|
|
182
|
+
dx = freq[1] - freq[0]
|
|
183
|
+
mean_val = np.abs(np.mean(Z))
|
|
184
|
+
if norm == "l2":
|
|
185
|
+
post_norm_energy = np.sum(np.abs(Z) ** 2) * (dx * dx)
|
|
186
|
+
else:
|
|
187
|
+
post_norm_energy = np.max(np.abs(Z))
|
|
188
|
+
|
|
189
|
+
logger.info(
|
|
190
|
+
f"Wavelet (m={m}, k={k}) | "
|
|
191
|
+
f"Mean (~0): {mean_val:.2e} | "
|
|
192
|
+
f"{norm.upper()}: {post_norm_energy:.4f} | "
|
|
193
|
+
f"Eigenvalue: {eigenvalue:.4f}"
|
|
194
|
+
)
|
|
195
|
+
# -------------------------
|
|
196
|
+
|
|
197
|
+
return Kx, Ky, Z
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def _generate_fourier_low_pass_filter(
|
|
201
|
+
size: int = 50,
|
|
202
|
+
sigma: float = 0.1,
|
|
203
|
+
norm: str = "l1",
|
|
204
|
+
freq_limit: int = 20,
|
|
205
|
+
verbose: bool = False,
|
|
206
|
+
) -> tuple[np.ndarray, np.ndarray]:
|
|
207
|
+
freq = np.linspace(-freq_limit, freq_limit, size)
|
|
208
|
+
Kx: np.ndarray
|
|
209
|
+
Ky: np.ndarray
|
|
210
|
+
Kx, Ky = np.meshgrid(freq, freq, indexing="ij")
|
|
211
|
+
Q = np.sqrt(Kx**2 + Ky**2)
|
|
212
|
+
|
|
213
|
+
sigma2 = sigma**2
|
|
214
|
+
Z = np.exp(-(sigma2) * (Q**2) / 2)
|
|
215
|
+
|
|
216
|
+
if norm == "l1":
|
|
217
|
+
z_max = np.max(np.abs(Z))
|
|
218
|
+
Z /= z_max
|
|
219
|
+
|
|
220
|
+
# --- DIAGNOSTIC PRINTS ---
|
|
221
|
+
if verbose:
|
|
222
|
+
dx = freq[1] - freq[0]
|
|
223
|
+
mean_val = np.abs(np.mean(Z))
|
|
224
|
+
if norm == "l2":
|
|
225
|
+
post_norm_energy = np.sum(np.abs(Z) ** 2) * (dx * dx)
|
|
226
|
+
else:
|
|
227
|
+
post_norm_energy = np.max(np.abs(Z))
|
|
228
|
+
|
|
229
|
+
logger.info(
|
|
230
|
+
f"Low pass (m=0, k=0) | "
|
|
231
|
+
f"Mean (>0): {mean_val:.2e} | "
|
|
232
|
+
f"{norm.upper()}: {post_norm_energy:.4f}"
|
|
233
|
+
)
|
|
234
|
+
# -------------------------
|
|
235
|
+
|
|
236
|
+
return Q, Z
|
fbscatnet/py.typed
ADDED
|
File without changes
|
fbscatnet/scatnet.py
ADDED
|
@@ -0,0 +1,468 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import re
|
|
3
|
+
import warnings
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import matplotlib.pyplot as plt
|
|
7
|
+
import numpy as np
|
|
8
|
+
import scipy.fft as spfft
|
|
9
|
+
from joblib import Parallel, delayed
|
|
10
|
+
from tqdm import tqdm
|
|
11
|
+
|
|
12
|
+
from .generate_bank import FourierBesselWaveletBank
|
|
13
|
+
from .logger_config import setup_logger
|
|
14
|
+
|
|
15
|
+
logger = setup_logger(__name__)
|
|
16
|
+
|
|
17
|
+
try:
|
|
18
|
+
with warnings.catch_warnings():
|
|
19
|
+
warnings.simplefilter("ignore", UserWarning)
|
|
20
|
+
import cupy as cp
|
|
21
|
+
|
|
22
|
+
try:
|
|
23
|
+
HAS_CUPY = cp.cuda.runtime.getDeviceCount() > 0
|
|
24
|
+
except Exception as exc:
|
|
25
|
+
logger.warning("CuPy is installed but no usable GPU was detected: %s", exc)
|
|
26
|
+
HAS_CUPY = False
|
|
27
|
+
except ImportError:
|
|
28
|
+
cp = None
|
|
29
|
+
HAS_CUPY = False
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _extract_k(key_str: str) -> int:
|
|
33
|
+
"""Extract the integer 'k' index encoded in a filter bank key string."""
|
|
34
|
+
match = re.search(r"k[=_ ]*(\d+)", key_str)
|
|
35
|
+
if match:
|
|
36
|
+
return int(match.group(1))
|
|
37
|
+
else:
|
|
38
|
+
logger.warning("Valid k key not found in '%s', defaulting to 0", key_str)
|
|
39
|
+
return 0
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class FourierBesselScatNet:
|
|
43
|
+
"""A scattering network based on Fourier-Bessel wavelets for generating image embeddings.
|
|
44
|
+
|
|
45
|
+
Attributes:
|
|
46
|
+
size (int): Image spatial size.
|
|
47
|
+
bank (FourierBesselWaveletBank): The bank of Fourier-Bessel wavelets.
|
|
48
|
+
num_filters (int): Total number of filters in the wavelet bank.
|
|
49
|
+
bank_keys (list[str]): List of string keys representing the filters.
|
|
50
|
+
low_pass (np.ndarray): Low-pass filter array retrieved from the bank.
|
|
51
|
+
final_features (np.ndarray): Generated feature embeddings from the last run.
|
|
52
|
+
"""
|
|
53
|
+
|
|
54
|
+
def __init__(self, bank: FourierBesselWaveletBank, backend: str = "cpu") -> None:
|
|
55
|
+
"""Initialise the FourierBesselScatNet.
|
|
56
|
+
|
|
57
|
+
Args:
|
|
58
|
+
bank (FourierBesselWaveletBank): A configured bank of Fourier-Bessel wavelets.
|
|
59
|
+
backend (str): Device to operate on
|
|
60
|
+
|
|
61
|
+
Raises:
|
|
62
|
+
ImportError: If CuPy requested but no GPU.
|
|
63
|
+
ValueError: If the bank contains no band-pass wavelets (i.e. only
|
|
64
|
+
the (0, 0) low-pass filter), since the scattering cascade has
|
|
65
|
+
nothing to filter with in that case.
|
|
66
|
+
"""
|
|
67
|
+
self.backend = backend.lower()
|
|
68
|
+
if self.backend == "gpu" and not HAS_CUPY:
|
|
69
|
+
raise ImportError(
|
|
70
|
+
"GPU backend was requested, but 'cupy' is not"
|
|
71
|
+
""
|
|
72
|
+
"installed or no compatible GPU was found."
|
|
73
|
+
)
|
|
74
|
+
if len(bank) <= 1:
|
|
75
|
+
raise ValueError(
|
|
76
|
+
"The wavelet bank contains no band-pass wavelets (only the "
|
|
77
|
+
"low-pass filter). FourierBesselScatNet needs at least one "
|
|
78
|
+
"wavelet - try constructing the bank with m > 1 and/or k > 1."
|
|
79
|
+
)
|
|
80
|
+
self.xp: Any = cp if self.backend == "gpu" else np
|
|
81
|
+
self.size = bank[0, 0].shape[0]
|
|
82
|
+
self.bank = bank
|
|
83
|
+
self.num_filters = len(bank)
|
|
84
|
+
self.bank_keys = list(bank.get_keys())
|
|
85
|
+
self.low_pass = self.xp.asarray(bank[0, 0], dtype=self.xp.float32)
|
|
86
|
+
self.final_features: Any | None = None
|
|
87
|
+
|
|
88
|
+
# Precompute everything that is constant across batches
|
|
89
|
+
xp = self.xp
|
|
90
|
+
self.order_1_keys = self.bank_keys[1:]
|
|
91
|
+
self.num_order_1_maps = len(self.order_1_keys)
|
|
92
|
+
|
|
93
|
+
self._k_map = {key: _extract_k(key) for key in self.order_1_keys}
|
|
94
|
+
|
|
95
|
+
self._order2_children: dict[str, list[str]] = {
|
|
96
|
+
key1: [key2 for key2 in self.order_1_keys if self._k_map[key2] < self._k_map[key1]]
|
|
97
|
+
for key1 in self.order_1_keys
|
|
98
|
+
}
|
|
99
|
+
self.num_order_2_maps = sum(len(v) for v in self._order2_children.values())
|
|
100
|
+
|
|
101
|
+
self._filters_1 = {
|
|
102
|
+
key: xp.asarray(bank[key], dtype=xp.complex64) for key in self.order_1_keys
|
|
103
|
+
}
|
|
104
|
+
self._low_pass_c = xp.asarray(bank[0, 0], dtype=xp.complex64)
|
|
105
|
+
|
|
106
|
+
# Stacked filter tensors for vectorised filtering.
|
|
107
|
+
# Shape: (num_order_1_maps, H, W)
|
|
108
|
+
self._filters_1_stack = xp.stack([self._filters_1[k] for k in self.order_1_keys])
|
|
109
|
+
|
|
110
|
+
# Per-key1 stacked children filters for vectorised order-2 filtering.
|
|
111
|
+
# Shape per entry: (num_children, H, W), or None if no valid children.
|
|
112
|
+
self._filters_2_stack_by_key1: dict[str, Any] = {
|
|
113
|
+
key1: (xp.stack([self._filters_1[k2] for k2 in children]) if children else None)
|
|
114
|
+
for key1, children in self._order2_children.items()
|
|
115
|
+
}
|
|
116
|
+
|
|
117
|
+
def _fft2(self, x: Any, xp: Any) -> Any:
|
|
118
|
+
"""Backend-aware 2D FFT (uses scipy on CPU for multi-threaded speed, cupy on GPU)."""
|
|
119
|
+
if self.backend == "gpu":
|
|
120
|
+
return xp.fft.fft2(x, axes=(-2, -1))
|
|
121
|
+
return spfft.fft2(x, axes=(-2, -1), workers=-1)
|
|
122
|
+
|
|
123
|
+
def _ifft2(self, x: Any, xp: Any) -> Any:
|
|
124
|
+
"""Backend-aware inverse 2D FFT, mirroring `_fft2`."""
|
|
125
|
+
if self.backend == "gpu":
|
|
126
|
+
return xp.fft.ifft2(x, axes=(-2, -1))
|
|
127
|
+
return spfft.ifft2(x, axes=(-2, -1), workers=-1)
|
|
128
|
+
|
|
129
|
+
def _filter_and_modulus(self, xp: Any, freq_signal: Any, filters: Any) -> Any:
|
|
130
|
+
"""Apply wavelet filter(s) in the frequency domain and take the spatial modulus.
|
|
131
|
+
|
|
132
|
+
Args:
|
|
133
|
+
xp: The array module to use (numpy or cupy) for this batch.
|
|
134
|
+
freq_signal: Input already in fftshifted frequency domain,
|
|
135
|
+
broadcastable against `filters`.
|
|
136
|
+
filters: One or a stack of frequency-domain filters to apply.
|
|
137
|
+
|
|
138
|
+
Returns:
|
|
139
|
+
The fftshifted frequency-domain modulus.
|
|
140
|
+
"""
|
|
141
|
+
filtered_fft = freq_signal * filters
|
|
142
|
+
shifted_freq = xp.fft.ifftshift(filtered_fft, axes=(-2, -1))
|
|
143
|
+
spatial_complex = self._ifft2(shifted_freq, xp)
|
|
144
|
+
modulus_spatial = xp.abs(spatial_complex).astype(xp.float32)
|
|
145
|
+
return xp.fft.fftshift(self._fft2(modulus_spatial, xp), axes=(-2, -1)).astype(xp.complex64)
|
|
146
|
+
|
|
147
|
+
def _smooth_and_pool(
|
|
148
|
+
self,
|
|
149
|
+
xp: Any,
|
|
150
|
+
modulus_fft: Any,
|
|
151
|
+
low_pass_c: Any,
|
|
152
|
+
batch_size: int,
|
|
153
|
+
n_maps: int,
|
|
154
|
+
d_size: int,
|
|
155
|
+
downsize: int,
|
|
156
|
+
) -> Any:
|
|
157
|
+
"""Low-pass smooth a modulus map and block-mean-pool it down to `d_size`.
|
|
158
|
+
|
|
159
|
+
Args:
|
|
160
|
+
xp: The array module to use (numpy or cupy) for this batch.
|
|
161
|
+
modulus_fft: Output of `_filter_and_modulus`, shape (batch, n_maps, H, W).
|
|
162
|
+
low_pass_c: Broadcastable complex low-pass filter.
|
|
163
|
+
batch_size: Number of samples in this batch.
|
|
164
|
+
n_maps: Number of filter channels being pooled.
|
|
165
|
+
d_size: Output spatial size after downsampling (H // downsize).
|
|
166
|
+
downsize: Block-mean pooling factor.
|
|
167
|
+
|
|
168
|
+
Returns:
|
|
169
|
+
Real-valued pooled features of shape (batch_size, d_size, d_size, n_maps).
|
|
170
|
+
"""
|
|
171
|
+
filtered_low_pass = modulus_fft * low_pass_c
|
|
172
|
+
smoothed_shifted = xp.fft.ifftshift(filtered_low_pass, axes=(-2, -1))
|
|
173
|
+
smoothed_spatial = self._ifft2(smoothed_shifted, xp)
|
|
174
|
+
return (
|
|
175
|
+
xp.real(smoothed_spatial)
|
|
176
|
+
.reshape(batch_size, n_maps, d_size, downsize, d_size, downsize)
|
|
177
|
+
.mean(axis=(3, 5))
|
|
178
|
+
.transpose(0, 2, 3, 1)
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
def generate_embeddings(
|
|
182
|
+
self,
|
|
183
|
+
data: np.ndarray,
|
|
184
|
+
downsize: int,
|
|
185
|
+
batch_size: int = 32,
|
|
186
|
+
use_multiprocessing: bool = False,
|
|
187
|
+
) -> np.ndarray:
|
|
188
|
+
"""Generate scattering network feature embeddings for a given dataset.
|
|
189
|
+
|
|
190
|
+
Args:
|
|
191
|
+
data (np.ndarray): Input image dataset of shape (num_samples, height, width).
|
|
192
|
+
downsize (int): Spatial downsampling factor via block mean pooling.
|
|
193
|
+
batch_size (int, optional): Number of samples per batch. Defaults to 32.
|
|
194
|
+
use_multiprocessing (bool, optional): Whether to use multiple CPU cores
|
|
195
|
+
(Ignored if backend='gpu'). Defaults to False.
|
|
196
|
+
|
|
197
|
+
Returns:
|
|
198
|
+
np.ndarray: Flattened feature embeddings of shape (num_samples, feature_dim).
|
|
199
|
+
"""
|
|
200
|
+
xp = self.xp
|
|
201
|
+
data = np.asarray(data, dtype=np.float32)
|
|
202
|
+
|
|
203
|
+
num_samples = data.shape[0]
|
|
204
|
+
d_size = int(self.size / downsize)
|
|
205
|
+
if (self.size % downsize) != 0:
|
|
206
|
+
raise ValueError(
|
|
207
|
+
f"Image size ({self.size}) must be exactly divisible by downsize ({downsize})."
|
|
208
|
+
)
|
|
209
|
+
|
|
210
|
+
order_1_keys = self.order_1_keys
|
|
211
|
+
num_order_1_maps = self.num_order_1_maps
|
|
212
|
+
num_order_2_maps = self.num_order_2_maps
|
|
213
|
+
|
|
214
|
+
def _process_batch(batch_data: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
|
215
|
+
b_xp: Any = cp if self.backend == "gpu" else np
|
|
216
|
+
b_data = (
|
|
217
|
+
b_xp.asarray(batch_data, dtype=b_xp.float32)
|
|
218
|
+
if self.backend == "gpu"
|
|
219
|
+
else batch_data
|
|
220
|
+
)
|
|
221
|
+
curr_batch_size = b_data.shape[0]
|
|
222
|
+
|
|
223
|
+
batch_fft = b_xp.fft.fftshift(self._fft2(b_data, b_xp), axes=(-2, -1)).astype(
|
|
224
|
+
b_xp.complex64
|
|
225
|
+
)
|
|
226
|
+
|
|
227
|
+
low_pass_c = self._low_pass_c
|
|
228
|
+
|
|
229
|
+
# ORDER 0: plain low-pass response, no modulus non-linearity.
|
|
230
|
+
low_pass_spatial = self._ifft2(
|
|
231
|
+
b_xp.fft.ifftshift(batch_fft * low_pass_c, axes=(-2, -1)), b_xp
|
|
232
|
+
)
|
|
233
|
+
low_pass_down = (
|
|
234
|
+
b_xp.real(low_pass_spatial)
|
|
235
|
+
.reshape(-1, d_size, downsize, d_size, downsize)
|
|
236
|
+
.mean(axis=(2, 4))
|
|
237
|
+
)
|
|
238
|
+
order_0_res = low_pass_down.reshape(curr_batch_size, -1)
|
|
239
|
+
|
|
240
|
+
# ORDER 1: filter with every wavelet, take modulus, smooth, pool.
|
|
241
|
+
bank_1_stack = self._filters_1_stack # (num_filters, H, W)
|
|
242
|
+
modulus_fft_1 = self._filter_and_modulus(
|
|
243
|
+
b_xp, batch_fft[:, None, :, :], bank_1_stack[None, :, :, :]
|
|
244
|
+
)
|
|
245
|
+
batch_pooled_order_1 = self._smooth_and_pool(
|
|
246
|
+
b_xp,
|
|
247
|
+
modulus_fft_1,
|
|
248
|
+
low_pass_c[None, None, :, :],
|
|
249
|
+
curr_batch_size,
|
|
250
|
+
num_order_1_maps,
|
|
251
|
+
d_size,
|
|
252
|
+
downsize,
|
|
253
|
+
)
|
|
254
|
+
|
|
255
|
+
# ORDER 2: for each order-1 channel, filter again with only the
|
|
256
|
+
# wavelets of strictly lower angular index k
|
|
257
|
+
batch_pooled_order_2 = b_xp.zeros(
|
|
258
|
+
(curr_batch_size, d_size, d_size, num_order_2_maps), dtype=b_xp.float32
|
|
259
|
+
)
|
|
260
|
+
order_2_idx = 0
|
|
261
|
+
|
|
262
|
+
for i, key1 in enumerate(order_1_keys):
|
|
263
|
+
children = self._order2_children[key1]
|
|
264
|
+
if not children:
|
|
265
|
+
continue
|
|
266
|
+
|
|
267
|
+
filters_2_stack = self._filters_2_stack_by_key1[key1] # (n_children, H, W)
|
|
268
|
+
mod_fft_1_single = modulus_fft_1[:, i] # (batch, H, W)
|
|
269
|
+
|
|
270
|
+
modulus_fft_2 = self._filter_and_modulus(
|
|
271
|
+
b_xp, mod_fft_1_single[:, None, :, :], filters_2_stack[None, :, :, :]
|
|
272
|
+
)
|
|
273
|
+
n_children = len(children)
|
|
274
|
+
down = self._smooth_and_pool(
|
|
275
|
+
b_xp,
|
|
276
|
+
modulus_fft_2,
|
|
277
|
+
low_pass_c[None, None, :, :],
|
|
278
|
+
curr_batch_size,
|
|
279
|
+
n_children,
|
|
280
|
+
d_size,
|
|
281
|
+
downsize,
|
|
282
|
+
)
|
|
283
|
+
batch_pooled_order_2[..., order_2_idx : order_2_idx + n_children] = down
|
|
284
|
+
order_2_idx += n_children
|
|
285
|
+
|
|
286
|
+
return (
|
|
287
|
+
order_0_res,
|
|
288
|
+
batch_pooled_order_1.reshape(curr_batch_size, -1),
|
|
289
|
+
batch_pooled_order_2.reshape(curr_batch_size, -1),
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
# Build batches index ranges
|
|
293
|
+
batches = [
|
|
294
|
+
data[start : min(start + batch_size, num_samples)]
|
|
295
|
+
for start in range(0, num_samples, batch_size)
|
|
296
|
+
]
|
|
297
|
+
|
|
298
|
+
# Execution Selection
|
|
299
|
+
if self.backend == "cpu" and use_multiprocessing:
|
|
300
|
+
nb_cpu: int = os.cpu_count() or 1
|
|
301
|
+
n_jobs = max(1, nb_cpu - 1) if os.cpu_count() and nb_cpu > 1 else 1
|
|
302
|
+
results = self._run_multiprocess(batches, _process_batch, n_jobs)
|
|
303
|
+
else:
|
|
304
|
+
results = []
|
|
305
|
+
for batch in tqdm(
|
|
306
|
+
batches, desc=f"Generating Embeddings on {self.backend.upper()}", unit="batch"
|
|
307
|
+
):
|
|
308
|
+
results.append(_process_batch(batch))
|
|
309
|
+
|
|
310
|
+
# Re-assemble outputs from batch results
|
|
311
|
+
order_0_list, order_1_list, order_2_list = zip(*results, strict=True)
|
|
312
|
+
|
|
313
|
+
order_0_features = xp.concatenate(order_0_list, axis=0)
|
|
314
|
+
first_order_features = xp.concatenate(order_1_list, axis=0)
|
|
315
|
+
second_order_features = xp.concatenate(order_2_list, axis=0)
|
|
316
|
+
|
|
317
|
+
final_features = xp.concatenate(
|
|
318
|
+
(order_0_features, first_order_features, second_order_features), axis=1
|
|
319
|
+
)
|
|
320
|
+
self.final_features = final_features
|
|
321
|
+
|
|
322
|
+
if self.backend == "gpu":
|
|
323
|
+
return np.asarray(final_features.get())
|
|
324
|
+
|
|
325
|
+
return np.asarray(final_features)
|
|
326
|
+
|
|
327
|
+
def _run_multiprocess(
|
|
328
|
+
self, batches: list[np.ndarray], _process_batch: Any, n_jobs: int
|
|
329
|
+
) -> list[Any]:
|
|
330
|
+
"""Run batches across multiple processes with a progress bar"""
|
|
331
|
+
try:
|
|
332
|
+
# Joblib generator evaluation wrapped inside tqdm
|
|
333
|
+
with tqdm(
|
|
334
|
+
total=len(batches), desc="Generating Embeddings (Multiprocessing)", unit="batch"
|
|
335
|
+
) as pbar:
|
|
336
|
+
job_results = Parallel(n_jobs=n_jobs, return_as="generator")(
|
|
337
|
+
delayed(_process_batch)(b) for b in batches
|
|
338
|
+
)
|
|
339
|
+
results = []
|
|
340
|
+
for res in job_results:
|
|
341
|
+
results.append(res)
|
|
342
|
+
pbar.update(1)
|
|
343
|
+
except TypeError:
|
|
344
|
+
# Fallback for older joblib versions without generator support
|
|
345
|
+
results = []
|
|
346
|
+
for res in tqdm(
|
|
347
|
+
Parallel(n_jobs=n_jobs)(delayed(_process_batch)(b) for b in batches),
|
|
348
|
+
desc="Generating Embeddings (Multiprocessing)",
|
|
349
|
+
unit="batch",
|
|
350
|
+
):
|
|
351
|
+
results.append(res)
|
|
352
|
+
|
|
353
|
+
return results
|
|
354
|
+
|
|
355
|
+
def save_embeddings(self, path: str) -> None:
|
|
356
|
+
"""Save the generated feature embeddings to a compressed .npz file."""
|
|
357
|
+
if self.final_features is None:
|
|
358
|
+
raise ValueError("No embeddings found. Run generate_embeddings() first.")
|
|
359
|
+
|
|
360
|
+
m, k, sigma = self.bank.summary(verbose=False)
|
|
361
|
+
os.makedirs(path, exist_ok=True)
|
|
362
|
+
save_path = rf"{path}/embedding_m{m}_k{k}_sigma{sigma}.npz"
|
|
363
|
+
|
|
364
|
+
features = self.final_features.get() if self.backend == "gpu" else self.final_features
|
|
365
|
+
np.savez_compressed(save_path, embedding=features)
|
|
366
|
+
|
|
367
|
+
logger.info("Embedding successfully saved to '%s'", save_path)
|
|
368
|
+
|
|
369
|
+
def visualise_maps(self, image: np.ndarray, downsize: int) -> None:
|
|
370
|
+
"""
|
|
371
|
+
Visualises Order 0 and Order 1 scattering maps.
|
|
372
|
+
|
|
373
|
+
Args:
|
|
374
|
+
image (np.ndarray): A single 2D image array (height, width).
|
|
375
|
+
downsize (int): Spatial downsampling factor. Set to 1 for no downsampling.
|
|
376
|
+
"""
|
|
377
|
+
|
|
378
|
+
if image.ndim != 2:
|
|
379
|
+
raise ValueError("Please provide a single 2D image array of shape (height, width).")
|
|
380
|
+
|
|
381
|
+
d_size = int(self.size / downsize)
|
|
382
|
+
|
|
383
|
+
# Add fake batch dimension to match FFT logic
|
|
384
|
+
batch = image[None, ...]
|
|
385
|
+
batch_fft = np.fft.fftshift(np.fft.fft2(batch, axes=(-2, -1)), axes=(-2, -1))
|
|
386
|
+
|
|
387
|
+
# Dictionaries to hold maps for plotting
|
|
388
|
+
maps = {}
|
|
389
|
+
|
|
390
|
+
# ORDER 0
|
|
391
|
+
low_pass_coeffs = batch_fft * self.low_pass
|
|
392
|
+
low_pass_spatial = np.real(
|
|
393
|
+
np.fft.ifft2(np.fft.ifftshift(low_pass_coeffs, axes=(-2, -1)), axes=(-2, -1))
|
|
394
|
+
)[0]
|
|
395
|
+
|
|
396
|
+
low_pass_down = low_pass_spatial.reshape(d_size, downsize, d_size, downsize).mean(
|
|
397
|
+
axis=(1, 3)
|
|
398
|
+
)
|
|
399
|
+
|
|
400
|
+
maps["Order 0 (Low Pass)"] = low_pass_down
|
|
401
|
+
|
|
402
|
+
# ORDER 1
|
|
403
|
+
order_1_keys = self.bank_keys[1:]
|
|
404
|
+
|
|
405
|
+
for key1 in order_1_keys:
|
|
406
|
+
wavelet_fft = self.bank[key1]
|
|
407
|
+
|
|
408
|
+
# Standard scattering cascade
|
|
409
|
+
filtered_fft = batch_fft * wavelet_fft
|
|
410
|
+
spatial_complex = np.fft.ifft2(
|
|
411
|
+
np.fft.ifftshift(filtered_fft, axes=(-2, -1)), axes=(-2, -1)
|
|
412
|
+
)
|
|
413
|
+
modulus_spatial = np.abs(spatial_complex)
|
|
414
|
+
|
|
415
|
+
modulus_fft = np.fft.fftshift(
|
|
416
|
+
np.fft.fft2(modulus_spatial, axes=(-2, -1)), axes=(-2, -1)
|
|
417
|
+
)
|
|
418
|
+
filtered_low_pass = modulus_fft * self.low_pass
|
|
419
|
+
|
|
420
|
+
smoothed_spatial = np.real(
|
|
421
|
+
np.fft.ifft2(np.fft.ifftshift(filtered_low_pass, axes=(-2, -1)), axes=(-2, -1))
|
|
422
|
+
)[0]
|
|
423
|
+
|
|
424
|
+
# Downsample
|
|
425
|
+
smoothed_down = smoothed_spatial.reshape(d_size, downsize, d_size, downsize).mean(
|
|
426
|
+
axis=(1, 3)
|
|
427
|
+
)
|
|
428
|
+
|
|
429
|
+
maps[f"Order 1 ({key1})"] = smoothed_down
|
|
430
|
+
|
|
431
|
+
# DYNAMIC GRID CALCULATION
|
|
432
|
+
num_maps = len(maps)
|
|
433
|
+
|
|
434
|
+
# Find the best exact integer factors
|
|
435
|
+
best_factor = 1
|
|
436
|
+
for i in range(1, int(np.sqrt(num_maps)) + 1):
|
|
437
|
+
if num_maps % i == 0:
|
|
438
|
+
best_factor = i
|
|
439
|
+
|
|
440
|
+
rows = best_factor
|
|
441
|
+
cols = num_maps // best_factor
|
|
442
|
+
|
|
443
|
+
# Fallback for prime numbers or extremely stretched grids
|
|
444
|
+
# If the aspect ratio is wider than 3:1, use a square-ish grid instead.
|
|
445
|
+
if cols / rows > 3:
|
|
446
|
+
cols = int(np.ceil(np.sqrt(num_maps)))
|
|
447
|
+
rows = int(np.ceil(num_maps / cols))
|
|
448
|
+
|
|
449
|
+
# PLOTTING
|
|
450
|
+
_, axes = plt.subplots(rows, cols, figsize=(cols * 3, rows * 3))
|
|
451
|
+
|
|
452
|
+
# Flatten the axes array so we can iterate through it easily
|
|
453
|
+
if isinstance(axes, np.ndarray):
|
|
454
|
+
axes = axes.flatten()
|
|
455
|
+
else:
|
|
456
|
+
axes = [axes] # Catches the edge case where num_maps == 1
|
|
457
|
+
|
|
458
|
+
for idx, (title, map_data) in enumerate(maps.items()):
|
|
459
|
+
axes[idx].imshow(map_data, cmap="inferno")
|
|
460
|
+
axes[idx].set_title(title, fontsize=10)
|
|
461
|
+
axes[idx].axis("off")
|
|
462
|
+
|
|
463
|
+
# Turn off the axes for any leftover empty subplots from the prime-number fallback
|
|
464
|
+
for empty_idx in range(num_maps, len(axes)):
|
|
465
|
+
axes[empty_idx].axis("off")
|
|
466
|
+
|
|
467
|
+
plt.tight_layout()
|
|
468
|
+
plt.show()
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: fbscatnet
|
|
3
|
+
Version: 1.0.0
|
|
4
|
+
Summary: Fourier-Bessel wavelet scattering transforms
|
|
5
|
+
Project-URL: Homepage, https://github.com/Smee18/FourierBesselWavelets
|
|
6
|
+
Project-URL: Repository, https://github.com/Smee18/FourierBesselWavelets.git
|
|
7
|
+
Project-URL: Issues, https://github.com/Smee18/FourierBesselWavelets/issues
|
|
8
|
+
Author-email: Marcel Venturotti <mv514@bath.ac.uk>
|
|
9
|
+
License: MIT
|
|
10
|
+
License-File: LICENSE
|
|
11
|
+
Keywords: fourier-bessel,image-processing,machine-learning,scattering-transform,wavelets
|
|
12
|
+
Classifier: Development Status :: 5 - Production/Stable
|
|
13
|
+
Classifier: Intended Audience :: Science/Research
|
|
14
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
15
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
16
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
17
|
+
Classifier: Programming Language :: Python :: 3.12
|
|
18
|
+
Classifier: Topic :: Scientific/Engineering :: Image Processing
|
|
19
|
+
Classifier: Topic :: Scientific/Engineering :: Mathematics
|
|
20
|
+
Requires-Python: >=3.10
|
|
21
|
+
Requires-Dist: joblib>=1.3
|
|
22
|
+
Requires-Dist: matplotlib>=3.7
|
|
23
|
+
Requires-Dist: numpy>=1.24
|
|
24
|
+
Requires-Dist: scipy>=1.10
|
|
25
|
+
Requires-Dist: tqdm>=4.65
|
|
26
|
+
Provides-Extra: dev
|
|
27
|
+
Requires-Dist: build; extra == 'dev'
|
|
28
|
+
Requires-Dist: mypy; extra == 'dev'
|
|
29
|
+
Requires-Dist: pre-commit; extra == 'dev'
|
|
30
|
+
Requires-Dist: pytest; extra == 'dev'
|
|
31
|
+
Requires-Dist: ruff; extra == 'dev'
|
|
32
|
+
Requires-Dist: scipy-stubs; extra == 'dev'
|
|
33
|
+
Requires-Dist: types-tqdm; extra == 'dev'
|
|
34
|
+
Provides-Extra: docs
|
|
35
|
+
Requires-Dist: sphinx; extra == 'docs'
|
|
36
|
+
Requires-Dist: sphinx-autobuild; extra == 'docs'
|
|
37
|
+
Requires-Dist: sphinx-rtd-theme; extra == 'docs'
|
|
38
|
+
Provides-Extra: gpu
|
|
39
|
+
Requires-Dist: cupy-cuda12x>=11.0; extra == 'gpu'
|
|
40
|
+
Description-Content-Type: text/markdown
|
|
41
|
+
|
|
42
|
+
# fbscatnet
|
|
43
|
+
|
|
44
|
+
[](https://pypi.org/project/fbscatnet/)
|
|
45
|
+
[](https://pypi.org/project/fbscatnet/)
|
|
46
|
+
[](https://opensource.org/licenses/MIT)
|
|
47
|
+
[](https://github.com/Smee18/FourierBesselWavelets/actions/workflows/ci.yml)
|
|
48
|
+
[](https://github.com/astral-sh/ruff)
|
|
49
|
+
|
|
50
|
+
`fbscatnet` is a high-performance Python library for computing **Fourier-Bessel wavelet scattering transforms**. It generates robust, feature embeddings from 2D images, making it an ideal feature extractor for computer vision, biomedical imaging, and physics-based machine learning. This is a brand new project of mine that I have been working on for some time. I would love any criticism, improvements, corrections to both the code and the maths.
|
|
51
|
+
|
|
52
|
+
## Features
|
|
53
|
+
|
|
54
|
+
- **Novel Wavelet:** New wavelet relying on Bessel basis functions.
|
|
55
|
+
- **Hardware Accelerated:** Seamlessly switch between multi-core CPU execution (`scipy`/`joblib`) and GPU acceleration (`cupy`).
|
|
56
|
+
- **Highly Optimized:** Vectorised filtering in the frequency domain for maximum throughput across large image batches.
|
|
57
|
+
- **ML-Ready:** Outputs flattened feature arrays directly compatible with `scikit-learn`, `xgboost`, or `pytorch`.
|
|
58
|
+
|
|
59
|
+
## Installation
|
|
60
|
+
|
|
61
|
+
Install the base package (CPU-only) via pip:
|
|
62
|
+
|
|
63
|
+
```bash
|
|
64
|
+
pip install fbscatnet
|
|
65
|
+
```
|
|
66
|
+
|
|
67
|
+
To enable **GPU acceleration**, install with the `gpu` extra (requires a CUDA-compatible GPU):
|
|
68
|
+
|
|
69
|
+
```bash
|
|
70
|
+
pip install fbscatnet[gpu]
|
|
71
|
+
```
|
|
72
|
+
|
|
73
|
+
## Quickstart
|
|
74
|
+
|
|
75
|
+
Extracting features from a dataset takes just a few lines of code:
|
|
76
|
+
|
|
77
|
+
```python
|
|
78
|
+
import numpy as np
|
|
79
|
+
import logging
|
|
80
|
+
from fbscatnet import FourierBesselWaveletBank, FourierBesselScatNet, logger_config
|
|
81
|
+
logger_config.enable_colored_logs(logging.DEBUG) # IMPORTANT: SET TO SEE LOGS
|
|
82
|
+
|
|
83
|
+
# 1. Create some dummy image data (e.g., 10 grayscale images of size 64x64)
|
|
84
|
+
images = np.random.rand(10, 64, 64)
|
|
85
|
+
|
|
86
|
+
# 2. Instantiate a Fourier-Bessel Wavelet Bank
|
|
87
|
+
# size: spatial dimension (64x64), m: angular order, k: radial roots
|
|
88
|
+
bank = FourierBesselWaveletBank(size=64, m=2, k=2, sigma=0.1)
|
|
89
|
+
|
|
90
|
+
# 3. Initialize the Scattering Network
|
|
91
|
+
# Use backend="gpu" if you installed with CuPy
|
|
92
|
+
net = FourierBesselScatNet(bank=bank, backend="cpu")
|
|
93
|
+
|
|
94
|
+
# 4. Generate feature embeddings
|
|
95
|
+
# downsize: spatial pooling factor (must evenly divide the image size)
|
|
96
|
+
features = net.generate_embeddings(images, downsize=4, use_multiprocessing=True)
|
|
97
|
+
|
|
98
|
+
print(f"Generated embeddings shape: {features.shape}")
|
|
99
|
+
# Output: (10, feature_dimension)
|
|
100
|
+
```
|
|
101
|
+
|
|
102
|
+
## Visualising Wavelet Maps
|
|
103
|
+
|
|
104
|
+
If you want to inspect how the wavelet filters are interacting with your image at the first scattering order, `fbscatnet` includes a built-in plotting tool:
|
|
105
|
+
|
|
106
|
+
```python
|
|
107
|
+
# Pass a single 2D image to visualize
|
|
108
|
+
single_image = images[0]
|
|
109
|
+
net.visualise_maps(single_image, downsize=4)
|
|
110
|
+
```
|
|
111
|
+
|
|
112
|
+
## API Overview
|
|
113
|
+
|
|
114
|
+
Sphinx documentation at: https://smee18.github.io/FourierBesselWavelets/
|
|
115
|
+
|
|
116
|
+
The `example` folder also contains a full pipeline, classifying MNIST using the library
|
|
117
|
+
The pdf of my notes taken along this project contains an in-depth explanation of the mathematics behind these wavelets. Here you can find all derivations, proofs and useful information. Some stuff might seem trivial but my goal is to really expose all aspects so anyone new to wavelet theory can understand the mathematics behind these Fourier Bessel wavelets.
|
|
118
|
+
|
|
119
|
+
## License
|
|
120
|
+
This project is licensed under the MIT License.
|
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
fbscatnet/__init__.py,sha256=WPRlnUw-NS_Xo0l-ihj0KrhlJF6Niqydc3nNN73Sjz0,339
|
|
2
|
+
fbscatnet/generate_bank.py,sha256=Qo_ymKoxqT82yjr0BGGdM2Iy9bc22309kEzidKXCwF8,8394
|
|
3
|
+
fbscatnet/logger_config.py,sha256=SPLTnoE05nkkN7fGTyIlcqrut7bNlfyD-KOBJT9OV24,1594
|
|
4
|
+
fbscatnet/math_utils.py,sha256=rzB_zoPlpdNgOx5O-pvzFGN4JeF_qg_OzoykkHiXIs4,6906
|
|
5
|
+
fbscatnet/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
|
|
6
|
+
fbscatnet/scatnet.py,sha256=nRZHlMRdk5pBvSJUZyJKfn5i3hFqW4NCYNBYwnth4ts,18560
|
|
7
|
+
fbscatnet-1.0.0.dist-info/METADATA,sha256=N4Yx345ZkKPQMx2dXg73orYKOfRb-tXqdVLl57OqBI8,5373
|
|
8
|
+
fbscatnet-1.0.0.dist-info/WHEEL,sha256=lCkmxWfQsSc9CfIClYeavTdQeEX2toPqufh9gI35EQA,87
|
|
9
|
+
fbscatnet-1.0.0.dist-info/licenses/LICENSE,sha256=ppcDcZqSfk-QMGl6mrPhtOb-5fVTK45M0rEN-WqOhEg,1074
|
|
10
|
+
fbscatnet-1.0.0.dist-info/RECORD,,
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Marcel Venturotti
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|