Skip to content

Models run in float64 and bfloat16; CodeBrain trains after inference_mode - #1246

Merged
bruAristimunha merged 8 commits into
braindecode:masterfrom
bruAristimunha:fix/model-dtypes
Oct 8, 2026
Merged

bruAristimunha merged 8 commits into
braindecode:masterfrom
bruAristimunha:fix/model-dtypes

Conversation

@bruAristimunha

@bruAristimunha bruAristimunha commented Oct 7, 2026 •

Copy link
Copy Markdown
Collaborator

model.to(torch.float64), torch.bfloat16 or torch.float16 failed in the forward of several models, and CodeBrain could not train after a first forward under torch.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) returns x in 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) and hilbert_freq.

Other root causes:

  • BENDR: start token built with torch.full without a dtype → x.dtype.
  • LUNA: NeRF channel-location encoding returned float32 → coordinates' dtype.
  • CodeBrain: .float() before the FFTs never cast back; sliding-window mask always float32 → follows the input. Training after inference_mode: the lazily initialised kernel_norm buffer was reassigned (an inference tensor) → updated in place.
  • REVE: default positions are a plain float32 attribute → Fourier features in float32, cast to the signal's dtype.
  • ZUNA: residual stream forced to float32 and fed to float64/bf16 linears → sub-layers get their weights' dtype.
  • SyncNet, DGCNN: tensors built in forward follow the input dtype.
  • Residual vector quantizer (BrainOmni, BrainTokenizer): k-means init used 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.equal to master.

Tests: test_integration.py::test_forward_in_dtype (every model × float64/bfloat16/float16; xfail with reason where PyTorch has no CPU kernel: EEGSym avg_pool3d in 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) and test_codebrain_trains_after_inference_mode.

…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.
Copilot AI balanced review requested due to automatic review settings October 7, 2026 15:36

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +490 to +491
self.kernel_norm.copy_(kernel.norm(dim=-1, keepdim=True).detach())
self.kernel_norm_initialized.fill_(True)
Comment on lines +540 to +546
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)
Copilot AI balanced review requested due to automatic review settings October 7, 2026 15:58

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🟡 Changes recommended

LUNA still fails in converted dtypes when callers provide float32 channel locations explicitly.

3 open findings

🧠 Review effort: Balanced

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

codecov Bot commented Oct 7, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 97.43590% with 2 lines in your changes missing coverage. Please review.
✅ Project coverage is 88.73%. Comparing base (030d24f) to head (a39c98b).

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:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

…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.
Copilot AI balanced review requested due to automatic review settings October 7, 2026 17:35

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

# not implement bfloat16 FFT kernels.
if input_dtype == torch.bfloat16:
x = x.float()
x = spectral_input(x)
Comment thread docs/whats_new.rst
Comment on lines +212 to +213
- 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
Comment thread docs/whats_new.rst
Comment on lines +224 to +225
(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.
Copilot AI balanced review requested due to automatic review settings October 7, 2026 17:53

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Comment thread docs/whats_new.rst
Comment on lines +216 to +217
- 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
Copilot AI balanced review requested due to automatic review settings October 7, 2026 21:42

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Comment thread docs/whats_new.rst
Comment on lines +231 to +235
``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)
Copilot AI balanced review requested due to automatic review settings October 8, 2026 07:11

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@bruAristimunha
bruAristimunha merged commit 5355344 into braindecode:master Oct 8, 2026
14 checks passed
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Oct 8, 2026
bruAristimunha added a commit to bruAristimunha/braindecode that referenced this pull request Oct 8, 2026
bruAristimunha added a commit that referenced this pull request Oct 8, 2026
…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
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants