Repository navigation
Models run on Intel Gaudi (HPU): portable ops for the remaining failures - #1249
Merged
bruAristimunha merged 10 commits intoOct 8, 2026
Merged
Conversation
…nce_mode FFTs on CPU have no bfloat16 kernel: BIOT, BrainBERT, CBraMod, CodeBrain, EEGDINO and LUNA transform in at least float32 and cast back. Tensors built in forward follow the input dtype (BENDR start token, CodeBrain attention mask, LUNA NeRF encoding, REVE position embedding, ZUNA residual stream). CodeBrain updates its lazily initialised kernel norm in place, so a first forward under torch.inference_mode() no longer breaks training. Float32 state dicts, outputs and gradients are unchanged.
…least, CPU for HPU PyTorch has no complex bfloat16 dtype and Intel Gaudi has no complex tensors. spectral_input(x) returns x in promote_types(x.dtype, float32), moved to the CPU when it lives on an HPU; real results go back with .to(x). Replaces the hand casts in BIOT, BrainBERT, CBraMod, CodeBrain, EEGDINO and LUNA, and is applied to Brant, ContraWR, DIVER1, EMG2QwertyNet, MAPA, MetaNeuromotorHand, SensingDynamics, FilterBankLayer (FBCNet, FBMSNet, FBLightConvNet, IFNet), GeneralizedGaussianFilter and hilbert_freq. SyncNet and DGCNN follow the input dtype. test_forward_in_dtype covers every model in float64, bfloat16 and float16.
…ripts EMG2QwertyNet passes the Spectrogram settings to torchaudio's functional spectrogram explicitly (TorchScript cannot read the submodule's pad). IFNet float16 passes, so it is not an expected failure. whats_new lists every model the helper covers.
test_foundation_models.py: keep both the CodeBrain inference-mode test and master's SleepFM tests. SleepFM/SleepFMStager pass test_forward_in_dtype.
…(Gaudi lazy compile)
Contributor
There was a problem hiding this comment.
🟡 Changes recommended
Float16 Hilbert output is not cast back to its input dtype, and the release note overstates supported dtype combinations.
2 open findings
What changed in this PR
Adds dtype-portable spectral operations and an Intel Gaudi lazy-mode workaround while preserving model output dtypes.
Changes:
- Introduces
spectral_inputfor portable FFT/STFT preprocessing. - Updates affected models and filter modules for dtype/device consistency.
- Adds broad dtype compatibility and regression tests.
| File | Description |
|---|---|
braindecode/functional/functions.py |
Adds spectral input promotion and updates Hilbert processing. |
braindecode/functional/__init__.py |
Exports spectral_input. |
braindecode/modules/filter.py |
Makes filter operations dtype/device portable. |
braindecode/models/bendr.py |
Types the start token from the input. |
braindecode/models/biot.py |
Makes STFT processing portable. |
braindecode/models/brainbert.py |
Promotes STFT input and restores output dtype. |
braindecode/models/brainomni.py |
Uses the Gaudi-compatible 4D LSTM permutation. |
braindecode/models/brant.py |
Makes band-power FFT processing portable. |
braindecode/models/cbramod.py |
Promotes spectral embedding input. |
braindecode/models/codebrain.py |
Fixes buffer initialization and spectral dtype handling. |
braindecode/models/contrawr.py |
Makes STFT processing portable. |
braindecode/models/dgcnn.py |
Matches identity-matrix dtype to adjacency. |
braindecode/models/diver1.py |
Makes spectral embedding portable. |
braindecode/models/eegdino.py |
Promotes FFT input and restores dtype. |
braindecode/models/emg2qwerty.py |
Reworks spectrogram computation for portability. |
braindecode/models/luna.py |
Corrects positional and frequency-feature dtypes. |
braindecode/models/mapa.py |
Makes multi-band STFT processing portable. |
braindecode/models/meta_neuromotor.py |
Moves spectral features through portable preprocessing. |
braindecode/models/reve.py |
Aligns positional embeddings with signal dtype. |
braindecode/models/sensingdynamics.py |
Makes low-pass FFT filtering portable. |
braindecode/models/syncnet.py |
Removes forced float32 convolution. |
braindecode/models/zuna.py |
Separates residual and sublayer precision. |
docs/api.rst |
Documents the new public helper. |
docs/whats_new.rst |
Adds release notes for dtype and HPU support. |
test/unit_tests/models/test_foundation_models.py |
Tests CodeBrain after inference mode. |
test/unit_tests/models/test_integration.py |
Adds model-wide dtype compatibility coverage. |
test/unit_tests/models/test_pretrained_compat.py |
Tests compatible pretrained models in additional dtypes. |
🧠 Review effort: Balanced
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| # not implement bfloat16 FFT kernels. | ||
| if input_dtype == torch.bfloat16: | ||
| x = x.float() | ||
| x = spectral_input(x) |
Comment on lines
+221
to
+223
| - Models now run after ``model.to(torch.float64)``, ``torch.bfloat16`` or | ||
| ``torch.float16``, and their FFT, STFT and filter-bank front ends run on Intel | ||
| Gaudi (HPU): the new :func:`braindecode.functional.spectral_input` gives these |
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## master #1249 +/- ##
==========================================
- Coverage 89.05% 89.02% -0.04%
==========================================
Files 157 157
Lines 19660 19398 -262
==========================================
- Hits 17509 17269 -240
+ Misses 2151 2129 -22 🚀 New features to boost your workflow:
|
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.


Per model, root cause -> equivalent op:
BrainOmni, BrainTokenizer (
_SEANetLSTM): Gaudi lazy mode fails to compile an LSTM that reads a 3D-permuted Conv1d output (synStatus 26;.contiguous(),transpose,einsumfail the same way) -> permute a 4D view,x.unsqueeze(-1).permute(2, 0, 1, 3).squeeze(-1). Same values; lazy forward now runs in fp32 and bf16.EMG2QwertyNet, MetaNeuromotorHand (rotation-invariant MLP): Gaudi eager mode returns wrong values for
Tensor.rollon a non-contiguous tensor (max-abs 4.6 on randn; the input is amovedimview) -> rollinputs.contiguous(). Eager fp32 now matches CPU (2.5e-6, was 1.6 on 2.3) and bf16 is finite (was NaN).Left as Habana issues (no portable one-line op; minimal repros kept, eager mode works unless noted; the
nn.SELUbackward gap is fixed in #1253):gaudi2_agu_config.cpp:338 size - 1 <= uint8 max).HPU check (Gaudi2, all 83 models, forward + one train step, fp32 + bf16, lazy and eager), compared with master before this PR: lazy, BrainOmni and BrainTokenizer forward now run (training on HPU: #1253); eager, EMG2QwertyNet and MetaNeuromotorHand now match CPU in fp32 and are finite in bf16 (no bf16 NaN left in eager). No model got worse.
Float32 outputs, gradients and state_dict on CPU are bit-identical to #1246 for every touched model; TorchScript tests pass.