Skip to content

[models] Add TFMTokenizer - #1202

Merged
bruAristimunha merged 28 commits into
braindecode:masterfrom
lindicaphxag-tech:feat/tfm-tokenizer-1201
Oct 9, 2026
Merged

bruAristimunha merged 28 commits into
braindecode:masterfrom
lindicaphxag-tech:feat/tfm-tokenizer-1201

Conversation

@lindicaphxag-tech

@lindicaphxag-tech lindicaphxag-tech commented Oct 2, 2026 •

Copy link
Copy Markdown
Contributor

Reviewer map

Fast path:

  1. braindecode/models/tfm_tokenizer.py — architecture, EMA VQ semantics, STFT/masking/token API;
  2. test/unit_tests/models/test_models.py — shape, gradient, EMA-state, config/state and edge-case contracts;
  3. registry/docs/NOTICE changes — standard integration and MIT attribution.

Two things are intentionally separated in the implementation:

  • released-checkpoint parity: all 191 mapped state entries, 100% token-ID agreement, zero max-abs error for audited forward internals/reconstruction and reconstruction-path gradients;
  • library-level corrections: EMA VQ gradient routing and never-selected-code handling are deliberate correctness fixes and are not presented as source parity.

No paper benchmark reproduction is claimed.

Summary

Adds TFMTokenizer, the self-supervised time-frequency motif tokenizer from Tokenizing Single-Channel EEG with Time-Frequency Motif Learning (ICLR 2026).

Part of #1201.

API contract

  • forward(X) remains Tensor-valued and returns the reconstructed magnitude spectrogram.
  • tokenize(...) exposes the full SSL payload: reconstruction, token IDs, straight-through quantized embeddings, pre-quantization embeddings, VQ loss, and target spectrogram.
  • No downstream classifier is included.

Reference and checkpoint fidelity

Reference: Jathurshan0330/TFM-Tokenizer@2d6da482b16dabbb2ebec808fa9f505fc6f367c4 (MIT).

Released checkpoint SHA-256:
f7929eed0f8275d08a7e18743fdd188c7ba5b4722151333b4d386ca05b84d6d9.

Pinned CPU parity assay:

  • all 191 checkpoint state entries mapped;
  • token-ID agreement: 100%;
  • STFT magnitudes, pre-quantization embeddings, quantized embeddings, and reconstructions: max abs error 0;
  • reconstruction-path gradients for the EEG input and first frequency/temporal convolution weights: max abs error 0.

Evidence:

  • the key numerical parity results are reproduced inline above; the underlying local audit artifacts are not required for library execution or tests

Deliberate library-level corrections

Two training-state issues are intentionally not presented as reference parity:

  1. EMA VQ gradient routing. The released straight-through formulation sends equal-and-opposite encoder gradients through the two VQ terms at the default commitment_cost=1.0. This port keeps the same scalar objective value, but makes the EMA codebook term value-only and leaves the commitment term responsible for the encoder gradient. The codebook remains EMA-only.
  2. Never-selected EMA codes. Unseen code vectors remain at their initialized locations until first assignment instead of being normalized by an epsilon-scale count. Selected centroids still follow the EMA update.

Both behaviors have focused regressions.

Braindecode integration

  • EEGModuleMixin integration and model registry entry;
  • summary table, API docs, categorization, reference bibliography, and What's New entry;
  • reference 2/2/8 encoder/decoder depths and STFT geometry;
  • complementary time-frequency masking;
  • config/state reconstruction coverage;
  • torch.compile integration coverage.

TorchScript is not claimed because the existing third-party LinearAttentionTransformer dependency uses a variadic forward signature that is not scriptable; the integration suite skips that check for this model only.

Claim boundary

This PR claims released-checkpoint parity for the tokenizer's mapped state, forward/token outputs, internal audited representations, and reconstruction gradients. It does not claim reproduction of the paper's pretraining/downstream benchmark metrics. Upstream workflows currently require maintainer approval for this fork, so no current-head upstream CI pass is claimed.

Copy link
Copy Markdown
Contributor Author

CI blocker follow-up on the current head:

  • Added the required Foundation Model + Attention/Transformer doc badges.
  • The default forward(X) is now Tensor-valued (reconstruction), while the complete self-supervised training payload is exposed explicitly through tokenize(...).
  • Removed the private _compute_spectrogram / _encode method chain from the scripted path so convert_model_to_plain() retains every method used by forward.
  • Applied the two docstrfmt changes reported by pre-commit.
  • Model-specific structured-output tests now call tokenize(...); a regression checks the default Tensor forward contract.

These changes address the four failures from the last upstream run without adding integration-test skips. Could the workflow be re-run on the current head when convenient?

Copy link
Copy Markdown
Contributor Author

Reference-fidelity correction on current head 5e6d18ef after re-auditing the released TFM-Tokenizer code:

  • The released tokenizer updates quantizer.embedding.weight through EMA; after nearest-neighbor lookup it applies quant_out = quant_in + (quant_out - quant_in).detach(), so the subsequent VQ loss does not send optimizer gradients into the codebook table.
  • The port now makes that contract explicit: lookup uses a detached codebook, encoder gradients still flow through the straight-through estimator, and the codebook remains EMA-only.
  • Regressions now require encoder-path gradients, no codebook optimizer gradient, a real EMA update in train mode, and no second codebook change from the optimizer step.
  • I also fixed a misplaced pytest.mark.parametrize decorator in the architecture-validation tests that would otherwise break collection when upstream CI is approved.

No architecture dimensions or reconstruction path changed; this is a correctness/fidelity cleanup before maintainer review.

Copy link
Copy Markdown
Contributor Author

EMA robustness follow-up on the current head fd494c01:

  • I re-audited the released TFM-Tokenizer EMA update and found a library-level edge case inherited from the reference implementation: cluster_size starts at zero while ema_weight starts from the random codebook. After the first training batch, never-selected codes are normalized by eps; with the reference defaults (codebook_size=8192, eps=1e-5) this can inflate an unseen code by roughly 1/eps and move it far outside the L2-normalized encoder space.
  • The port now leaves never-selected code vectors at their initialized locations until their first assignment. EMA statistics are still updated exactly as before, and occupied codes still use the same EMA centroid normalization.
  • A focused regression uses more codebook entries than assignments and requires all unselected entries to remain unchanged, finite, and near their initialized scale.

This deliberately improves numerical robustness rather than claiming checkpoint parity with the research repository; the tokenizer architecture, VQ objective, straight-through estimator, and checkpoint state schema are unchanged.

Copy link
Copy Markdown
Contributor Author

VQ gradient audit on the current head: with the straight-through tensor used in both VQ terms, the EMA-only port sent equal-and-opposite gradients into the encoder when commitment_cost=1.0 (the reference default). I changed only the gradient routing, not the objective value: the codebook-distance term now uses the raw detached EMA codebook vector, the commitment term trains the encoder, reconstruction still uses the straight-through path, and the codebook remains EMA-only. The regression now requires a strictly non-zero encoder gradient rather than only checking grad is not None.

Copy link
Copy Markdown
Contributor Author

One more VQ-gradient invariant is now explicit on head e9a8138a.

The released code applies the straight-through estimator before evaluating both VQ MSE terms. If
q_st = z + (q - z).detach(), then the two encoder gradients are opposite:

  • ∇z MSE(q_st, z.detach()) ∝ +(q - z)
  • ∇z MSE(q_st.detach(), z) ∝ -(q - z)

At the released/default commitment_cost=1.0, those contributions cancel exactly. That makes the nominal VQ objective produce zero encoder gradient even though the codebook itself is EMA-updated.

The current port keeps the same scalar objective value but evaluates the codebook term from the detached EMA codebook vector itself. The codebook term is therefore value-only (the codebook stays EMA-only), while the commitment term supplies the encoder gradient. A focused regression now requires nonzero frequency- and temporal-encoder gradients at commitment_cost=1.0 and still requires no optimizer gradient on the codebook embedding.

This is a deliberate library-level correctness fix rather than a checkpoint-parity claim; reconstruction architecture and EMA state remain unchanged.

Copy link
Copy Markdown
Contributor Author

Pinned official-checkpoint parity is now available for the current PR head abb99259f8da6a02e0c0366e7feb39363e3a617c.

