Repository navigation
Models run in float64 and bfloat16; CodeBrain trains after inference_mode - #1246
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.
There was a problem hiding this comment.
Warning
Copilot couldn't run its full agentic review because it didn't start before the timeout. Make sure your repository has a runner available, or add a copilot-code-review.yml file specifying one with the runs-on attribute. See the docs for more details.
Copilot review overview
2 open findings
What changed in this PR
Fixes dtype robustness across several braindecode models so they can run correctly after model.to(torch.float64) / model.to(torch.bfloat16) (especially CPU FFT/STFT paths), and ensures CodeBrain can train even if its first forward happened under torch.inference_mode().
Changes:
- Promote FFT/STFT inputs to at least float32 on CPU and cast results back to the model/input dtype across multiple models.
- Update CodeBrain’s lazy buffer initialization to be in-place, avoiding inference tensors ending up in buffers.
- Add regression tests for forward-in-dtype compatibility and CodeBrain training after inference-mode forward; document the fixes in “What’s New”.
| File | Description |
|---|---|
| test/unit_tests/models/test_pretrained_compat.py | Adds a dtype-parameterized forward test for all COMPAT models. |
| test/unit_tests/models/test_foundation_models.py | Adds regression test for CodeBrain training after torch.inference_mode() forward. |
| docs/whats_new.rst | Documents dtype / inference-mode training fixes for affected models. |
| braindecode/models/zuna.py | Adjusts residual/activation dtypes to avoid forcing float32 residual into float64/bfloat16 linears. |
| braindecode/models/reve.py | Ensures positional embeddings follow signal dtype while keeping Fourier positions float32. |
| braindecode/models/luna.py | Promotes FFT to float32+ on CPU and casts back; casts NeRF encoding back to coord dtype. |
| braindecode/models/eegdino.py | Promotes rfft to float32+ on CPU and casts magnitude back. |
| braindecode/models/codebrain.py | Promotes FFTs appropriately, casts outputs back, makes kernel-norm buffer init in-place, and aligns attention mask dtype. |
| braindecode/models/cbramod.py | Promotes rfft to float32+ on CPU and casts magnitude back. |
| braindecode/models/brainbert.py | Promotes spectrogram FFT inputs to float32+ and casts final magnitude back to input dtype. |
| braindecode/models/biot.py | Runs STFT in float32+ on CPU and casts magnitude back. |
| braindecode/models/bendr.py | Ensures constructed start token follows input dtype to prevent unintended promotion. |
🧠 Review effort: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| self.kernel_norm.copy_(kernel.norm(dim=-1, keepdim=True).detach()) | ||
| self.kernel_norm_initialized.fill_(True) |
| dtype = self.attention.wq.weight.dtype | ||
| residual_dtype = torch.promote_types(dtype, torch.float32) | ||
| input_tensor = input_tensor.to(residual_dtype) | ||
| hidden_states = input_tensor + self.attention_norm_post( | ||
| self.attention( | ||
| self.attention_norm(input_tensor), rotary_cosine, rotary_sine | ||
| ).float() | ||
| self.attention_norm(input_tensor).to(dtype), rotary_cosine, rotary_sine | ||
| ).to(residual_dtype) |
| encoded = torch.cat([encoded, pad], dim=-1) | ||
| return encoded | ||
| # Sin/cos run in float32 at least (float32 frequency bands); return coords' dtype. | ||
| return encoded.to(coords.dtype) |
Codecov Report❌ Patch coverage is Additional details and impacted files@@ Coverage Diff @@
## master #1246 +/- ##
==========================================
+ Coverage 88.71% 88.73% +0.01%
==========================================
Files 157 157
Lines 19654 19676 +22
==========================================
+ Hits 17437 17459 +22
Misses 2217 2217 🚀 New features to boost your workflow:
|
…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.
There was a problem hiding this comment.
🟡 Changes recommended
hilbert_freq still returns float32 for float16 input, and the release notes overstate compatibility.
6 open findings
Cast supplied channel locations to signal device and dtype Fordtype=torch.bfloat16,residual_dtypebecomestorch.float32, soinput_tensoris float32…copy_requiresself.kernel_normto already be allocated with the exact same shape as… Restore float16 dtype after spectral input promotion · New Correct universal float16 compatibility claim · New Correct HPU fallback description for DFT APIs · New
🧠 Review effort: Balanced
| # not implement bfloat16 FFT kernels. | ||
| if input_dtype == torch.bfloat16: | ||
| x = x.float() | ||
| x = spectral_input(x) |
| - 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 |
| (FBCNet, FBMSNet, FBLightConvNet, IFNet), :class:`braindecode.modules.GeneralizedGaussianFilter` | ||
| and :func:`braindecode.functional.hilbert_freq`. Tensors built inside ``forward`` of |
test_foundation_models.py: keep both the CodeBrain inference-mode test and master's SleepFM tests. SleepFM/SleepFMStager pass test_forward_in_dtype.
There was a problem hiding this comment.
🟡 Changes recommended
Float16 hilbert_freq output is incorrectly left as float32, and the release note overstates full-model float16 support.
7 open findings
Cast supplied channel locations to signal device and dtype Fordtype=torch.bfloat16,residual_dtypebecomestorch.float32, soinput_tensoris float32…copy_requiresself.kernel_normto already be allocated with the exact same shape as… Restore float16 dtype after spectral input promotion Narrow float16 support claim to spectral front ends · New Correct HPU fallback description for DFT APIs Correct universal float16 compatibility claim
🧠 Review effort: Balanced
| - 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 |
…M) into fix/model-dtypes
There was a problem hiding this comment.
Warning
Copilot couldn't run its full agentic review because it didn't start before the timeout. Make sure your repository has a runner available, or add a copilot-code-review.yml file specifying one with the runs-on attribute. See the docs for more details.
Copilot review overview
9 open findings
Cast supplied channel locations to signal device and dtype Fordtype=torch.bfloat16,residual_dtypebecomestorch.float32, soinput_tensoris float32…copy_requiresself.kernel_normto already be allocated with the exact same shape as… The wording reads like all listed dtypes (incl.float64) are supported on Intel Gaudi/HPU, but… · New This test assumesmodel(x)returns atorch.Tensorwith.dtype, but some models may return a… · New Restore float16 dtype after spectral input promotion Narrow float16 support claim to spectral front ends Correct HPU fallback description for DFT APIs Correct universal float16 compatibility claim
🧠 Review effort: Lite
| ``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 | ||
| ops a float32 (at least) input, on the CPU for HPU tensors since PyTorch has no | ||
| complex bfloat16 and Gaudi no complex dtype; the real result is cast back with | ||
| ``.to(x)``. Used by :class:`braindecode.models.BIOT`, :class:`braindecode.models.BrainBERT`, |
| model = model.to(dtype).eval() | ||
| x = torch.randn(2, len(gkw["chs_info"]), gkw["n_times"], dtype=dtype) | ||
| with torch.no_grad(): | ||
| y = model(x) |
…aindecode#1250, braindecode#1252) into fix/hpu-models
…res (#1249) * FIX models run in float64 and bfloat16; CodeBrain trains after inference_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. * DOC set PR number in whats_new * FIX one spectral_input helper for FFT/STFT/filter inputs: float32 at 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. * FIX BrainBERT window on the spectral input's device; EMG2QwertyNet scripts 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. * FIX BrainOmni/BrainTokenizer SEANet LSTM input permuted as a 4D view (Gaudi lazy compile) * DOC whats_new for #1249 * FIX EMG2QwertyNet/MetaNeuromotorHand roll a contiguous input (Gaudi eager roll bug) * DOC whats_new: drop the #1246 entry duplicated by the merge



model.to(torch.float64),torch.bfloat16ortorch.float16failed in the forward of several models, and CodeBrain could not train after a first forward undertorch.inference_mode().One helper for spectral ops. PyTorch has no complex bfloat16 (CPU pocketfft/MKL and MPS reject bf16/fp16 FFTs) and Gaudi has no complex dtype.
braindecode.functional.spectral_input(x)returnsxin at least float32 (on the CPU for HPU tensors); callers cast real results back with.to(x). It is a plain function so scripted models still compile. Used by BIOT, BrainBERT, Brant, CBraMod/CSBrain, CodeBrain, ContraWR, DIVER1, EEGDINO, EMG2QwertyNet, LUNA, MAPA, MetaNeuromotorHand, SensingDynamics,FilterBankLayer(FBCNet, FBMSNet, FBLightConvNet, IFNet),GeneralizedGaussianFilter(EEGMiner) andhilbert_freq.Other root causes:
torch.fullwithout a dtype →x.dtype..float()before the FFTs never cast back; sliding-window mask always float32 → follows the input. Training afterinference_mode: the lazily initialisedkernel_normbuffer was reassigned (an inference tensor) → updated in place.forwardfollow the input dtype.torch.cdist, which has no bf16/fp16 kernel → distances in at least float32.Float32 is unchanged: outputs and every gradient of all registry models are
torch.equalto master.Tests:
test_integration.py::test_forward_in_dtype(every model × float64/bfloat16/float16; xfail with reason where PyTorch has no CPU kernel: EEGSymavg_pool3din bf16/fp16, and fp16 overflow in EEGMiner, the filter-bank models and LUNA),test_pretrained_compat.py::test_forward_in_dtype(every pretrained model × float64/bfloat16) andtest_codebrain_trains_after_inference_mode.