Skip to content

tdt-loss: add TDT loss kernel - #882

Open
ebezzam wants to merge 12 commits into
huggingface:mainfrom
ebezzam:tdt_loss
Open

ebezzam wants to merge 12 commits into
huggingface:mainfrom
ebezzam:tdt_loss

Conversation

@ebezzam

@ebezzam ebezzam commented May 19, 2026

Copy link
Copy Markdown

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

@ebezzam
ebezzam requested review from danieldk and drbh as code owners May 19, 2026 07:51
@github-actions

Copy link
Copy Markdown

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.

@github-actions github-actions Bot closed this May 19, 2026
@ebezzam ebezzam changed the title Initial commit, missing .cu file. Add TDT loss kernel May 19, 2026
@danieldk danieldk reopened this May 19, 2026
@danieldk danieldk changed the title Add TDT loss kernel tdt-loss: add TDT loss kernel May 19, 2026
@drbh drbh added area: build-system build.toml, Nix flakes, packaging, and kernel-builder integration area: docs README, CARD.md, guides, and repository documentation area: tests Tests, validation code, and benchmark harnesses new-kernel A brand-new kernel package is added size: L Diff <= 1000 lines stale No update in 30 days type: feature New functionality / capability feature New functionality / capability cuda NVIDIA CUDA kernels and removed stale No update in 30 days area: build-system build.toml, Nix flakes, packaging, and kernel-builder integration area: docs README, CARD.md, guides, and repository documentation area: tests Tests, validation code, and benchmark harnesses size: L Diff <= 1000 lines type: feature New functionality / capability labels Jun 30, 2026
ebezzam and others added 7 commits September 28, 2026 14:26
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>
Deep-unlearning and others added 3 commits September 28, 2026 17:14
…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>
@Deep-unlearning

Copy link
Copy Markdown

Hi @ebezzam, I picked this up and wrote the missing CUDA sources from scratch (the original .cu files were never published, only the compiled .so).

Branch: https://github.com/Deep-unlearning/kernels-community/tree/tdt-loss-cuda (your commit, rebased onto current main, plus mine)

What's in it

  • tdt_loss/logprobs.cu: log-softmax over the vocabulary for each lattice node, plus one fused kernel for the gradients w.r.t. the token and duration logits
  • tdt_loss/lattice.cu: alpha/beta passes over the lattice (float64 accumulation for precision on long label sequences, no 1024-label limit)
  • reads fp32/fp16/bf16 logits directly, and no copy for slices of the joint output (logits[..., :V])
  • supports all reductions: mean_volume, mean_batch, mean, sum, none
  • tests load the kernel with get_kernel, have kernels_ci markers, and there's a flake.lock, per AGENTS.md

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):

GPU PyTorch (fp32) kernel (fp32 / bf16)
L4 1252 ms 53 / 39 ms
A100 1259 ms 9.5 / 8.6 ms
H200 993 ms 5.0 / 4.9 ms

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.

Fine-tuning loss: CUDA kernel vs PyTorch TDT loss

(Run with layerdrop=0: the converted Parakeet configs have layerdrop=0.1 while NeMo trains without stochastic depth, which occasionally drops the last encoder layer and gives ~1000 loss spikes with either loss. I'll fix that separately in transformers.)

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 main)

The nix build in CI will be its first run through kernel-builder (I compiled with PyTorch's JIT builder; the flake.lock is copied from rotary).

The transformers side is ready at https://github.com/Deep-unlearning/transformers/tree/tdt-loss-kernel (a fast-forward of huggingface/transformers#46048).

@github-actions github-actions Bot added chore Version bumps, releases, misc maintenance and removed cuda NVIDIA CUDA kernels feature New functionality / capability new-kernel A brand-new kernel package is added labels Sep 28, 2026
…pping

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

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.

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

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

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.
@Deep-unlearning

Copy link
Copy Markdown

Status / next steps

  • CUDA kernels, tests (43 passing on L4/A100/H200), benchmarks, README
  • layers.TDTLoss for the kernels exchange pattern (thanks @vasqu)
  • Registered in kernel-maintainers.json and check_kernel_freshness.py
  • Build: could someone with write access run /kernel-bot build tdt-loss? This is the first kernel-builder/nix build of this kernel (I could only test with the PyTorch JIT builder).
  • Review and merge, then publish kernels-community/tdt-loss v1
  • Then merge Add TDT loss kernel transformers#46048 (it needs the published kernel for use_kernels=True)

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).

@vasqu

vasqu commented Sep 28, 2026

Copy link
Copy Markdown
Collaborator

/kernel-bot build-and-stage tdt-loss

@github-actions

github-actions Bot commented Sep 28, 2026 •

Copy link
Copy Markdown

Build request processed.

Command: /kernel-bot build-and-stage tdt-loss
Mode: build and stage
Target branch: pr-882
PR head SHA: 3b9907c58a98e3874b16b454a72fe8fb2a8b67c7
Workflows: build.yaml, build-mac.yaml, build-windows.yaml

Dispatched (1):

Hub uploads:

@vasqu

vasqu commented Sep 28, 2026

Copy link
Copy Markdown
Collaborator

Using the staged build so you can use it via hf hub as well

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

chore Version bumps, releases, misc maintenance

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants