Skip to content

Model & filter fidelity fixes (#1067–#1071) + AugmentedDataLoader expansion & augmentation speedups - #1073

Merged
bruAristimunha merged 6 commits into
masterfrom
fix/model-fidelity-and-augmentation-perf
Jun 24, 2026
Merged

bruAristimunha merged 6 commits into
masterfrom
fix/model-fidelity-and-augmentation-perf

Conversation

@bruAristimunha

Copy link
Copy Markdown
Collaborator

Two themes: (1) correctness/fidelity fixes for four models and the IIR filter bank, and (2) the AugmentedDataLoader fixed-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

AugmentedDataLoader n_augmentation + augmentation speedups (#1070)

Feature. AugmentedDataLoader(..., n_augmentation=N) expresses a fixed set-expansion: each batch keeps its clean originals and appends N independently transformed copies ((1+N)x), e.g. the EEG-Inception MI 6x training set (n_augmentation=5). Default 0 keeps the current in-place behavior — fully backwards compatible.

Speedups (all verified equivalent; benchmark in benchmarks/bench_augmentation_growth.py):

change effect
collate closure → picklable _AugmentationCollate class enables num_workers>0 (the closure couldn't pickle); near-linear worker scaling (2 wk → 2.01x, 4 wk → 3.63x)
Transform.forward all-True fast path skips the boolean gather/scatter when every sample is augmented (probability>=1.0); 1.24x, benefits all transforms, numerically identical
mask_encoding vectorized drops the per-batch Python loop → single scatter broadcast over channels; 2.9x, identical
channels_shuffle gather replaces (B,C,C) one-hot matrices + matmul with a gather; 1.8x, bit-for-bit identical and seed-reproducible (the per-sample rng.permutation is kept on purpose so seeded outputs don't change)

Test plan

  • pytest test/unit_tests/augmentation/ — 129 passed (includes new n_augmentation tests).
  • pytest test/unit_tests/models/test_models.py -k eegitnet and the generic test_integration.py harness (torchscript / export / compile / batch-1) for EEGITNet, EEGSym, MSVTNet.
  • pytest test/unit_tests/models/test_modules.py (FilterBankLayer) and FBCNet/FBMSNet forward.
  • All equivalences (mask_encoding, channels_shuffle, the forward fast path) checked against the previous implementation; channels_shuffle is bit-identical.

Closes #1067, closes #1068, closes #1069, closes #1070, closes #1071.

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.
Copilot AI review requested due to automatic review settings June 24, 2026 17:23

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge 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 👍 / 👎.

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.

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 filtfilt recursion 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.forward fast path, channels_shuffle gather, vectorized mask_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.

Comment on lines +254 to +255
if n_augmentation < 0:
raise ValueError("n_augmentation must be a non-negative integer.")
Comment on lines +204 to +219
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)
Comment on lines +189 to +195
# 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)
@bruAristimunha
bruAristimunha merged commit a7ee6eb into master Jun 24, 2026
18 of 19 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

2 participants