Skip to content

Add TDT loss kernel - #46048

Draft
ebezzam wants to merge 89 commits into
huggingface:mainfrom
ebezzam:tdt_loss_kernel
Draft

ebezzam wants to merge 89 commits into
huggingface:mainfrom
ebezzam:tdt_loss_kernel

Conversation

@ebezzam

@ebezzam ebezzam commented May 19, 2026 •

Copy link
Copy Markdown
Contributor

CPU CI GPU run-slow

What does this PR do?

Use the kernels-community/tdt-loss CUDA kernel for the TDT loss of ParakeetForTDT, for faster training.

The kernel is swapped in with the usual kernels exchange pattern (as for rotary_pos_emb):

  • tdt_loss is decorated with use_kernel_forward_from_hub("tdt_loss") and attached to ParakeetForTDT with use_kernelized_func
  • the tdt_loss mapping points to the TDTLoss layer of kernels-community/tdt-loss (version 1), for training and inference (it does not support torch.compile)
  • so it is used with ParakeetForTDT.from_pretrained(..., use_kernels=True); the PyTorch implementation stays the default

Speed (loss fwd+bwd, batch 4, 200 frames, 40 labels, vocab 8193): 1259 ms → 9.5 ms on A100, 993 ms → 5.0 ms on H200. Fine-tuning parakeet-tdt-0.6b-v3 on LibriSpeech gives the same loss curve as the PyTorch loss, with a full training step going from 1882 ms to 213 ms on A100 and peak memory from 33.6 to 21.8 GiB.

Tests: test_use_kernels checks that use_kernels=True swaps tdt_loss for the kernel (in inference and training) with the same loss. The kernel's correctness (against the PyTorch loss, a naive loop version and float64) is tested in kernels-community.

Merge order: this needs kernels-community/tdt-loss to be published first (huggingface/kernels-community#882), otherwise use_kernels=True can't load it. Tested on GPU with the staged build of the kernel (kernels-staging/tdt-loss@pr-882).

cc @eustlb @vasqu

🤖 Generated with Claude Code

Hainan Xu and others added 30 commits February 20, 2026 09:45
Implement Token-and-Duration Transducer (TDT) decoding for Parakeet models,
extending the existing CTC-only support. This adds ParakeetForTDT with greedy
TDT decoding in generate(), per-token timestamp generation, and full
integration with AutoModelForTDT, processors, and ASR pipeline.
- Use -100 label padding for training (HF convention)
- Fix timestamp recording in inner blank-seeking loop
- Add max_symbols_per_step guard matching NeMo
- Clean up decoding loop
- Add TDT training example to docs
- Use setUpClass for TDT integration tests
@ebezzam
ebezzam marked this pull request as draft May 19, 2026 02:33
"finegrained-fp8": {"repo_id": "kernels-community/finegrained-fp8", "version": 1},
"deep-gemm": {"repo_id": "kernels-community/deep-gemm", "version": 1},
"sonic-moe": {"repo_id": "kernels-community/sonic-moe", "revision": "ep-support"},
"tdt-loss": {"repo_id": "eustlb/tdt-loss", "revision": "v1"},

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Comment on lines 736 to 738
Verify that ParakeetForTDT loss matches NeMo's TDT loss (sigma=0) for both
the CUDA kernel and the pure PyTorch implementation.
reproducer: https://gist.github.com/883ea42bf7d8ce2af42f3055627476a7

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Should we also test for sigma != 0?

@HuggingFaceDocBuilderDev

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

Deep-unlearning and others added 6 commits September 28, 2026 14:36
# Conflicts:
#	src/transformers/integrations/hub_kernels.py
#	tests/models/parakeet/test_modeling_parakeet.py
- Point the hub kernel mapping to kernels-community/tdt-loss (version 1).
- Dispatch to the kernel for CUDA inputs, for all reductions and sigma values.
- Test the kernel against the NeMo fixtures and the PyTorch implementation
  (including sigma != 0), and run the integration loss test with both.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Train mode updates the BatchNorm running statistics, which changed the eval
loss of the second backend.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…ntegration test loops

- Any error while loading the kernel (e.g. the Hub is unreachable) now falls
  back to the PyTorch implementation instead of failing the loss.
- Remove a duplicated forward/backward in the integration test and free the
  outputs inside each backend's subtest.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Decorate `tdt_loss` with `use_kernel_forward_from_hub` and attach it to
ParakeetForTDT with `use_kernelized_func`, so that `use_kernels=True`
swaps in the `TDTLoss` layer of kernels-community/tdt-loss (as for rotary),
instead of loading the kernel automatically.
},
},
"tdt_loss": {
"cuda": LayerRepository(repo_id="kernels-community/tdt-loss", layer_name="TDTLoss", version=1),

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.

iirc your kernel did not support compile right? lets give proper flags

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.

Right, it supports backward but not torch.compile (can_torch_compile = False on the layer). Done in 3f02948: the mapping now registers it for Mode.TRAINING and Mode.INFERENCE only.

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.

Imo no need to test it on this side as the kernel should test it for itself no? Meaning kernels community ensures that the loss is correct 🤔

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.

Makes sense, the kernel tests in kernels-community already check it against a PyTorch reference, a naive loop version and float64. Done in 3f02948: I removed the loss-correctness tests here and restored the integration test, keeping only test_use_kernels, which checks that use_kernels=True actually swaps tdt_loss for the kernel (inference and training) with the same loss.

Also tested against your staged build (kernels-staging/tdt-loss@pr-882) on an L4: the kernels-community tests pass (43), the Parakeet tests pass with use_kernels=True (including the slow integration test and the full test file), and the speed matches my earlier numbers (loss fwd+bwd 1285 ms → 53 ms at Parakeet size).

… the kernel exchange test

- Register the kernel for training and inference only (it does not support torch.compile).
- The kernel's correctness is tested in kernels-community: keep only the test that
  use_kernels=True swaps tdt_loss for it, and restore the integration test.
@github-actions

Copy link
Copy Markdown
Contributor

[For maintainers] Suggested jobs to run (before merge)

run-slow: parakeet

@github-actions

Copy link
Copy Markdown
Contributor

CI recap

Dashboard: View test results in Grafana
Latest run: 36543082143:1
Result: success | Jobs: 16 | Tests: 182,426 | Failures: 0 | Duration: 15h 49m

@require_torch
@require_torch_gpu
@require_kernels
class TDTLossKernelTest(unittest.TestCase):

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.

lets make this a test under the model instead of a separate class pls

),
},
},
"tdt_loss": {

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.

We might have to drop the kernel for when we make use kernels true as default. best to make it toch compile friendly from the get go

This branch has not been deployed

No deployments
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.

6 participants