Bug
FilterBankLayer._apply_iir (braindecode/modules/filter.py) downcasts the float64 filter coefficients to the input dtype before torchaudio.functional.filtfilt:
filtered = filtfilt(
x,
a_coeffs=a_coeffs.type_as(x).to(x.device), # float64 coeffs -> float32
b_coeffs=b_coeffs.type_as(x).to(x.device),
clamp=False,
)
For an IIR band-pass with poles near the unit circle (e.g. a 4th-order Chebyshev-II bank, as used by FBMSNet/FBCNet), the recursion diverges in float32 and the whole forward becomes NaN from the first batch (train loss NaN, accuracy at chance). Verified: NaN in float32, stable in float64.
Reproduce
import torch
from braindecode.modules.filter import FilterBankLayer
fb = FilterBankLayer(n_chans=22, sfreq=250, method="iir",
iir_params=dict(order=4, ftype="cheby2", rs=30, output="ba"))
out = fb(torch.randn(4, 22, 1000)) # float32
print(torch.isnan(out).any()) # tensor(True)
Fix
Run the recursion in float64 and cast back (the coefficient buffers are already float64):
orig_dtype = x.dtype
filtered = filtfilt(x.double(),
a_coeffs=a_coeffs.double().to(x.device),
b_coeffs=b_coeffs.double().to(x.device), clamp=False)
return filtered.to(orig_dtype).unsqueeze(1)
This makes any IIR filter bank (cheby2 etc.) usable. (Separately, the IIR path is restricted to 'ba' coefficients — SOS would be more robust for higher orders, but the float64 cast already fixes the NaN.) Happy to PR.
Bug
FilterBankLayer._apply_iir(braindecode/modules/filter.py) downcasts the float64 filter coefficients to the input dtype beforetorchaudio.functional.filtfilt:For an IIR band-pass with poles near the unit circle (e.g. a 4th-order Chebyshev-II bank, as used by
FBMSNet/FBCNet), the recursion diverges in float32 and the whole forward becomes NaN from the first batch (train loss NaN, accuracy at chance). Verified: NaN in float32, stable in float64.Reproduce
Fix
Run the recursion in float64 and cast back (the coefficient buffers are already float64):
This makes any IIR filter bank (cheby2 etc.) usable. (Separately, the IIR path is restricted to
'ba'coefficients — SOS would be more robust for higher orders, but the float64 cast already fixes the NaN.) Happy to PR.