spd_learn.functional.compute_gabor_wavelet#

spd_learn.functional.compute_gabor_wavelet(tt: Tensor, foi: Tensor, fwhm: Tensor, sfreq: float = 250.0, scaling: str = 'oct', dtype: dtype = torch.complex64, min_foi_oct: float = -2.0, max_foi_oct: float = 6.0, min_fwhm_oct: float = -6.0, max_fwhm_oct: float = 1.0) → Tensor[source]#

Compute a complex Gabor (Morlet) wavelet filterbank.

Creates a bank of complex-valued Gabor wavelets with learnable center frequencies and temporal resolutions. The wavelets are L2-normalized and optionally scaled for octave-based frequency analysis.

Parameters:
  • tt (Tensor) – Time vector with shape (kernel_length,), typically centered at 0.

  • foi (Tensor) – Center frequencies in octaves (log2 Hz) with shape (n_wavelets,). For example, foi=3.0 corresponds to 2^3 = 8 Hz.

  • fwhm (Tensor) – Full Width at Half Maximum in octaves with shape (n_wavelets,). Controls the temporal resolution of each wavelet.

  • sfreq (float, default=250.0) – Sampling frequency in Hz.

  • scaling (str, default="oct") – Scaling method. If “oct”, applies octave-based amplitude scaling to achieve constant energy per octave.

  • dtype (torch.dtype, default=torch.complex64) – Output data type for the complex wavelets.

  • min_foi_oct (float, default=-2.0) – Minimum clamp value for center frequencies (in octaves).

  • max_foi_oct (float, default=6.0) – Maximum clamp value for center frequencies (in octaves).

  • min_fwhm_oct (float, default=-6.0) – Minimum clamp value for FWHM (in octaves).

  • max_fwhm_oct (float, default=1.0) – Maximum clamp value for FWHM (in octaves).

Returns:

Complex wavelet filterbank with shape (n_wavelets, kernel_length).

Return type:

Tensor

Notes

The Gabor wavelet is defined as:

\[\psi(t) = \exp(2\pi i f t) \cdot \exp\left(-\frac{4 \ln 2 \cdot t^2}{h^2}\right)\]

where \(f\) is the center frequency and \(h\) is the FWHM of the Gaussian envelope.

See also

WaveletConv

Module wrapper using this function.

References

See [Paillard et al., 2025] for details on learnable wavelet filterbanks.

Examples

>>> import torch
>>> from spd_learn.functional import compute_gabor_wavelet
>>> # Create time vector for 0.5s kernel at 250 Hz
>>> tt = torch.linspace(-0.25, 0.25, 125)
>>> # Frequencies: 4, 8, 16 Hz in octave notation
>>> foi = torch.tensor([2.0, 3.0, 4.0])
>>> fwhm = torch.tensor([-2.0, -3.0, -4.0])
>>> wavelets = compute_gabor_wavelet(tt, foi, fwhm, sfreq=250)
>>> wavelets.shape
torch.Size([3, 125])