Coverage for src/crispatt/_stft.py: 38%
30 statements
« prev ^ index » next coverage.py v7.15.3, created at 2026-09-03 17:23 +1200
« prev ^ index » next coverage.py v7.15.3, created at 2026-09-03 17:23 +1200
1import warnings
3import attrs
4import numpy as np
5import scipy.signal as signal
8def _round_power_of_2(number: float) -> int:
9 # if number <= 0:
10 # return
11 return 2 ** np.ceil(np.log2(number)).astype(int)
14@attrs.define
15class LocalFT:
16 x: np.ndarray
17 stft: signal.ShortTimeFFT
19 @classmethod
20 def from_stft_params(
21 cls,
22 x: np.ndarray,
23 window_length: float,
24 hop_length: float,
25 padding_factor: int,
26 ):
27 spacing = x[1] - x[0] # Spacing is expected to be regular for FFT.
28 stft = cls._compute_stft(
29 spacing, window_length, hop_length, padding_factor=padding_factor
30 )
31 return cls(x, stft)
33 @classmethod
34 def _compute_stft(
35 cls,
36 spacing: float,
37 window_length: float,
38 hop_length: float,
39 padding_factor: int,
40 ):
41 window_width = _round_power_of_2(window_length / spacing)
42 hop = np.floor(hop_length / spacing).astype(int)
43 if hop < 1:
44 warnings.warn("STFT hop lower than 1, set to 1.", stacklevel=2)
45 hop = 1
46 window = np.hanning(window_width)
47 stft = signal.ShortTimeFFT(
48 window,
49 hop=hop,
50 fs=1 / spacing,
51 mfft=window_width * padding_factor,
52 scale_to="magnitude",
53 )
54 return stft
56 def envelope_from_fft(self, deflection: np.ndarray):
57 n_samples = self.x.size
58 tmask = (
59 self.stft.t(n_samples) >= self.stft.lower_border_end[0] * self.stft.T
60 ) & (
61 self.stft.t(n_samples)
62 <= self.stft.upper_border_begin(n_samples)[0] * self.stft.T
63 )
64 tott = self.stft.t(n_samples)
65 masked_t = tott[tmask]
66 masked_magnitude = np.max(np.abs(self.stft.stft(deflection)), axis=0)[tmask]
67 return masked_t, masked_magnitude