Skip to content

[ENH] EMG2QwertyNet: built-in SpecAugment + feature-extraction flags - #1015

Merged
bruAristimunha merged 7 commits into
braindecode:masterfrom
bruAristimunha:emg2qwerty-built-in-spec-augment
May 9, 2026
Merged

bruAristimunha merged 7 commits into
braindecode:masterfrom
bruAristimunha:emg2qwerty-built-in-spec-augment

Conversation

@bruAristimunha

@bruAristimunha bruAristimunha commented May 9, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Move emg2qwerty's SpecAugment from a Lightning callback into EMG2QwertyNet itself, and expose the encoder representation through two BIOT-style feature-extraction flags so downstream wrappers (e.g. neuroai's DownstreamWrapperModel) can drop in without a custom hook.

  • spec_augment (init kwarg, default False) — installs a parameter-free _SpecAugment submodule between the log-spectrogram and BatchNorm. Active only in train(); default builds an nn.Identity, so state-dict layout and existing checkpoints are bit-identical. Knobs match the paper (Sivakumar et al. NeurIPS 2024, Sec 5.2): n_time_masks=3, time_mask_param=25, n_freq_masks=2, freq_mask_param=4, spec_augment_prob=1.0. Per-(sample × band × electrode) IID masking via torchaudio.transforms.{Time,Frequency}Masking(iid_masks=True) on the 5-D (B, num_bands, electrodes, freq, T_spec) tensor — same recipe as the upstream emg2qwerty/transforms.py:SpecAugment dataset transform. mask_value stays a 0-D on-device tensor (no host sync on GPU).
  • return_features=True (runtime kwarg) → {"features": (B, T_out, num_features), "cls_token": None} dict, BIOT / signal-JEPA convention.
  • return_feature=True (init kwarg, BIOT legacy) → default forward returns (emissions, features) tuple, so YAML-driven wrappers can pick up features via model_output_key=1 without runtime kwargs.

Test plan

pytest test/unit_tests/models/test_models.py -k emg2qwerty — 6 passing (4 pre-existing + 2 new):

  • test_emg2qwerty_spec_augment_contract — eval determinism, train mutation + stochasticity, per-(B, band, electrode) IID (NOT band-shared, regression guard), zero state-dict footprint, constructor validation, from_config round-trip.
  • test_emg2qwerty_feature_flags — all three forward paths (default tensor / runtime dict / init tuple), runtime-wins-over-init, get_config round-trip, get_output_shape invariant.

Notes for reviewers

  • All review comments resolved (5 threads). Codex's "band-shared masking" P2 was reverted after re-auditing upstream emg2qwerty/transforms.py — the paper recipe is per-electrode IID; band-shared was a 16× under-augmentation regression.
  • _SpecAugment is _-prefixed and intentionally not re-exported — only the constructor flag is public, mirroring _LogSpectrogram / _SpectrogramNorm.
  • No new dependencies; reuses torchaudio.transforms already imported for the STFT.

Move the upstream SpecAugment masking from a Lightning callback
(neuralbench's SpecAugmentCallback, registered as a forward hook on
``model.spectrogram``) into ``EMG2QwertyNet`` itself as a parameter-free
``_SpecAugment`` submodule that runs only in ``train()`` mode. Constructor
flag ``spec_augment=False`` (default) installs an ``nn.Identity`` so the
state-dict layout and existing checkpoints stay untouched; setting
``spec_augment=True`` exposes the paper-recipe knobs (``n_time_masks``,
``time_mask_param``, ``n_freq_masks``, ``freq_mask_param``,
``spec_augment_prob``) and applies the mask in-place between the
log-spectrogram and the per-channel BatchNorm. Masking remains IID per
``(sample x band)`` and shared across the electrodes within a band, with
mask value set to the per-window mean (matches the upstream callback).

Adds three regression tests: train/eval determinism contrast, identity
default, and constructor argument validation.
Copilot AI review requested due to automatic review settings May 9, 2026 14:40

@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: 037272ab69

ℹ️ 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".

Comment thread braindecode/models/emg2qwerty.py Outdated

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 makes EMG2QwertyNet self-contained by moving SpecAugment from an external (Lightning-side) mechanism into the model as an internal, parameter-free submodule that only applies masking in train() mode.

Changes:

  • Add a spec_augment constructor flag (default False) that installs either _SpecAugment or an nn.Identity passthrough to preserve checkpoint/state-dict layout when disabled.
  • Implement _SpecAugment (train-only time/frequency masking on the log-spectrogram) using torchaudio.transforms masking ops and configurable masking hyperparameters.
  • Add unit tests covering train-only behavior, disabled default behavior, config round-trip, state-dict key stability, and parameter validation.

Reviewed changes

Copilot reviewed 2 out of 2 changed files in this pull request and generated 2 comments.

File Description
braindecode/models/emg2qwerty.py Adds built-in SpecAugment module and wires it into the forward path between spectrogram computation and normalization/encoder.
test/unit_tests/models/test_models.py Adds tests for SpecAugment train-only behavior, default disabled behavior, validation, state-dict layout, and config round-tripping.

💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.

Comment thread braindecode/models/emg2qwerty.py Outdated
Comment thread test/unit_tests/models/test_models.py Outdated
Add ``return_features: bool = False`` to ``EMG2QwertyNet.forward`` so
the encoder can drop into neuroai's
``DownstreamWrapperModel(model_output_key="features")`` (and any other
caller following the BIOT / signal-JEPA convention) without needing a
custom forward hook.

When ``True``, ``forward`` returns
``{"features": (batch, T_out, num_features), "cls_token": None}`` —
the output of the TDS-Conv encoder, batch-first, before
``self.final_layer``. ``cls_token`` is fixed to ``None`` because
TDS-Conv has no ``[CLS]``; the key is kept for parity with the rest of
the family. Default ``False`` keeps the existing
``(batch, T_out, n_outputs)`` emissions tensor untouched, so existing
CTC training loops, tests, and ``get_output_shape`` are unaffected.

Adds a regression test pinning the dict layout, the batch-first
``(B, T_out, num_features)`` shape, and the invariant that applying
``self.final_layer`` to the returned features reproduces the default
emissions.
- Codex P2: ``torchaudio.transforms.TimeMasking(iid_masks=True)`` on a
  ``(B*num_bands, electrodes, freq, T)`` tensor draws one mask per
  ``(batch, electrode)`` pair, not per ``(sample × band)`` — every
  electrode within a band saw an independent mask, contradicting the
  comment and pushing the augmentation strength ~32× past the paper
  recipe. Replace the torchaudio call with manual band-shared masking:
  vectorised per-row width / start sampling, ``masked_fill`` broadcast
  across the electrode axis. Same shape, correct semantics.

- Copilot: ``flat.mean().item()`` forced a host round-trip on every
  training forward when SpecAugment was enabled. Mask value is now a
  0-D on-device tensor (``flat.mean()``); ``masked_fill`` accepts it
  directly, no sync.

- Copilot: the train-mode regression test trusted seeds 1 / 3 to draw
  non-zero ``n_t`` and ``n_f`` from ``torch.randint(n+1, ())``. With
  the manual-masking rewrite the RNG sequence shifts, so the test now
  monkeypatches ``torch.randint`` for the size-`()` calls (forcing two
  time and two frequency masks per call) and additionally pins the
  band-sharing contract on a uniquely-valued tensor: every electrode
  within a ``(B, band)`` row must see the exact same mask pattern.

Adds a what's-new entry covering both the built-in SpecAugment and
the ``return_features`` flag from the previous commit.
Copilot AI review requested due to automatic review settings May 9, 2026 15:00
@bruAristimunha

Copy link
Copy Markdown
Collaborator Author

Thanks for the reviews — pushed b355eb2b addressing all three:

@chatgpt-codex (P2: per-electrode masks) — Confirmed by direct check: torchaudio.transforms.TimeMasking(iid_masks=True) on a (B*num_bands, electrodes, freq, T) 4-D tensor samples one mask per (batch, electrode) pair, so each of the 16 electrodes in a band drew its own independent mask. The comment and the original neuralbench callback were both wrong about the semantics. Replaced the torchaudio path with manual band-shared masking (_band_shared_axis_mask static method): per-row vectorised width / start sampling, masked_fill broadcasts the mask across the electrode axis. Pinned by a new contract assertion in test_emg2qwerty_spec_augment_train_only that builds a uniquely-valued tensor and verifies every electrode within each (B, band) row sees the exact same mask pattern.

@copilot (.item() GPU sync) — Mask value is now flat.mean() (0-D on-device tensor) and masked_fill accepts it directly. No more host round-trip per training forward.

@copilot (seed-dependent test) — The train-mode regression now monkeypatches torch.randint for the size-() calls so both n_t and n_f come back as 2 deterministically, regardless of global RNG / torchaudio versions.

Also added a what's-new entry covering both built-in SpecAugment and the return_features flag from the previous commit.

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

Copilot reviewed 3 out of 3 changed files in this pull request and generated 2 comments.

Comment thread braindecode/models/emg2qwerty.py Outdated
Comment thread docs/whats_new.rst Outdated
…e flag

Re-audited upstream emg2qwerty (Meta, NeurIPS 2024) before locking in
the previous Codex-driven rewrite. Two fixes here.

1. SpecAugment masking semantics — revert band-sharing.

   ``emg2qwerty/transforms.py:SpecAugment.__call__`` applies SpecAugment
   as a *dataset-level* transform on a single sample shaped
   ``(T, bands, electrodes, freq)``. After ``movedim(0, -1)`` the input
   is 4-D ``(bands, electrodes, freq, T)`` and ``torchaudio``'s
   ``iid_masks=True`` samples one mask per leading axis index — i.e.
   per ``(band, electrode)`` pair, with each electrode in a band drawing
   its own mask. The previous neuralbench callback's ``(B*num_bands,
   electrodes, freq, T)`` reshape produced the same per-``(sample × band
   × electrode)`` independence; only its comment was wrong. The Codex
   review's claim that the paper recipe is band-shared and the original
   was "32× too aggressive" doesn't match the upstream code.

   The previous commit's manual band-shared masking (one mask per
   ``(sample × band)``, broadcast over all 16 electrodes) was therefore
   16× *less* aggressive than the paper. Replace it with a direct
   ``torchaudio.transforms.TimeMasking(iid_masks=True)`` /
   ``FrequencyMasking(iid_masks=True)`` call on the 5-D
   ``(B, num_bands, electrodes, freq, T_spec)`` tensor (no
   reshape gymnastics, no manual masking). ``mask_along_axis_iid``
   draws one mask per leading-axis triple, so the result is per-
   ``(B × num_bands × electrodes)`` independence — exactly the upstream
   recipe scaled to a batched tensor. ``mask_value=spec.mean()``
   stays a 0-D on-device tensor; ``torchaudio`` handles ``Tensor``
   ``mask_value`` natively (``torch.where`` path), so there is still no
   host round-trip on GPU. Update the test contract to assert
   per-``(B, band, electrode)`` independence — i.e. that masking is
   NOT band-shared — and pin it deterministically by monkeypatching
   ``torch.randint``.

2. ``return_feature`` (init kwarg, BIOT-style) for config-driven wrappers.

   Confirmed via direct trace that ``DownstreamWrapper.build`` calls
   ``model(**dummy_batch)`` with no runtime kwargs — and
   ``build_dummy_batch`` only seeds the first forward parameter. So the
   runtime ``return_features=True`` (plural) flag added in 3a831aa is
   inert when neuroai's wrapper drives the model from YAML
   (``model(x)`` returns the emissions tensor; ``tensor["features"]``
   raises). Mirror :class:`BIOT`'s legacy contract: a ``return_feature``
   (singular) constructor kwarg makes ``forward`` return a
   ``(emissions, features)`` tuple by default, so a YAML such as
   ``kwargs: {return_feature: true}`` + ``model_output_key: 1`` picks up
   the ``(B, T_out, num_features)`` encoder representation without a
   call-site change. Runtime ``return_features=True`` still wins (returns
   the dict) for programmatic users; both round-trip via
   ``get_config`` / ``from_config`` and ``get_output_shape`` keeps
   reporting the emissions shape regardless of which flag is set.
@bruAristimunha

Copy link
Copy Markdown
Collaborator Author

Self-correction on Codex's earlier P2 — pushed `55b78ec1`.

Re-audited the upstream `emg2qwerty/transforms.py` to confirm the actual paper recipe before locking in the band-shared rewrite. Findings:

```python

emg2qwerty/transforms.py — applied at the dataset level on a SINGLE sample

def call(self, specgram: torch.Tensor) -> torch.Tensor:
# (T, ..., C, freq) -> (..., C, freq, T)
x = specgram.movedim(0, -1)
...
x = self.time_mask(x, mask_value=self.mask_value) # iid_masks=True
```

The single-sample input is `(T, bands, electrodes, freq)`. After `movedim(0, -1)` it's 4-D `(bands, electrodes, freq, T)`, and `torchaudio`'s `iid_masks=True` samples one mask per leading axis pair → per `(band, electrode)`. Each electrode in a band gets its own mask. Verified empirically too. The previous neuralbench callback's `(B*num_bands, electrodes, freq, T)` reshape produced the same per-`(sample × band × electrode)` independence; only its comment was wrong.

So the band-shared "fix" was a regression — it made the augmentation 16× less aggressive than the paper. Reverted to per-`(sample × band × electrode)` iid masking via a direct `torchaudio` call on the 5-D tensor (`mask_along_axis_iid` draws one mask per leading axis triple). Kept the on-device `mask_value` (no host sync) and the deterministic `torch.randint` monkeypatch in the test. Test contract is now an assertion that masking is NOT band-shared.

Also fixed a separate issue surfaced while double-checking the `return_features` integration with neuroai. `DownstreamWrapper.build` calls `model(**dummy_batch)` where `dummy_batch` only contains the first forward parameter — it does not pass `return_features=True` at runtime. So the runtime dict flag was inert under neuroai's YAML-driven flow. Added a `return_feature` (singular, init kwarg, BIOT-style legacy) that makes `forward` return a `(emissions, features)` tuple by default, so a config such as `kwargs: {return_feature: true}` + `model_output_key: 1` works out of the box. Runtime `return_features=True` still wins for programmatic users.

Apologies @chatgpt-codex-connector — your concern that masking was per-`(B × electrode)` was factually correct, but the conclusion that the paper is band-shared isn't supported by the upstream code. The paper recipe is per-electrode independent, and that's what's in this commit.

@chatgpt-codex-connector

Copy link
Copy Markdown

To use Codex here, create an environment for this repo.

``[park2019specaug]_`` is defined only in the EMG2QwertyNet docstring,
not in ``whats_new.rst`` — Sphinx reports an undefined-citation error
when building the changelog. Replace with an inline parenthetical
("Park et al., Interspeech 2019") so the entry stands on its own.
Copilot AI review requested due to automatic review settings May 9, 2026 15:26

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

Copilot reviewed 3 out of 3 changed files in this pull request and generated 3 comments.

Comment thread braindecode/models/emg2qwerty.py Outdated
Comment thread braindecode/models/emg2qwerty.py
Comment thread docs/whats_new.rst
5 tests (211 lines) → 2 tests (87 lines) with the same contracts.
``test_emg2qwerty_spec_augment_contract`` folds the disabled-default
``nn.Identity`` check, the constructor validation, the train/eval
mutation contrast, the per-(B, band, electrode) independence assertion,
the parameter-free state-dict invariant, and the ``from_config`` round-trip
into one. The mock-patched mask-count helper is hoisted out of the
function body and the per-electrode independence loop is replaced by a
vectorised flatten-and-compare. ``test_emg2qwerty_feature_flags`` covers
all three forward paths (default tensor, runtime dict, init tuple) plus
runtime-wins-over-init, ``get_config`` round-trip, and
``get_output_shape`` invariants in one pass over a shared input.

No behaviour change in the model. All 6 emg2qwerty tests still pass.
@bruAristimunha bruAristimunha changed the title [ENH] EMG2QwertyNet: built-in SpecAugment on the log-spectrogram [ENH] EMG2QwertyNet: built-in SpecAugment + feature-extraction flags May 9, 2026
- Lazy ``features`` materialisation. ``forward`` was always calling
  ``encoded.transpose(0, 1).contiguous()`` even on the default
  emissions-only path, paying an extra transpose + copy per training
  step regardless of whether the caller set either feature flag. Move
  the transpose into the two branches that actually return the
  features tensor; the default path now stays identical to pre-flag
  behaviour.

- Device-aware RNG in ``_SpecAugment``. The probabilistic gate
  (``torch.rand((), device="cpu").item() >= self.prob``) and the
  mask-count draws (``torch.randint(n+1, ())``) defaulted to CPU,
  which mixed CPU RNG with the device-side RNG that ``torchaudio``'s
  ``mask_along_axis_iid`` uses internally — ``torch.cuda.manual_seed``
  alone would make the gate non-deterministic. Route both through
  ``x.device`` so a single device-appropriate seed reproduces the
  whole augmentation. The ``.item()`` calls are still required for
  the Python ``if``/``range`` bounds (control-flow sync is
  unavoidable).
Copilot AI review requested due to automatic review settings May 9, 2026 15:44
@bruAristimunha

Copy link
Copy Markdown
Collaborator Author

Pushed `d02ac8af` addressing all three @copilot-pull-request-reviewer comments:

  1. Lazy `features` materialisation — `forward()` no longer transposes `encoded` on the default emissions-only path; the transpose + contiguous copy now happens only inside the two branches that actually return the features tensor. Default path is back to pre-flag behaviour.
  2. Device-aware RNG — both the prob gate (`torch.rand((), device=x.device)`) and the mask-count draws (`torch.randint(n+1, (), device=x.device)`) now use `x.device`. Single seed (`torch.manual_seed` or `torch.cuda.manual_seed`) reproduces the augmentation, and torchaudio's internal device-side RNG stays in the same stream as our Python-level gate. `.item()` is still required to bind the Python `if`/`range`, so a host sync remains for control flow only.
  3. Changelog ↔ PR description alignment — the PR description was rewritten in a previous step (the old "shared across electrodes within a band" wording was a stale draft from the band-shared revert). Both the description and `whats_new.rst:45` now say per-`(sample × band × electrode)` IID.

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

Copilot reviewed 3 out of 3 changed files in this pull request and generated 2 comments.

Comment on lines +673 to +680
spec = x.movedim(0, -1).contiguous()
# 0-D on-device tensor — ``masked_fill`` / ``torch.where`` accept it
# without a host sync.
mask_value = spec.mean()
n_t = int(torch.randint(self.n_time_masks + 1, (), device=x.device).item())
for _ in range(n_t):
spec = self.time_mask(spec, mask_value=mask_value)
n_f = int(torch.randint(self.n_freq_masks + 1, (), device=x.device).item())
Comment on lines +3541 to +3543
size = kwargs.get("size", args[1] if len(args) >= 2 else None)
if size == ():
return torch.tensor(2, dtype=torch.long)
@bruAristimunha
bruAristimunha merged commit a792446 into braindecode:master May 9, 2026
17 checks passed
@bruAristimunha
bruAristimunha deleted the emg2qwerty-built-in-spec-augment branch May 10, 2026 10:27
@bruAristimunha bruAristimunha mentioned this pull request May 19, 2026
6 tasks
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