Conversation
|
Hi @ebezzam, thanks for your interest in contributing! This project requires that pull request authors are vouched, and you are not in the list of vouched users. This PR will be closed automatically. See https://github.com/huggingface/kernels-community/blob/main/CONTRIBUTING.md for more details. |
Implement the TDT loss kernels from scratch: - logprobs.cu: log-softmax gather over the vocabulary (online softmax, one block per lattice node) and a fused gradient kernel w.r.t. the logits. - lattice.cu: alpha/beta recursions over anti-diagonals, one block per sample with threads striding over u (no U <= 1024 limit). Logits are read in float32/float16/bfloat16 with arbitrary strides on the (batch, T, U) dimensions, so slices of the joint output need no copy. Supports all NeMo reductions (mean_volume, mean_batch, mean, sum, none). Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
alpha + beta - log_likelihood cancels in float32 when the log-likelihood is large (long targets). The recursions are ~2% of the runtime, so accumulate them in float64; log-probs and logits stay in float32 or lower. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Only the log-values are float64; exp/log act on differences in float32, and each node uses a two-pass logsumexp instead of chained log_add calls. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…rkers and flake.lock Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
- Device-side asserts for out-of-range lengths and targets (like PyTorch indexing), and reject negative durations. - __launch_bounds__ on all kernels so the largest block always launches. - Cache the durations tensor per device instead of copying it every call. - Drop the unused log-likelihood output of the beta recursion. - Move the PyTorch reference to tests/reference.py, used by the benchmark too. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
|
Hi @ebezzam, I picked this up and wrote the missing CUDA sources from scratch (the original Branch: https://github.com/Deep-unlearning/kernels-community/tree/tdt-loss-cuda (your commit, rebased onto current What's in it
Tests: 42 pass on L4 (sm89), A100 (sm80) and H200 (sm90). They check loss and gradients against the transformers PyTorch implementation, a naive loop version and a float64 reference. Benchmark (loss fwd+bwd at Parakeet size: batch 4, 200 frames, 40 labels, vocab 8193):
Fine-tune check (parakeet-tdt-0.6b-v3 on LibriSpeech, 300 steps, batch 8, bf16): same loss curve as the PyTorch loss (identical first-step loss, median per-step difference 0.0017). Full training step 1882 ms → 213 ms on A100 (1305 → 144 ms on H200), peak memory 33.6 → 21.8 GiB. (Run with I don't have push access here, so to bring it into this PR: git fetch https://github.com/Deep-unlearning/kernels-community.git tdt-loss-cuda
git push -f origin FETCH_HEAD:tdt_loss(force needed because the branch was rebased onto The nix build in CI will be its first run through kernel-builder (I compiled with PyTorch's JIT builder; the The transformers side is ready at https://github.com/Deep-unlearning/transformers/tree/tdt-loss-kernel (a fast-forward of huggingface/transformers#46048). |
9c3bdd4 to
f711276
Compare
…pping Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
There was a problem hiding this comment.
One important thing is to move this not into init imo and have a proper layers file (wrapped into nn module) so that we can use the exchange pattern as usual (e.g. see rope)
https://github.com/huggingface/kernels-community/blob/main/rotary/torch-ext/rotary/layers.py
There was a problem hiding this comment.
Thanks! Done in 3b9907c: the implementation is now in loss.py, and layers.py has a stateless TDTLoss layer with the same signature as transformers.loss.loss_tdt.tdt_loss (same idea as rotary/layers.py).
On the transformers side (huggingface/transformers#46048), tdt_loss is now decorated with use_kernel_forward_from_hub("tdt_loss") and attached to ParakeetForTDT with use_kernelized_func, with a tdt_loss → kernels-community/tdt-loss:TDTLoss mapping. So it's swapped in via use_kernels=True instead of being loaded automatically. Tested on GPU with a local build: the kernel matches the PyTorch loss and NeMo, and the Parakeet integration test passes with both.
Move the implementation from __init__.py to loss.py, add layers.TDTLoss (a stateless layer with the same signature as transformers' tdt_loss) and use the transformers argument order in the functional API.
|
Status / next steps
The "Comment with build instructions" check fails with a 403 from the bot's token. It does the same on other fork PRs (e.g. #1182). |
|
/kernel-bot build-and-stage tdt-loss |
|
Build request processed. Command: Dispatched (1):
Hub uploads: |
|
Using the staged build so you can use it via hf hub as well |

What does this PR do?
Goal of this PR is to add a kernel that has been drafted by @eustlb for the recently merged Parakeet TDT.
It's a very rough draft with a lot of Claude help as it's my first time working with kernels 😅
TODO
Perhaps we can see how other kernels used in Transformers were integrated. As the kernel is for a function, maybe rotary_pos_emb is a good example? Here are the files in this repo.
Here is a draft PR for adding/using this kernel in Transformers: huggingface/transformers#46048
cc @vasqu