Repository navigation
Model & filter fidelity fixes (#1067–#1071) + AugmentedDataLoader expansion & augmentation speedups - #1073
Conversation
The IIR path downcast float64 coefficients to the input dtype before filtfilt, so band-passes with poles near the unit circle (cheby2 banks used by FBMSNet/FBCNet) diverged to NaN in float32. Run the recursion in float64 and cast back. Closes #1067
Match the authors' EEG-ITNet: max-norm(1.0) on the 3 inception spatial depthwise convs, max-norm(0.25) on the final dense, and dimensionality-reduction width 14 (was doubled to 28). Closes #1068
forward dropped the main head and returned only branch predictions, making the paper's joint deep-supervision loss unreproducible. Return (main_logits, stacked_branch_logits), matching the source. Add MSVTNet to the torchscript skip-list (polymorphic return, as for EEGDINO). Closes #1069
The inception/residual spatial convs were dense (groups=1, bias=True); the authors' unit_dconv uses grouped/depthwise convs with no bias. Closes #1071
- AugmentedDataLoader: add n_augmentation for fixed set-expansion (keep the clean originals + N augmented copies); default 0 keeps current behavior. Closes #1070 - Replace the collate closure with a picklable _AugmentationCollate class, enabling num_workers>0. - Transform.forward: fast path that skips the boolean gather/scatter when every sample is augmented (probability>=1.0). - Vectorize mask_encoding (drop the per-batch Python loop) and channels_shuffle (gather instead of one-hot matrices + matmul; output bit-identical and seed-reproducible). - Add benchmarks/bench_augmentation_growth.py.
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: 5e9972905d
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
| aug_X, aug_y = self.transform(X, y) | ||
| xs.append(aug_X) | ||
| ys.append(aug_y) | ||
| return torch.cat(xs), torch.cat(ys) |
There was a problem hiding this comment.
Handle tuple targets during batch expansion
When n_augmentation > 0 is used with a transform that returns tuple targets, such as the existing Mixup transform ((y_a, y_b, lam)), ys contains the original tensor label batch followed by tuple labels from augmented copies. The final torch.cat(ys) therefore raises a TypeError before yielding a batch, even though Transform.forward otherwise supports tuple targets; the expansion path needs to concatenate tuple components consistently or reject these transforms explicitly.
Useful? React with 👍 / 👎.
There was a problem hiding this comment.
Pull request overview
This PR improves model/filter fidelity for several architectures (IIR filter stability, MSVTNet feature returns, EEGITNet/EEGSym architectural fidelity) and expands the augmentation subsystem with a fixed-expansion n_augmentation mode plus several augmentation performance optimizations (picklable collate, fast paths, vectorization).
Changes:
- Fix FilterBankLayer IIR instability by running
filtfiltrecursion in float64 and casting back. - Bring multiple models closer to their reference implementations (EEGITNet constraints + DR width, MSVTNet returning main+branch logits, EEGSym grouped spatial convs).
- Add
AugmentedDataLoader(n_augmentation=...)and speed up augmentation ops (Transform.forwardfast path,channels_shufflegather, vectorizedmask_encoding), plus a benchmark script.
Reviewed changes
Copilot reviewed 10 out of 10 changed files in this pull request and generated 3 comments.
Show a summary per file
| File | Description |
|---|---|
braindecode/modules/filter.py |
Run IIR filtering in float64 to prevent float32 NaNs, cast back to input dtype. |
braindecode/models/eegitnet.py |
Add max-norm constraints + correct DR width; update weight mapping for parametrized final layer. |
braindecode/models/msvtnet.py |
Return (main_logits, branch_logits) when return_features=True. |
braindecode/models/eegsym.py |
Match reference by using grouped/depthwise spatial convs without bias. |
braindecode/augmentation/base.py |
Add picklable collate and n_augmentation fixed-expansion mode; add Transform fast path. |
braindecode/augmentation/functional.py |
Speedups: gather-based channels_shuffle, vectorized mask_encoding, drop unused import. |
test/unit_tests/augmentation/test_base.py |
Add tests for n_augmentation behavior and validation. |
test/unit_tests/models/test_integration.py |
Skip TorchScript for MSVTNet due to polymorphic return type. |
benchmarks/bench_augmentation_growth.py |
New standalone benchmark for augmentation scaling and worker scaling. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| if n_augmentation < 0: | ||
| raise ValueError("n_augmentation must be a non-negative integer.") |
| def __call__(self, batch): | ||
| collated = default_collate(batch) | ||
| X, y = collated[:2] | ||
| if self.device is not None: | ||
| X = X.to(self.device) | ||
| y = y.to(self.device) | ||
| if self.n_augmentation == 0: | ||
| return (*self.transform(X, y), *collated[2:]) | ||
| # Fixed expansion: keep the clean originals, then append | ||
| # ``n_augmentation`` independently transformed copies. | ||
| xs, ys = [X], [y] | ||
| for _ in range(self.n_augmentation): | ||
| aug_X, aug_y = self.transform(X, y) | ||
| xs.append(aug_X) | ||
| ys.append(aug_y) | ||
| return torch.cat(xs), torch.cat(ys) |
| # Return BOTH the main-head logits and the stacked branch logits, as | ||
| # in the authors' implementation (``return x, bx``). This is required | ||
| # for the paper's joint deep-supervision loss | ||
| # ``lambda * CE(main) + (1 - lambda) * sum_i CE(branch_i)`` and lets | ||
| # callers score on the main head. | ||
| # x: [batch_size, n_classes]; branches: [n_branches, batch_size, n_classes] | ||
| return x, torch.stack(branch_preds) |
Two themes: (1) correctness/fidelity fixes for four models and the IIR filter bank, and (2) the
AugmentedDataLoaderfixed-expansion feature plus a set of augmentation speedups. Each commit is scoped to one issue (the augmentation feature and the perf work share the augmentation subsystem and are in one commit). Happy to split into per-issue PRs if preferred.Model & filter fidelity fixes
FilterBankLayerIIR NaNs in float32 (FilterBankLayer IIR path NaNs in float32 (coefficients downcast before filtfilt) #1067)._apply_iirdowncast the float64 coefficients to the input dtype beforefiltfilt, so band-passes with poles near the unit circle (cheby2 banks used byFBMSNet/FBCNet) diverged to NaN in float32. Run the recursion in float64 and cast back.EEGITNetmissing constraints + doubled DR width (EEGITNet: missing source max-norm constraints + doubled dimensionality-reduction width #1068). Apply the authors'max_norm(1.0)to the 3 inception spatial depthwise convs andmax_norm(0.25)to the final dense (via the existingLinearWithConstraint/MaxNormParametrizeprimitives), and set the dimensionality-reduction conv width to 14 (it was doubled to 28). The weight-loadingmappingnow points atfinal_layer.parametrizations.weight.original(same convention ascodebrain/atcnet).MSVTNetdrops main head withreturn_features=True(MSVTNet.forward returns branch predictions only (drops main head) when return_features=True #1069).forwardreturned only branch predictions, making the paper's joint deep-supervision loss unreproducible. It now returns(main_logits, stacked_branch_logits), matching the source. MSVTNet is added to the torchscript skip-list (polymorphic return, exactly asEEGDINO); export/compile/default-forward are unaffected (they trace the actual path).EEGSymdense spatial convs (EEGSym spatial convolutions are dense (groups=1, bias=True) instead of depthwise/grouped no-bias #1071). The inception/residual spatial convs were dense (groups=1, bias=True); the authors'unit_dconvuses grouped/depthwise convs with no bias. Nowgroups=out_channels, bias=False.AugmentedDataLoader
n_augmentation+ augmentation speedups (#1070)Feature.
AugmentedDataLoader(..., n_augmentation=N)expresses a fixed set-expansion: each batch keeps its clean originals and appendsNindependently transformed copies ((1+N)x), e.g. the EEG-Inception MI 6x training set (n_augmentation=5). Default0keeps the current in-place behavior — fully backwards compatible.Speedups (all verified equivalent; benchmark in
benchmarks/bench_augmentation_growth.py):_AugmentationCollateclassnum_workers>0(the closure couldn't pickle); near-linear worker scaling (2 wk → 2.01x, 4 wk → 3.63x)Transform.forwardall-True fast pathprobability>=1.0); 1.24x, benefits all transforms, numerically identicalmask_encodingvectorizedchannels_shufflegather(B,C,C)one-hot matrices + matmul with agather; 1.8x, bit-for-bit identical and seed-reproducible (the per-samplerng.permutationis kept on purpose so seeded outputs don't change)Test plan
pytest test/unit_tests/augmentation/— 129 passed (includes newn_augmentationtests).pytest test/unit_tests/models/test_models.py -k eegitnetand the generictest_integration.pyharness (torchscript / export / compile / batch-1) for EEGITNet, EEGSym, MSVTNet.pytest test/unit_tests/models/test_modules.py(FilterBankLayer) and FBCNet/FBMSNet forward.forwardfast path) checked against the previous implementation;channels_shuffleis bit-identical.Closes #1067, closes #1068, closes #1069, closes #1070, closes #1071.