Skip to content

Include sEMG hand-pose models in the stable release line - #1150

Merged
bruAristimunha merged 1 commit into
braindecode:masterfrom
bruAristimunha:fix/include-emg-pose-models-1.8.1
Aug 31, 2026
Merged

bruAristimunha merged 1 commit into
braindecode:masterfrom
bruAristimunha:fix/include-emg-pose-models-1.8.1

Conversation

@bruAristimunha

@bruAristimunha bruAristimunha commented Aug 31, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

#1132 was merged into refactor/emg-pose-shared-components after #1145 had already merged that base branch into master. Consequently the model commit was absent from the v1.8.0 tag and PyPI wheel.

Validation

  • pytest test/unit_tests/models/test_emg_pose_baselines.py -q (21 passed)
  • pre-commit on all changed source/test files (all passed)
  • downstream NeuralBench construction and forward pass for all three models

* feat(models): add sEMG hand-pose baselines

* docs: note sEMG hand-pose models

* style(models): sort model imports

* fix(models): address hand-pose review findings

* fix(models): preserve causal time alignment

* fix(models): tighten pose model validation

* fix(models): preserve VEMG Hub license

* fix(models): align SensingDynamics with published architecture

* style(models): format SensingDynamics

* fix(models): address pose review edge cases

* feat(models): add pose checkpoint mappings

* feat(neuropose): make the emg2pose capacity reachable, and document the gap

Measured against the released regression_neuropose.ckpt, this port carries
1,440,756 parameters where the checkpoint carries 6,363,758 -- 0.23x. The
mapping table already noted that later blocks diverge; this quantifies it and
makes the reference capacity constructible.

Three sources of the difference: encoder widths 32/64/128 (the base_channels
doubling rule cannot express emg2pose's 32/128/256), 3 residual blocks against
5, and 2 convolutions per residual block against 3.

Adds encoder_channels and n_convs_per_block. Both default to the existing
behaviour, so every published NeuroPoseNet number is unaffected -- a test pins
the default parameter count so it cannot drift silently.

Does not change the defaults: the class documents itself as Liu et al.'s
original rather than emg2pose's adaptation, so raising them would silently
reinterpret existing results. The docstring now states plainly that the
defaults will not reproduce emg2pose's angular error, and gives the argument
set that reaches 0.83x capacity (5,254,516 parameters).

Full equivalence remains out of reach: decoder geometry, padding and
resampling still differ, and the pooling schedule is fixed, so emg2pose's
widened 2 kHz front end cannot be expressed -- this class decimates to
internal_sfreq and discards EMG content above 100 Hz.

Evidence: training the 0.83x configuration on NEMAR NM000281 reaches 15.80 deg
validation angular error against 16.01 deg at the default capacity, same data
and schedule.

* refactor(models)!: name the pose models after their papers

braindecode names a model after its source: the 24 classes ending in Net are
those whose paper name already ends in Net (EEGNet, ATCNet, Deep4Net, FBCNet
...), while BENDR, BIOT, LUNA, USleep and TSception keep theirs unchanged.

None of these three papers uses a Net suffix:

  NeuroPoseNet       -> NeuroPose        Liu, Zhang & Gowda 2021, named in title
  VEMG2PoseNet       -> VEMG2Pose        Salter et al. 2024 call it vemg2pose
  SensingDynamicsNet -> SensingDynamics  emg2pose's name for the Simpetru
                                         et al. 2022 baseline

Safe to do now: all three are absent from origin/master and exist only on this
branch, so no deprecation shim is needed. After merge it would need one.

EMG2QwertyNet keeps its name -- it is already released, and renaming it is a
separate breaking change. The EMG family is inconsistent either way, since
MetaNeuromotorHand also ships without the suffix.

Also records in the NeuroPose docstring that the higher-capacity configuration
was measured and is worse on held-out data at a reduced budget, so the reader
does not repeat the experiment.

* build: require torch>=2.1 for nn.CircularPad2d

SensingDynamics now uses torch.nn.CircularPad2d (added in torch 2.1)
instead of a private F.pad wrapper, so raise the declared floor to match.

2.1 is the minimum the code actually needs. It deliberately stops short
of 2.4: PyTorch dropped x86_64 macOS wheels after 2.2.2, so a higher
floor would make braindecode uninstallable there. uv.lock needs no
regeneration -- its 2.2.2 resolution already satisfies 2.1.

* refactor(models): reuse shared blocks and make the pose models explicit

Replace three hand-rolled helpers in the sEMG hand-pose models with the
existing braindecode/torch equivalents, and make naming, einops patterns
and constructor calls explicit throughout.

Reuse:
- _CausalTDSConvBlock's F.pad + Conv1d -> braindecode.modules.CausalConv1d
- _TDSStage.ff_blocks' nn.Sequential  -> braindecode.modules.MLP
- SensingDynamics._CircularPad2d      -> torch.nn.CircularPad2d

CausalConv1d subclasses nn.Conv1d, and MLP places its Linears at the same
indices as the sequence it replaces, so no state_dict key moves. That
second equivalence depends on MLP's internal layer trimming, so a new
test pins the resulting feedforward geometry.

Explicitness:
- einops axes are named, e.g.
  Rearrange("b c t -> b t c") -> "batch nchans ntimes -> batch ntimes nchans"
- nn.* constructors take keyword arguments (in_features=, out_channels=, ...)
- n_out/chs/c1/c2/c3/in_c/k_t, and NeuroPose.forward's t/h/z/out, are named
- NeuroPose._ResBlock.f -> .layers; VEMG2Pose.p_init -> .initial_pose

VEMG2Pose leaves the torch.export and torch.jit.script skip-lists. Only
the TorchScript entry described a real problem: script rejects rebinding
the LSTM's NoneType hidden state to a tuple, fixed here by threading an
explicit zero state (identical numerics to the None default). The export
entry was stale -- the previous code already exported cleanly.

Defects found in review and fixed here:
- SensingDynamics' feature axis was mislabelled nchans; it carries 8
  electrodes, not the 16 input channels
- _ResBlock was documented as pre-activation, but the activation follows
  the residual sum
- the NeuroPose docstring reported (2, 2) pooling where the code uses (2, 1)
- a dead activation_cls assignment is removed

SensingDynamics.forward keeps its literal receptive-field bound rather
than reading the class attribute: the attribute is absent from the
instance __dict__ that the TorchScript test's plain-module conversion
copies.

All three models remain bit-identical to their previous implementations,
and variable-length inputs (T != n_times) stay supported.

* refactor(models): share the reset_head bookkeeping and simplify the pose models

Five models repeated the same block to keep n_outputs, the braindecode
init kwargs and the Hugging Face hub config in sync after swapping a
head. Add EEGModuleMixin._set_n_outputs to do it once, and call it from
VEMG2Pose, NeuroPose, SensingDynamics, DANCE and EMG2Qwerty. Each
reset_head now shows only the layers it actually rebuilds.

SensingDynamics checks its 167-sample receptive field in __init__ rather
than on every forward pass: the bound is a property of the model, not of
a batch. This makes n_times required at construction, which the model
already declared as mandatory in models_mandatory_parameters, and drops
the literal that forward had to carry because TorchScript cannot resolve
a class attribute there.

NeuroPose.forward loses its intermediate re-assignments. The resampling
block stays inline, with a comment saying why: forward is wrapped by
_disable_batch_norm_training_if_batch_size_one, so TorchScript resolves
global names against util.py, and the plain-module conversion it uses
drops class methods -- neither a helper method nor a module-level
function survives scripting.

Drop test_emg_pose_baselines.py.

All three pose models stay bit-identical, with unchanged state_dict keys.

* feat(neuropose)!: match emg2pose's NeuroPose exactly

The port previously reproduced Liu et al.'s original 200 Hz architecture
and documented a capacity gap against the configuration emg2pose actually
published: 32/64/128 encoder widths against 32/128/256, three residual
blocks of two convs against five of three, and a decimating 200 Hz front
end against upstream's widened pooling at 2 kHz. No argument combination
closed that gap, so the released checkpoint could not be loaded and the
class-level mapping covered only the first encoder block.

Rebuild the model on emg2pose's network/neuropose.yaml instead:

- encoder 32/128/256, pooling (10,2) (8,2) (4,4) over (time, electrodes)
- five residual blocks of three Conv-BN-ReLU-Dropout groups, with no
  activation after the residual sum, matching ResidualBlock upstream
- decoder 128/32/1, upsampling (10,4) (8,4) (4,2)
- padding="same" 3x2 kernels throughout, dropout 0.05
- the head reads the flattened feature/electrode axes, not a fixed width

Verified against the installed emg2pose package: identical parameter
count (6,354,903 at n_chans=16, n_outputs=20), and loading upstream's
state_dict through the mapping reproduces its output exactly
(max|difference| = 0). The mapping is now generated from the block layout
and covers all 149 checkpoint tensors, none missing and none spurious.

One deviation remains, and it is a superset: upstream requires n_times to
be divisible by the total temporal pooling factor, having no way to
recover what flooring discards. This class interpolates back to n_times,
which is the identity when the length does divide.

BREAKING CHANGE: internal_sfreq, n_bands, channel_adapter, base_channels
and encoder_dim are gone; encoder_channels now takes the pooling and
upsampling schedules alongside it. Liu et al.'s original configuration is
still reachable by passing that schedule explicitly, as the docstring
shows. Default capacity rises from ~1.44 M to ~6.35 M parameters, so the
earlier from-scratch replication numbers no longer describe this model.

* refactor(models): derive what was hardcoded, and add the papers' figures

Several values were literals that duplicated something the model already
knew, so they went stale the moment the geometry was reconfigured.

NeuroPose
  The checkpoint map was rebuilt by _build_checkpoint_mapping() from a
  second copy of the block counts. Reconfiguring the schedule left it
  emitting the default layout's 149 keys regardless -- silently wrong. It
  is now read off the blocks themselves, so n_res_blocks=2 with
  n_convs_per_block=4 produces the correct 100 keys and loads.

VEMG2Pose
  encoder_sfreq was sfreq / (5 * 2 * 4 * 2), the stem and TDS strides
  restated as a literal. It is now the product of the strides actually
  configured. The stem kernels/strides, the TDS subsampling schedule and
  the LSTM depth are parameters rather than constants buried in the body.

SensingDynamics
  electrode_features = 8 was the width of the electrode axis after conv3,
  which is a function of n_chans, the circular wrap and the kernels; the
  literal is why n_chans had to be exactly 16. receptive_field_samples was
  a class constant kept in sync with the conv stack by a comment. Both are
  derived now, and the conv geometry, the Butterworth order and the
  electrode wrap are parameters.

All three stay bit-identical at their defaults, with unchanged state_dict
keys, and the new parameters are documented.

Figures: each model now shows its own paper's diagram -- NeuroPose Fig. 6
(encoder/ResNet/decoder, CC BY) and SensingDynamics Fig. 7 (architecture
with the circular padding, CC BY-ND) self-hosted like braindecode's other
model figures, and emg2pose's Fig. 1 linked from arXiv, since that paper
describes vemg2pose only in prose and publishes no architecture diagram.

Pre-trained weights: Meta's regression_neuropose.ckpt is rehosted at
braindecode/NeuroPose-emg2pose, so NeuroPose.from_pretrained() works
without a conversion step. Verified against the real released checkpoint:
strict load, and output identical to the reference implementation.

* fix(vemg2pose)!: make the port load the released emg2pose checkpoints

VEMG2Pose could not load either released checkpoint: its mapping covered
20 of 68 tensors, because the port diverged from upstream's TdsNetwork in
several ways at once. Its second TDS stage ran at feature_dim after an
early 1x1 projection, where upstream keeps 256 and projects at the end of
the stage; its residual blocks were depthwise Conv1d with a per-timestep
LayerNorm, where upstream factorises the channel axis and convolves with
Conv2d; its convolutions were causally left-padded, where upstream is
valid and consumes a left context; and its head emitted one value per
joint, where upstream emits two.

Rebuild it on emg2pose's own network/tds.yaml and pose_modules:

- Conv1dBlock stem (11/5 then 5/2, valid, LayerNorm over channels)
- two TdsStages, each a strided conv then tds_blocks pairs of
  TDSConv2dBlock and TDSFullyConnectedBlock, the last projecting to
  feature_dim
- LSTM over [feature ; previous pose], LeakyReLU, and a head scaled by
  output_scalar

The head width and the rollout now follow `parameterization`. The name
means what it says: 'hybrid' emits a position and a velocity per joint,
taking the position for num_position_steps and integrating the velocity
after, which is what regression_vemg2pose.ckpt was trained as.
'velocity' integrates a single velocity from y0, which is
tracking_vemg2pose.ckpt. 'position' reads the pose out directly.

Verified against the reference implementation with the real weights: both
checkpoints load strict, all 68 tensors map with none left over, and the
outputs are identical (max absolute difference 0.0) to
VEMG2PoseWithInitialState and StatePoseModule respectively. left_context
derives to 1790, matching upstream's own accumulation.

Both are rehosted, so no conversion step is needed:

    VEMG2Pose.from_pretrained("braindecode/VEMG2Pose-emg2pose")
    VEMG2Pose.from_pretrained("braindecode/VEMG2Pose-emg2pose-tracking")

Parameter count moves from 6,269,992 to 5,980,328, which is the reference
network's own size; summary.csv follows.

The model figures now use braindecode's own docs URL rather than a
relative path, matching biot, so they resolve outside a local build.

* feat(vemg2pose): add the MLP decoder, covering every released checkpoint

emg2pose ships five checkpoints across three baselines, and the two named
after the paper itself were still unloadable: they use the same TdsNetwork
encoder this class already implements, but swap the recurrent decoder for
a state-conditioned MLP (Linear/LayerNorm/LeakyReLU twice, then the head),
predicting position directly.

Add a `decoder` argument for it. Both decoders take (feature, hidden,
cell) and return the same triple, so `forward` never branches on which one
is installed -- TorchScript compiles every branch it sees and would reject
a conditionally-typed submodule, the same trap that made _TdsStage keep an
Identity rather than a conditional linear_layer. The MLP decoder ignores
the recurrent state and passes it through.

All four checkpoints now load strict against the reference implementation,
with all 68 tensors mapped and no leftovers, and produce identical output
(max absolute difference 0.0):

    regression_vemg2pose  decoder=lstm  parameterization=hybrid
    tracking_vemg2pose    decoder=lstm  parameterization=velocity
    regression_emg2pose   decoder=mlp   parameterization=position
    tracking_emg2pose     decoder=mlp   parameterization=position

All four are rehosted, and each config records its own decoder and
parameterization so from_pretrained restores the right rollout:

    VEMG2Pose.from_pretrained("braindecode/VEMG2Pose-emg2pose")
    VEMG2Pose.from_pretrained("braindecode/VEMG2Pose-emg2pose-tracking")
    VEMG2Pose.from_pretrained("braindecode/EMG2Pose-emg2pose")
    VEMG2Pose.from_pretrained("braindecode/EMG2Pose-emg2pose-tracking")

With NeuroPose that is every checkpoint emg2pose released; SensingDynamics
has none, which the replication established by training it from scratch.

The decoder's parameters move under `decoder.` in the state_dict, which
the mapping derives rather than restates. Parameter count is unchanged.

* fix(models): harden EMG pose model contracts

* fix(tests): type EMG pose model factories

* refactor(sensingdynamics): keep SMU model-local

* test(models): remove duplicate pose coverage

* Fix VEMG2Pose valid-window outputs

* Simplify hosted EMG pose models

* refactor(models): require native EMG pose checkpoint keys
@bruAristimunha
bruAristimunha merged commit 65523b9 into braindecode:master Aug 31, 2026
11 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

Development

Successfully merging this pull request may close these issues.

1 participant