Skip to content

FilterBankLayer IIR path NaNs in float32 (coefficients downcast before filtfilt) #1067

Description

@bruAristimunha

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions