Add TDT loss kernel - #46048
Add TDT loss kernel#46048ebezzam wants to merge 89 commits into
Conversation
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
Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com>
| "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"}, |
There was a problem hiding this comment.
| 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 |
There was a problem hiding this comment.
Should we also test for sigma != 0?
|
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. |
# 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), |
There was a problem hiding this comment.
iirc your kernel did not support compile right? lets give proper flags
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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 🤔
There was a problem hiding this comment.
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.
|
[For maintainers] Suggested jobs to run (before merge) run-slow: parakeet |
CI recapDashboard: View test results in Grafana |
| @require_torch | ||
| @require_torch_gpu | ||
| @require_kernels | ||
| class TDTLossKernelTest(unittest.TestCase): |
There was a problem hiding this comment.
lets make this a test under the model instead of a separate class pls
| ), | ||
| }, | ||
| }, | ||
| "tdt_loss": { |
There was a problem hiding this comment.
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
What does this PR do?
Use the
kernels-community/tdt-lossCUDA kernel for the TDT loss ofParakeetForTDT, for faster training.The kernel is swapped in with the usual kernels exchange pattern (as for
rotary_pos_emb):tdt_lossis decorated withuse_kernel_forward_from_hub("tdt_loss")and attached toParakeetForTDTwithuse_kernelized_functdt_lossmapping points to theTDTLosslayer ofkernels-community/tdt-loss(version 1), for training and inference (it does not supporttorch.compile)ParakeetForTDT.from_pretrained(..., use_kernels=True); the PyTorch implementation stays the defaultSpeed (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-v3on 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_kernelschecks thatuse_kernels=Trueswapstdt_lossfor 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-lossto be published first (huggingface/kernels-community#882), otherwiseuse_kernels=Truecan'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