Repository navigation
augmentation: add BandRotation for surface-EMG wristband layouts - #1013
bruAristimunha merged 8 commits into
Conversation
Per-band circular roll along the channel axis plus inter-band temporal jitter, for inputs laid out as ``(B, num_bands * electrodes_per_band, T)``. Models small wristband rotation between sessions and relative timing noise between two arms; introduced in the emg2qwerty paper (Sivakumar et al., NeurIPS 2024). * New ``braindecode.augmentation.functional.band_rotation`` (low-level function on (X, y) tensors). * New ``braindecode.augmentation.BandRotation`` Transform wrapping it for use with ``AugmentedDataLoader``. * 4 functional unit tests (shape, seed reproducibility, no-op path, channel-count validation, circular-roll value preservation) plus a ``test_set_params`` parametrize entry to verify Skorch round-trip. * Bench-checked: a vectorized ``torch.gather`` collapsing the per-band loop is ~16% slower for the typical num_bands=2 case on CPU (gather index tensor exceeds what two contiguous rolls touch); the loop only loses past num_bands>=8. Loop kept with a comment recording the benchmark.
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: fe7b33f624
ℹ️ 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".
* Tighten the channel-count error message (drop redundant middle-term
expansion).
* Compress the gather-benchmark comment from 7 lines to 3.
* Add a Transform-level seed-reproducibility test; the existing tests
only covered the functional.
* Trim a few word-of-filler patterns from docstrings + changelog
("Originally introduced" → "Introduced", "additionally gets" →
"also gets", "the underlying" → drop).
There was a problem hiding this comment.
Pull request overview
This PR introduces a new augmentation for surface-EMG wristband-style channel layouts by adding BandRotation (Transform) and band_rotation (functional) to model per-band circular electrode misalignment plus inter-band temporal jitter, and integrates it into public exports, tests, and the changelog.
Changes:
- Add
band_rotationfunctional augmentation andBandRotationtransform wrapper. - Add unit tests for
band_rotationand includeBandRotationin transform parameter round-trip testing. - Document the new augmentation in
docs/whats_new.rstand export it frombraindecode.augmentation.
Reviewed changes
Copilot reviewed 6 out of 6 changed files in this pull request and generated 4 comments.
Show a summary per file
| File | Description |
|---|---|
braindecode/augmentation/functional.py |
Adds band_rotation implementation (per-band channel roll + optional temporal jitter). |
braindecode/augmentation/transforms.py |
Adds BandRotation transform wrapper and wiring to functional op. |
braindecode/augmentation/__init__.py |
Exports BandRotation from the augmentation package. |
test/unit_tests/augmentation/test_functional.py |
Adds focused unit tests for band_rotation. |
test/unit_tests/augmentation/test_transforms.py |
Adds BandRotation to set_params round-trip coverage. |
docs/whats_new.rst |
Adds changelog entry under Enhancements for the new augmentation. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
* Validate ``band_offsets`` is non-empty and ``max_temporal_jitter >= 0`` with explicit ValueError messages, instead of letting numpy raise its less actionable internal errors. * Clarify that the temporal jitter applies to band 1 only regardless of ``num_bands`` (functional + Transform docstrings). * Note in the Transform docstring that one set of parameters applies to the whole transformed sub-batch — a property users typically expect to read at the Transform level, not the functional one. * Add unit tests for the two new error paths plus a comment on the value-set comparison test cautioning against scaling it up.
Addresses 5 inline comments from chatgpt-codex-connector and copilot-pull-request-reviewer: * (codex) Wrap-around discontinuity from ``torch.roll`` on the time axis: add ``circular_jitter: bool = True`` parameter so users can opt into a crop-and-pad shift instead. Default stays paper-faithful. * (copilot) Fix ``random_state`` docstring type drift in the Transform: ``numpy.random.Generator`` → ``numpy.random.RandomState`` to match what ``check_random_state`` actually returns. * (copilot) Validate ``BandRotation`` constructor parameters up-front (positive ``num_bands``/``electrodes_per_band``, non-empty ``band_offsets``, non-negative ``max_temporal_jitter``) so config mistakes raise on Transform build, not on the first batch. * (copilot) Tighten functional-level validation to match: positive ``num_bands``/``electrodes_per_band`` plus a check that ``band_offsets`` contains integers. * (copilot) Make the seed-reproducibility test deterministic by forcing ``band_offsets=(2,)`` instead of relying on the RNG happening to sample a non-zero offset. New test coverage: parametrized invalid-param rejection + a circular-vs-zero-pad jitter contrast that pins down the boundary where the two modes diverge.
Co-authored-by: Copilot Autofix powered by AI <[email protected]>
Co-authored-by: Copilot Autofix powered by AI <[email protected]>
* Validate that ``band_offsets`` contains integers in ``BandRotation.__init__`` so misconfigurations fail at Transform construction time instead of on the first batch (matches the functional ``band_rotation`` check). * Make ``test_band_rotation_circular_jitter_wraps_vs_zero_pads`` conditional on the observed ``n_diff``: when the RNG happens to draw ``shift == 0`` both modes are no-ops, so we assert outputs equal ``X`` and return; otherwise we keep the boundary-region invariants. This decouples the test from the sampler's exact draw sequence. * Replace the misleading "use a ``set``-based equality check" comment in ``test_band_rotation_circular_roll_preserves_values`` with a pointer to a multiset-aware alternative (``torch.unique(..., return_counts=True)`` / histograms) — sets silently drop duplicates and would miss value-multiplicity bugs.
Normalize ``band_offsets`` to a tuple before the non-empty truth test in ``band_rotation``. Previously ``if not band_offsets:`` would raise ``ValueError: The truth value of an array with more than one element is ambiguous`` if a caller passed a ``numpy.ndarray``, masking the intended labelled error message. ``np.asarray(band_offsets)`` later in the function is unaffected by the early ``tuple(...)`` conversion.
Summary
braindecode.augmentation.BandRotation(and the underlyingband_rotationfunction) — per-band circular roll along the channel axis plus inter-band temporal jitter, for surface-EMG inputs laid out as(B, num_bands * electrodes_per_band, T).Transform/AugmentedDataLoaderpipeline; alphabetised in the__init__exports and changelog updated.Why this lives in braindecode
The augmentation is model-agnostic at the tensor level: any model that consumes a
(B, num_bands * electrodes_per_band, T)EMG layout can use it. Keeping it inbraindecode.augmentationlets non-Lightning users (Skorch / plain DataLoader) pick it up without dragging in a callback framework.API
Performance note
A vectorized
torch.gathercollapsing the per-band loop into a single kernel was benchmarked and is ~16% slower for the typicalnum_bands=2case on CPU (the gather index tensor exceeds what two contiguoustorch.rollcalls touch); only pastnum_bands >= 8does the gather win. The loop is kept with a comment recording the benchmark so a future contributor doesn't redo the experiment.Test plan
test/unit_tests/augmentation/test_functional.py(shape + seed reproducibility, no-op path, channel-count validation, circular-roll value preservation).BandRotationadded to thetest_set_paramsparametrize list intest_transforms.pyto exercise the Skorchset_paramsround-trip.116 passed(was 111 before; +5 from this PR).Changelog
Entry added under
Enhancementsindocs/whats_new.rst.