Repository navigation
[ENH] EMG2QwertyNet: built-in SpecAugment + feature-extraction flags - #1015
bruAristimunha merged 7 commits into
Conversation
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.
There was a problem hiding this comment.
💡 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".
There was a problem hiding this comment.
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_augmentconstructor flag (defaultFalse) that installs either_SpecAugmentor annn.Identitypassthrough to preserve checkpoint/state-dict layout when disabled. - Implement
_SpecAugment(train-only time/frequency masking on the log-spectrogram) usingtorchaudio.transformsmasking 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.
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.
|
Thanks for the reviews — pushed @chatgpt-codex (P2: per-electrode masks) — Confirmed by direct check: @copilot ( @copilot (seed-dependent test) — The train-mode regression now monkeypatches Also added a what's-new entry covering both built-in SpecAugment and the |
…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.
|
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 sampledef call(self, specgram: torch.Tensor) -> torch.Tensor: 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. |
|
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.
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.
- 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).
|
Pushed `d02ac8af` addressing all three @copilot-pull-request-reviewer comments:
|
| 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()) |
| size = kwargs.get("size", args[1] if len(args) >= 2 else None) | ||
| if size == (): | ||
| return torch.tensor(2, dtype=torch.long) |
Summary
Move emg2qwerty's SpecAugment from a Lightning callback into
EMG2QwertyNetitself, and expose the encoder representation through two BIOT-style feature-extraction flags so downstream wrappers (e.g. neuroai'sDownstreamWrapperModel) can drop in without a custom hook.spec_augment(init kwarg, defaultFalse) — installs a parameter-free_SpecAugmentsubmodule between the log-spectrogram and BatchNorm. Active only intrain(); default builds annn.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 viatorchaudio.transforms.{Time,Frequency}Masking(iid_masks=True)on the 5-D(B, num_bands, electrodes, freq, T_spec)tensor — same recipe as the upstreamemg2qwerty/transforms.py:SpecAugmentdataset transform.mask_valuestays 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 viamodel_output_key=1without 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_configround-trip.test_emg2qwerty_feature_flags— all three forward paths (default tensor / runtime dict / init tuple), runtime-wins-over-init,get_configround-trip,get_output_shapeinvariant.Notes for reviewers
emg2qwerty/transforms.py— the paper recipe is per-electrode IID; band-shared was a 16× under-augmentation regression._SpecAugmentis_-prefixed and intentionally not re-exported — only the constructor flag is public, mirroring_LogSpectrogram/_SpectrogramNorm.torchaudio.transformsalready imported for the STFT.