diff --git a/src/transformers/integrations/hub_kernels.py b/src/transformers/integrations/hub_kernels.py index 76344e5453e5..b407bba7645c 100644 --- a/src/transformers/integrations/hub_kernels.py +++ b/src/transformers/integrations/hub_kernels.py @@ -290,6 +290,7 @@ def register_kernel_mapping_transformers(*args, **kwargs): "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"}, } _KERNEL_MODULE_MAPPING: dict[str, ModuleType | None] = {} diff --git a/src/transformers/loss/loss_tdt.py b/src/transformers/loss/loss_tdt.py index 9493c56cc6a2..3e474fb20ea4 100644 --- a/src/transformers/loss/loss_tdt.py +++ b/src/transformers/loss/loss_tdt.py @@ -20,6 +20,17 @@ logger = logging.get_logger(__name__) +def _load_tdt_kernel(): + """Try to load the TDT loss CUDA kernel from the Hub. Returns None on failure.""" + from ..integrations.hub_kernels import lazy_load_kernel + + kernel = lazy_load_kernel("tdt-loss") + if kernel is None or not hasattr(kernel, "tdt_loss"): + logger.warning_once("Falling back to pure PyTorch implementation.") + return None + return kernel + + def tdt_loss( token_logits: torch.Tensor, duration_logits: torch.Tensor, @@ -38,6 +49,9 @@ def tdt_loss( the token prediction head and the duration prediction head. It uses vectorized anti-diagonal processing for efficiency: all (t, u) pairs on each anti-diagonal t+u=n are computed in parallel as batched tensor operations. + When the ``kernels-community/tdt-loss`` CUDA kernel is installed, it is used automatically for GPU tensors, + Falls back to the pure PyTorch implementation otherwise. + Args: token_logits: Token logits of shape `(batch, T, U+1, vocab_size+1)`. duration_logits: Duration logits of shape `(batch, T, U+1, num_durations)`. @@ -53,6 +67,20 @@ def tdt_loss( Scalar loss tensor (or per-example losses if `reduction="none"`). """ + kernel = _load_tdt_kernel() if token_logits.is_cuda else None + if kernel is not None and hasattr(kernel, "tdt_loss"): + durations_t = torch.tensor(durations, dtype=torch.int32, device=token_logits.device) + return kernel.tdt_loss( + token_logits, + duration_logits, + targets, + logit_lengths, + target_lengths, + durations_t, + blank_token_id, + sigma, + reduction, + ) if reduction not in ("mean", "sum", "none"): raise ValueError(f'Invalid reduction mode "{reduction}". Expected one of "mean", "sum", or "none".') diff --git a/tests/models/parakeet/test_modeling_parakeet.py b/tests/models/parakeet/test_modeling_parakeet.py index f2324dc4d733..2c6d219797aa 100644 --- a/tests/models/parakeet/test_modeling_parakeet.py +++ b/tests/models/parakeet/test_modeling_parakeet.py @@ -16,7 +16,9 @@ import json import tempfile import unittest +from contextlib import nullcontext from pathlib import Path +from unittest.mock import patch from transformers import is_datasets_available, is_torch_available from transformers.testing_utils import cleanup, require_torch, slow, torch_device @@ -731,9 +733,12 @@ def test_tdt_model_integration_timestamps(self): @slow def test_tdt_model_integration_loss(self): """ - Verify that ParakeetForTDT loss matches NeMo's TDT loss (sigma=0). + 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 """ + from transformers.loss.loss_tdt import _load_tdt_kernel + RESULTS_PATH = Path(__file__).parent.parent.parent / "fixtures/parakeet/expected_loss_tdt.json" with open(RESULTS_PATH, "r") as f: raw_data = json.load(f) @@ -754,20 +759,33 @@ def test_tdt_model_integration_loss(self): ) inputs.to(model.device) - # Forward in eval mode — check loss matches NeMo - model.eval() - with torch.no_grad(): - outputs = model(**inputs) - self.assertIsNotNone(outputs.loss, "Loss must be computed when labels are provided") - self.assertEqual(outputs.logits.dim(), 4, "Training logits must be 4D (B, T, U+1, V+D)") - torch.testing.assert_close(outputs.loss.cpu(), EXPECTED_MEAN_LOSS, rtol=1e-3, atol=1e-3) - - # Backward — verify gradients flow - del outputs - torch.cuda.empty_cache() - model.train() - model.zero_grad() - outputs = model(**inputs) - outputs.loss.backward() - n_with_grad = sum(1 for p in model.parameters() if p.grad is not None) - self.assertGreater(n_with_grad, 0, "No gradients after backward") + # Test both backends: kernel (if available) and pure PyTorch + has_kernel = _load_tdt_kernel() is not None + backends = [ + ("kernel", None), + ("torch", patch("transformers.loss.loss_tdt._load_tdt_kernel", return_value=None)), + ] + if not has_kernel: + backends = backends[1:] # skip kernel test when not installed + + for backend_name, ctx in backends: + with self.subTest(backend=backend_name): + ctx_manager = ctx if ctx is not None else nullcontext() + with ctx_manager: + # Forward in eval mode — check loss matches NeMo + model.eval() + with torch.no_grad(): + outputs = model(**inputs) + self.assertIsNotNone(outputs.loss, "Loss must be computed when labels are provided") + self.assertEqual(outputs.logits.dim(), 4, "Training logits must be 4D (B, T, U+1, V+D)") + torch.testing.assert_close(outputs.loss.cpu(), EXPECTED_MEAN_LOSS, rtol=1e-3, atol=1e-3) + + # Backward — verify gradients flow + del outputs + torch.cuda.empty_cache() + model.train() + model.zero_grad() + outputs = model(**inputs) + outputs.loss.backward() + n_with_grad = sum(1 for p in model.parameters() if p.grad is not None) + self.assertGreater(n_with_grad, 0, "No gradients after backward")