Reference: Jathurshan0330/TFM-Tokenizer@2d6da482b16dabbb2ebec808fa9f505fc6f367c4; official 2x2x8 checkpoint SHA256 f7929eed0f8275d08a7e18743fdd188c7ba5b4722151333b4d386ca05b84d6d9.

CPU assay results:

  • all 191 checkpoint state entries mapped;
  • token agreement: 1.000;
  • spectrogram / pre-quant embeddings / straight-through quantized embeddings / reconstruction: max abs error 0;
  • reconstruction-path gradients for raw EEG input, first frequency conv, and first temporal conv: max abs error 0.

Run: https://github.com/lindicaphxag-tech/lindicaphxag-tech/actions/runs/37114458814

The EMA-VQ objective-gradient correction remains an explicitly documented library-level deviation and is intentionally excluded from the reference-parity oracle rather than presented as reference behavior.

Copy link
Copy Markdown
Contributor Author

@bruAristimunha @julien-gadonneix — this is now at the point where I would value a maintainer pass on the model/API contract rather than further self-hardening. The current head keeps the default forward(X) Tensor-valued, isolates the full SSL payload in tokenize(...), makes the EMA-only codebook gradient route explicit, and includes checkpoint/config/EMA regressions. I’ll avoid additional scope unless review identifies a concrete issue.

Copy link
Copy Markdown
Contributor Author

Follow-up on the EMA regression: the failure exposed a real update bug rather than a flaky test. embedding.weight[occupied].copy_(...) was writing into the temporary tensor produced by boolean advanced indexing, so selected EMA centroids never reached the actual codebook. Head eaf26ab1 switches that one line to indexed assignment; unseen entries still remain unchanged. Independent current-head validation is green: https://github.com/lindicaphxag-tech/lindicaphxag-tech/actions/runs/37116304219. I’ll hold the PR here for maintainer review.

Copy link
Copy Markdown
Contributor Author

Review-history cleanup only: I squashed the branch to a single commit before maintainer review. The Git tree is unchanged (c6e5121cf25a8f4c2bb91ac5ed18fd87822c8a61), so the existing focused/parity evidence remains applicable. I’ll hold this head unless review identifies a concrete issue.

@codecov

codecov Bot commented Oct 4, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 92.07921% with 8 lines in your changes missing coverage. Please review.
✅ Project coverage is 89.19%. Comparing base (83a1e67) to head (37afb53).
⚠️ Report is 2 commits behind head on master.

Additional details and impacted files
@@            Coverage Diff             @@
##           master    #1202      +/-   ##
==========================================
+ Coverage   89.16%   89.19%   +0.02%     
==========================================
  Files         158      159       +1     
  Lines       19641    19741     +100     
==========================================
+ Hits        17513    17607      +94     
- Misses       2128     2134       +6     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@bruAristimunha bruAristimunha added model Adds a new model needs-replication Model PR: paper number must be replicated (NeuralBench) before merge labels Oct 5, 2026
lindicaphxag-tech and others added 15 commits October 6, 2026 20:42
Conflict in docs/whats_new.rst: rebuilt from master's file plus the TFMTokenizer entry.
The private _EMAVectorQuantizer duplicated EMACodebook. With epsilon=1e-5, no dead-code expiry and no k-means init, EMACodebook is the reference EMA update; the uniform init is kept. Never-selected codes now follow the reference update again. Buffers are renamed (embed, embed_avg, inited); the loading snippet maps the released keys. EMACodebook skips the inited host read when k-means init is off.
bfloat16 and HPU inputs now get their spectrogram in float32 on the CPU. The model ships released weights, so the braindecode#1252 coverage test needs a COMPAT entry (channel-agnostic, 200 Hz).

@bruAristimunha bruAristimunha left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Verified: TFMTokenizer passes master's model contract, integration and pretrained-compat suites, including the new CPU accelerator and low-precision checks, in our CPU test jobs, and its tokens match the authors' code and released checkpoint exactly. Thanks @lindicaphxag-tech for the careful TFM-Tokenizer port!

@bruAristimunha
bruAristimunha merged commit 0e47114 into braindecode:master Oct 9, 2026
23 of 24 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

model Adds a new model needs-replication Model PR: paper number must be replicated (NeuralBench) before merge

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants