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

1import warnings 

2 

3import attrs 

4import numpy as np 

5import scipy.signal as signal 

6 

7 

8def _round_power_of_2(number: float) -> int: 

9 # if number <= 0: 

10 # return 

11 return 2 ** np.ceil(np.log2(number)).astype(int) 

12 

13 

14@attrs.define 

15class LocalFT: 

16 x: np.ndarray 

17 stft: signal.ShortTimeFFT 

18 

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) 

32 

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 

55 

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