diff --git a/.github/workflows/test-gpu.yml b/.github/workflows/test-gpu.yml index a334e11ab..397bd74a9 100644 --- a/.github/workflows/test-gpu.yml +++ b/.github/workflows/test-gpu.yml @@ -73,7 +73,7 @@ jobs: --find-links=https://data.pyg.org/whl/torch-2.14.0+cu130.html \ --index-strategy=unsafe-best-match \ --group=test \ - . + '.[cudnn]' - name: GPU sanity check run: | diff --git a/README.md b/README.md index 788755733..239c11488 100644 --- a/README.md +++ b/README.md @@ -33,6 +33,8 @@ pip install git+https://github.com/NVIDIA/structured-data-models.git > [!NOTE] > For CUDA workloads, we highly recommend installing [`cudf`](https://docs.rapids.ai/install) as an additional dependency to keep dataframe-style operations on GPU and avoid unnecessary data movement. +The optional `cudnn` extra (`pip install "structured-data-models[cudnn]"`) adds cuDNN variable-length attention for padded inputs, enabled via `sdm.nn.enable_cudnn_varlen()` (Linux only; on other platforms the extra installs nothing and the boolean-mask path is used). + ## Model Families **Tabular Foundation Models:** diff --git a/docs/source/install.md b/docs/source/install.md index 5d88936b5..2625f03af 100644 --- a/docs/source/install.md +++ b/docs/source/install.md @@ -10,3 +10,5 @@ pip install git+https://github.com/NVIDIA/structured-data-models.git ```{note} For CUDA workloads, we highly recommend installing [`cudf`](https://docs.rapids.ai/install) as an additional dependency to keep dataframe-style operations on GPU and avoid unnecessary data movement. ``` + +The optional `cudnn` extra (`pip install "structured-data-models[cudnn]"`) adds cuDNN variable-length attention for padded inputs, enabled via `sdm.nn.enable_cudnn_varlen()` (Linux only; on other platforms the extra installs nothing and the boolean-mask path is used). diff --git a/pyproject.toml b/pyproject.toml index 0b7b94e25..2560bcf55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,11 @@ dependencies = [ "typing_extensions", ] +[project.optional-dependencies] +cudnn = [ + "nvidia-cudnn-frontend>=1.23,<1.29; sys_platform == 'linux'", +] + [project.urls] Download = "https://github.com/NVIDIA/structured-data-models/releases" Homepage = "https://github.com/NVIDIA/structured-data-models" diff --git a/sdm/nn/__init__.py b/sdm/nn/__init__.py index 1be3d7235..2274074a4 100644 --- a/sdm/nn/__init__.py +++ b/sdm/nn/__init__.py @@ -3,6 +3,7 @@ """Neural network modules for structured data models.""" +from sdm.nn._cudnn_varlen import cudnn_varlen_stats, enable_cudnn_varlen from sdm.nn.rope import RotaryEmbedding from sdm.nn.glu import SwiGLU from sdm.nn.softplus import SoftplusScale @@ -13,6 +14,8 @@ __all__ = [ + "cudnn_varlen_stats", + "enable_cudnn_varlen", "RotaryEmbedding", "SwiGLU", "SoftplusScale", diff --git a/sdm/nn/_cudnn_varlen.py b/sdm/nn/_cudnn_varlen.py new file mode 100644 index 000000000..67d8deb5e --- /dev/null +++ b/sdm/nn/_cudnn_varlen.py @@ -0,0 +1,408 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""cuDNN-frontend variable-length SDPA for padded key/value streams.""" + +import math +import threading +from typing import Any + +import torch +import torch.nn.functional as F +from torch import Tensor + +from sdm._warnings import warn_once + +try: + import cudnn as _cudnn_fe # ty: ignore[unresolved-import] +except Exception: # noqa: BLE001 - the frontend dlopens libcudnn at import + _cudnn_fe = None # pragma: no cover - optional dependency + +_enabled = False +# Serializes handle/stream binding and graph-cache access: the module +# targets single-stream serving; concurrent callers are correct but +# serialized (the lock only orders host-side enqueue). +_lock = threading.Lock() +# Execution graphs are cached per (batch, heads, lengths, head-dim, +# dtype, device) shape, matching the bucketed-serving design where the +# shape set is finite; the valid lengths remain runtime tensor data, so +# varying them never rebuilds a graph. +_graph_cache: dict[tuple[Any, ...], Any] = {} +_handles: dict[int, Any] = {} +_UNPROBED = object() + + +def enable_cudnn_varlen(enabled: bool = True) -> bool: + """Toggle the cuDNN variable-length attention path. + + When enabled, eligible attention calls that provide + ``seqused_key_value`` replace the boolean-mask path in + :class:`~sdm.nn.SDPA`: instead of materializing a mask and routing + to the masked kernels, attention runs on cuDNN's native + padding-mask support with per-batch valid lengths (``seq_len_kv``), + which bounds the computation to the valid region. + + When cuDNN cannot build a plan for a shape, the op degrades to an + equivalent masked fallback at runtime (probed once per shape, + negatively cached, warned once), so ineligible shapes, a missing + ``nvidia-cudnn-frontend``, or graph-build failures keep today's + behavior; :func:`cudnn_varlen_stats` reports how many probed + shapes built an execution graph versus degraded. The path targets + inference (gradient-enabled calls are ineligible). The op runs + outside :func:`torch.nn.functional.scaled_dot_product_attention`, + so :func:`torch.nn.attention.sdpa_kernel` and + :func:`torch.backends.cuda.enable_cudnn_sdp` do not reach it; + disable this path to return to PyTorch's backend selection. Under + CUDA graphs (``mode="reduce-overhead"`` or a manual capture) a shape + must have run once eagerly - the compiler's warmup run suffices - + before it is recorded; a shape first met inside a capture replays + the masked fallback and is probed on its next eager call. + + One execution graph is cached per distinct eligible shape (batch, + heads, query/key lengths, head dim, dtype, device) and reused for + every call at that shape; disabling the path drops the cache entries. + This is the same per-shape policy as PyTorch's own cuDNN SDPA graph + cache. Each graph costs roughly 2 MB of host memory and a few tens + of milliseconds to build, so pair the path with shape bucketing + (``seqused_train`` padding, as the serving recipe does); a stream of + unbucketed table shapes pays that cost once per new shape. Treat the + toggle as a setup-time switch: the frontend keeps about 1 MB of host + memory per dropped graph, so disabling and re-enabling around every + request rebuilds the graphs and slowly grows the process. Compiled + callers guard on the toggle and recompile when it changes; a + ``mode="reduce-overhead"`` graph recorded while the path was enabled + keeps replaying its recorded cuDNN kernels after such a cycle without + probing again, so :func:`cudnn_varlen_stats` no longer counts those + shapes. + + Args: + enabled: Whether eligible ``seqused_key_value`` attention calls + should use cuDNN's native padding-mask support instead of a + boolean mask. Enabling has no effect when the optional + ``nvidia-cudnn-frontend`` package (installable via the + ``cudnn`` extra: ``pip install 'structured-data-models[cudnn]'``) + is unavailable; disabling clears the cached execution + graphs. + + Returns: + Whether the path is active after the call. + """ + global _enabled + _enabled = bool(enabled) and _cudnn_fe is not None + if not _enabled: + with _lock: + _graph_cache.clear() + return _enabled + + +def cudnn_varlen_stats() -> tuple[int, int]: + """Return ``(built, degraded)`` engagement counts for probed shapes. + + ``built`` counts shapes serving on a cuDNN execution graph; + ``degraded`` counts shapes whose graph build failed and were + negatively cached, so their calls take the masked fallback. Both + reset to zero when the path is disabled via + :func:`enable_cudnn_varlen`. The counts reflect host-side probes: + a CUDA graph replay (``mode="reduce-overhead"`` or a manual capture) + that was recorded before such a reset re-executes the recorded + kernels without probing and is not counted again. + + Returns: + The ``(built, degraded)`` pair. + """ + with _lock: + graphs = list(_graph_cache.values()) + built = sum(graph is not None for graph in graphs) + return built, len(graphs) - built + + +def is_available() -> bool: + """Return whether the optional cuDNN frontend is importable.""" + return _cudnn_fe is not None + + +def _shape_eligible( + query: Tensor, # [B, Q, H, D] (flattened batch, pre-transpose layout) + num_query_heads: int, + num_key_value_heads: int, +) -> bool: + """The shape/dtype half of :func:`eligible`, independent of device/grad.""" + return ( + query.dtype in (torch.bfloat16, torch.float16) + and num_query_heads == num_key_value_heads + and query.size(-1) % 8 == 0 + and query.size(-1) <= 128 + and query.size(0) <= 65535 + ) + + +def eligible( + query: Tensor, # [B, Q, H, D] (flattened batch, pre-transpose layout) + key: Tensor, # [B, KV, H, D] + num_query_heads: int, + num_key_value_heads: int, +) -> bool: + """Return whether this call can take the variable-length path.""" + # The graph derives its dtype and device from ``query`` alone and then + # binds ``key``/``value`` to whatever arrives, so a mismatched key + # would be reinterpreted instead of rejected. The boolean-mask path + # raises in that case; keep the two consistent. ``value`` and the + # count tensor are not visible here; the op itself routes their + # mismatches to the masked fallback, which raises the same way. + return ( + _enabled + and _cudnn_fe is not None + and query.is_cuda + and key.dtype == query.dtype + and key.device == query.device + and not torch.is_grad_enabled() + and _shape_eligible(query, num_query_heads, num_key_value_heads) + ) + + +class _Graph: + """One built cuDNN execution graph for a fixed shape.""" + + def __init__( + self, + B: int, + H: int, + Q: int, + KV: int, + D: int, + dtype: torch.dtype, + device: torch.device, + ) -> None: + assert _cudnn_fe is not None + fe_dtype = ( + _cudnn_fe.data_type.BFLOAT16 + if dtype == torch.bfloat16 + else _cudnn_fe.data_type.HALF + ) + index = device.index or 0 + with torch.cuda.device(device): + if index not in _handles: + _handles[index] = _cudnn_fe.create_handle() + handle = _handles[index] + self.handle = handle + graph = _cudnn_fe.pygraph( + io_data_type=fe_dtype, + intermediate_data_type=_cudnn_fe.data_type.FLOAT, + compute_data_type=_cudnn_fe.data_type.FLOAT, + handle=handle, + ) + + # Logical BHSD dims over the caller's [B, S, H, D] contiguous + # storage (BSHD physical layout - cuDNN's preferred layout). + def _tensor(name: str, S: int) -> Any: + return graph.tensor( + name=name, + dim=(B, H, S, D), + stride=(S * H * D, D, H * D, 1), + data_type=fe_dtype, + ) + + self.q = _tensor("q", Q) + self.k = _tensor("k", KV) + self.v = _tensor("v", KV) + self.seq_q = graph.tensor( + name="seq_q", + dim=(B, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=_cudnn_fe.data_type.INT32, + ) + self.seq_kv = graph.tensor( + name="seq_kv", + dim=(B, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=_cudnn_fe.data_type.INT32, + ) + out, _ = graph.sdpa( + name="sdpa", + q=self.q, + k=self.k, + v=self.v, + is_inference=True, + attn_scale=1.0 / math.sqrt(D), + use_padding_mask=True, + seq_len_q=self.seq_q, + seq_len_kv=self.seq_kv, + ) + out.set_output(True).set_dim((B, H, Q, D)).set_stride( + (Q * H * D, D, H * D, 1) + ).set_data_type(fe_dtype) + self.out = out + graph.validate() + graph.build_operation_graph() + graph.create_execution_plans([_cudnn_fe.heur_mode.A]) + graph.check_support() + graph.build_plans() + self.graph = graph + self.workspace_size = max(graph.get_workspace_size(), 1) + + +def _get_graph_or_none( + B: int, + H: int, + Q: int, + KV: int, + D: int, + dtype: torch.dtype, + device: torch.device, +) -> Any: + """Build (or fetch) the execution graph; ``None`` when unsupported. + + A failed build is negatively cached and warned once per shape, so an + unsupported shape degrades to the masked fallback instead of raising + mid-serving. Callers must hold the module lock. + """ + graph_key = (B, H, Q, KV, D, str(dtype), device.index or 0) + cached = _graph_cache.get(graph_key, _UNPROBED) + if cached is not _UNPROBED: + return cached + if device.type == "cuda" and torch.cuda.is_current_stream_capturing(): + # Handle creation and plan building under CUDA graph capture would + # invalidate the capture; a shape first met inside a capture takes + # the masked fallback and is probed on its next eager call. + return None + try: + graph = _Graph(B, H, Q, KV, D, dtype, device) + except Exception as exc: # noqa: BLE001 + _graph_cache[graph_key] = None + warn_once( + key=f"cudnn-varlen-unsupported-{graph_key}", + message=( + f"cuDNN variable-length attention unavailable for shape " + f"B={B} H={H} Q={Q} KV={KV} D={D} ({exc}); using the " + f"masked fallback for this shape." + ), + stacklevel=2, + ) + return None + _graph_cache[graph_key] = graph + return graph + + +def _masked_fallback( + query: Tensor, # [B, Q, H, D] + key: Tensor, # [B, KV, H, D] + value: Tensor, # [B, KV, H, D] + seqused_key_value: Tensor, # [B] +) -> Tensor: + """Boolean-mask attention identical to the SDPA seqused path.""" + kv_index = torch.arange(key.size(1), device=key.device) + mask = kv_index.view(1, 1, 1, -1) < seqused_key_value.view(-1, 1, 1, 1) + out = F.scaled_dot_product_attention( + query=query.transpose(-3, -2), + key=key.transpose(-3, -2), + value=value.transpose(-3, -2), + attn_mask=mask, + ).transpose(-3, -2) + # The op's fake registration promises a contiguous [B, Q, H, D] + # result; a transposed view here would violate that stride contract + # and crash compiled callers on the degrade path. + return out.contiguous() + + +@torch.library.custom_op("sdm::cudnn_varlen_sdpa", mutates_args=()) +def cudnn_varlen_sdpa( + query: Tensor, # [B, Q, H, D] contiguous + key: Tensor, # [B, KV, H, D] contiguous + value: Tensor, # [B, KV, H, D] contiguous + seqused_key_value: Tensor, # [B] int32 +) -> Tensor: # [B, Q, H, D] + """Variable-length SDPA on cuDNN's native padding-mask support. + + Counts are bounded to ``[0, KV]`` before they reach cuDNN, which + binds them as raw sequence lengths: an over-range count saturates + like the boolean mask instead of reading past the valid keys. A count + of 0 follows cuDNN's zero-valid-length semantics. + + Shape notation: ``D`` here is the per-head channel dimension, + written ``C`` in :mod:`sdm.nn.attention`. + """ + assert _cudnn_fe is not None + B, Q, H, D = query.shape + KV = key.size(1) + # The graph derives its geometry from ``query`` (plus ``key``'s + # length) and binds the remaining tensors raw, so a mismatched + # value/key geometry, batch, count dtype/length, or a + # foreign-device count would be reinterpreted instead of rejected. + # :func:`eligible` cannot see ``value`` or the count; route their + # mismatches to the masked fallback, which raises the same errors + # as the boolean-mask path. The graph also declares fixed BSHD + # strides and binds raw storage, so non-contiguous inputs take the + # fallback too. All checks are host-side metadata (no device sync). + if ( + not query.is_contiguous() + or not key.is_contiguous() + or not value.is_contiguous() + or key.dtype != query.dtype + or key.device != query.device + or value.dtype != query.dtype + or value.device != query.device + or value.size() != key.size() + or key.size(0) != B + or key.size(-1) != D + or key.size(-2) != H + or seqused_key_value.device != query.device + or seqused_key_value.dtype != torch.int32 + or seqused_key_value.numel() != B + ): + return _masked_fallback(query, key, value, seqused_key_value) + # The lock serializes handle/stream binding and graph-cache access: + # the module targets single-stream serving; concurrent callers are + # correct but serialized. The support probe lives inside the op + # (opaque to compilation), so unsupported shapes degrade to the + # masked fallback at runtime without touching the traced graph. + with _lock: + graph = _get_graph_or_none(B, H, Q, KV, D, query.dtype, query.device) + if graph is None: + return _masked_fallback(query, key, value, seqused_key_value) + out = torch.empty_like(query) + # Allocated per call: the caching allocator makes this cheap and + # stream-ordered, so callers on distinct streams never share + # device-side scratch (the lock only orders host-side enqueue). + workspace = torch.empty( + graph.workspace_size, device=query.device, dtype=torch.uint8 + ) + # The query-length constant is built per call rather than cached + # with the graph: a cached device tensor allocated during a CUDA + # graph warmup run would live in the graph's private pool. + seq_q = torch.full( + (B, 1, 1, 1), Q, dtype=torch.int32, device=query.device + ) + seq_kv = seqused_key_value.clamp(min=0, max=KV).view(B, 1, 1, 1) + # cuDNN requires the handle's device to be current at call time; + # guard it like the build in :class:`_Graph` already does. + with torch.cuda.device(query.device): + stream = torch.cuda.current_stream(query.device) + _cudnn_fe.set_stream( + handle=graph.handle, stream=stream.cuda_stream + ) + graph.graph.execute( + { + graph.q: query, + graph.k: key, + graph.v: value, + graph.seq_q: seq_q, + graph.seq_kv: seq_kv, + graph.out: out, + }, + workspace, + handle=graph.handle, + ) + return out + + +# The fake implementation keeps compiled callers on ``fullgraph=True``; +# the runtime support probe lives inside the op body (opaque to +# compilation), so one compiled graph stays valid whether a shape +# builds an execution graph or degrades to the masked fallback. Both +# real paths return a contiguous tensor, so the fake promises the same +# layout regardless of the input strides. +@cudnn_varlen_sdpa.register_fake +def _( + query: Tensor, key: Tensor, value: Tensor, seqused_key_value: Tensor +) -> Tensor: + return torch.empty_like(query, memory_format=torch.contiguous_format) diff --git a/sdm/nn/attention.py b/sdm/nn/attention.py index 0d85996b7..a6c186879 100644 --- a/sdm/nn/attention.py +++ b/sdm/nn/attention.py @@ -14,7 +14,7 @@ from sdm._memory import chunk_memory_limit from sdm.cache import KVCacheEntry -from sdm.nn import QueryScaling +from sdm.nn import QueryScaling, _cudnn_varlen class SDPA(torch.nn.Module): @@ -80,7 +80,12 @@ def forward( value: The value tensor with shape ``[..., KV, Hkv, C]``. seqused_key_value: Valid key/value lengths with shape ``[...]`` and :external+torch:ref:`torch.int32 ` dtype. Counts - above ``KV`` act as ``KV``. + above ``KV`` act as ``KV``. With + :func:`~sdm.nn.enable_cudnn_varlen`, eligible calls + (``bfloat16``/``float16``, ``Hq == Hkv``, ``C % 8 == 0``, + ``C <= 128``, at most ``65535`` batch elements ``[...]``, + no gradients, default ``scale``) run on cuDNN's native + padding mask instead of a boolean mask. attn_mask: Boolean attention mask with shape ``[..., Q, KV]``. Entries set to ``True`` participate in attention. @@ -154,7 +159,29 @@ def forward( if seqused_key_value is not None: seqused_key_value = seqused_key_value.expand(batch_shape) - seqused_key_value = seqused_key_value.reshape(-1).unsqueeze(-1) + seqused_key_value = seqused_key_value.reshape(-1) + # The variable-length op bakes in the default ``1 / sqrt(C)`` + # attention scale, so a custom `scale` stays on the + # boolean-mask path rather than being silently ignored. + if self.scale is None and _cudnn_varlen.eligible( + query, + key, + num_query_heads=self.num_query_heads, + num_key_value_heads=self.num_key_value_heads, + ): + # cuDNN's native padding-mask support bounds attention + # to the valid key/value region instead of masking it + # (see enable_cudnn_varlen). The op saturates out-of-range + # counts, so the model contract (degrade, do not raise) + # matches the boolean mask below. + out = _cudnn_varlen.cudnn_varlen_sdpa( + query=query.contiguous(), + key=key.contiguous(), + value=value.contiguous(), + seqused_key_value=seqused_key_value, + ) + return out.view(batch_shape + out.size()[-3:]) + seqused_key_value = seqused_key_value.unsqueeze(-1) key_index = torch.arange(key.size(-3), device=key.device) attn_mask = key_index.unsqueeze(0) < seqused_key_value attn_mask = attn_mask.unsqueeze(-2).expand(-1, query.size(-3), -1) diff --git a/test/models/tabiclv2/test_model.py b/test/models/tabiclv2/test_model.py index c232e5088..d30c17d05 100644 --- a/test/models/tabiclv2/test_model.py +++ b/test/models/tabiclv2/test_model.py @@ -11,7 +11,13 @@ from sdm.models import TabICLv2 from sdm.models.tabiclv2.model import _TabICLv2 from sdm.models.tabiclv2.row_embedding import RowEmbedding -from sdm.nn import Attention, TransformerBlock +from sdm.nn import ( + Attention, + TransformerBlock, + _cudnn_varlen, + cudnn_varlen_stats, + enable_cudnn_varlen, +) from sdm.testing import onlyCUDA, onlyFullTest, withCUDA @@ -1214,3 +1220,115 @@ def test_autocast_compile(device: torch.device) -> None: predicted.numerical, expected_predicted.numerical, ) + + +@pytest.mark.cuda +def test_cudnn_varlen_toggle_and_degrade() -> None: + torch.manual_seed(0) + model = TabICLv2(pretrained=False) + _randomize_residual_exits(model) + + R_context, R_query, C = 11, 6, 5 + x_context = torch.randn(R_context, C) + x_query = torch.randn(R_query, C) + y_context = torch.randint(0, 10, (R_context, 1)) + + # Pad in-context rows exactly as bucketed serving does. Padded targets + # repeat real ones so padding cannot widen the class set (which would + # change the number of output columns). + x_context_padded = torch.cat([x_context, torch.full((5, C), 123.0)]) + y_context_padded = torch.cat([y_context, y_context[:5]]) + seqused_train = torch.tensor(R_context, dtype=torch.int32) + + expected = model( + x_context_padded, + y_context_padded, + x_query, + recipe=Recipe(), + seqused_train=seqused_train, + ) + + # Without the optional dependency (or on CPU) the boolean-mask path + # keeps serving exactly, and disabling always reports inactive. + active = enable_cudnn_varlen(True) + if not _cudnn_varlen.is_available(): + assert not active + try: + out = model( + x_context_padded, + y_context_padded, + x_query, + recipe=Recipe(), + seqused_train=seqused_train, + ) + torch.testing.assert_close( + out.numerical, + expected.numerical, + atol=1e-3, + rtol=1e-3, + ) + finally: + assert enable_cudnn_varlen(False) is False + + +@onlyCUDA +def test_cudnn_varlen_equivalence() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda") + + torch.manual_seed(0) + model = TabICLv2(pretrained=False, device=device).to(torch.bfloat16) + _randomize_residual_exits(model) + + R_context, R_query, C = 96, 32, 8 + x_context = torch.randn(R_context, C, device=device, dtype=torch.bfloat16) + x_query = torch.randn(R_query, C, device=device, dtype=torch.bfloat16) + y_context = torch.randint(0, 10, (R_context, 1), device=device) + + # Pad in-context rows so the call takes the `seqused_key_value` branch + # the variable-length path replaces; padded targets repeat real ones. + x_context_padded = torch.cat( + [ + x_context, + torch.full((32, C), 123.0, device=device, dtype=torch.bfloat16), + ] + ) + y_context_padded = torch.cat([y_context, y_context[:32]]) + seqused_train = torch.tensor(R_context, dtype=torch.int32, device=device) + + expected = model( + x_context_padded, + y_context_padded, + x_query, + recipe=Recipe(), + seqused_train=seqused_train, + ) + # The toggle must actually engage (the dependency is installed on + # this CI job), and at least one execution graph must be built by + # the enabled forward - otherwise a silent permanent degrade to the + # bit-identical masked fallback would keep this equivalence check + # green without ever running the real kernel. + assert enable_cudnn_varlen(True) is True + try: + out = model( + x_context_padded, + y_context_padded, + x_query, + recipe=Recipe(), + seqused_train=seqused_train, + ) + assert cudnn_varlen_stats()[0] >= 1, ( + "cuDNN varlen path never engaged (all shapes degraded)" + ) + finally: + enable_cudnn_varlen(False) + + # The variable-length kernels differ from the masked kernels, so + # allow kernel-switch-scale noise (same class as the padding tests). + torch.testing.assert_close( + out.numerical, + expected.numerical, + atol=1e-2, + rtol=1e-2, + ) diff --git a/test/nn/test_cudnn_varlen.py b/test/nn/test_cudnn_varlen.py new file mode 100644 index 000000000..5ee0b54b0 --- /dev/null +++ b/test/nn/test_cudnn_varlen.py @@ -0,0 +1,631 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import copy +import warnings + +import pytest +import torch +from torch.nn.attention import SDPBackend, sdpa_kernel + +from sdm.nn import SDPA, TransformerBlock, _cudnn_varlen +from sdm.testing import onlyCUDA + +# Unit tests for the sdm.nn._cudnn_varlen gating and custom op (the +# model-mediated toggle and equivalence tests live in +# test/models/tabiclv2/test_model.py). The eligibility and degrade tests +# are CPU-safe by design but additionally marked `cuda`, so the GPU CI +# job - where `nvidia-cudnn-frontend` is installed - also exercises them +# against the real dependency; the remaining tests run the kernel and +# need CUDA plus the optional package. + + +@pytest.mark.cuda +def test_cudnn_varlen_shape_eligibility() -> None: + def shape_ok( + batch: int = 4, + head_dim: int = 64, + num_heads: int = 8, + num_key_value_heads: int = 8, + dtype: torch.dtype = torch.bfloat16, + ) -> bool: + query = torch.empty(batch, 2, num_heads, head_dim, dtype=dtype) + return _cudnn_varlen._shape_eligible( + query, + num_query_heads=num_heads, + num_key_value_heads=num_key_value_heads, + ) + + assert shape_ok() + assert shape_ok(dtype=torch.float16) + # Full precision stays on the boolean-mask path. + assert not shape_ok(dtype=torch.float32) + # Grouped-query attention stays on the boolean-mask path. + assert not shape_ok(num_key_value_heads=2) + # Head dims must be multiples of eight... + assert not shape_ok(head_dim=36) + # ...with 128 the inclusive ceiling (136 is a multiple of eight, + # isolating the ceiling gate). + assert shape_ok(head_dim=128) + assert not shape_ok(head_dim=136) + # Batch 65535 is the inclusive ceiling. + assert shape_ok( + batch=65535, + head_dim=8, + num_heads=1, + num_key_value_heads=1, + ) + assert not shape_ok( + batch=65536, + head_dim=8, + num_heads=1, + num_key_value_heads=1, + ) + + +@pytest.mark.cuda +def test_cudnn_varlen_eligibility_gates() -> None: + query = torch.empty(4, 32, 8, 64) + key = torch.empty(4, 48, 8, 64) + + # A disabled flag short-circuits everything. + _cudnn_varlen.enable_cudnn_varlen(False) + assert not _cudnn_varlen.eligible( + query, key, num_query_heads=8, num_key_value_heads=8 + ) + + _cudnn_varlen.enable_cudnn_varlen(True) + try: + # CPU tensors are never eligible (also covers environments + # without the optional dependency, where enabling is inert). + assert not _cudnn_varlen.eligible( + query, key, num_query_heads=8, num_key_value_heads=8 + ) + if torch.cuda.is_available() and _cudnn_varlen.is_available(): + base = query.cuda().bfloat16() + base_k = key.cuda().bfloat16() + # Eligibility mirrors serving: inference (no-grad) context. + with torch.no_grad(): + assert _cudnn_varlen.eligible( + base, base_k, num_query_heads=8, num_key_value_heads=8 + ) + # The graph takes dtype and device from ``query``, so a + # key that differs in either must not be bound to it. + assert not _cudnn_varlen.eligible( + base, + base_k.float(), + num_query_heads=8, + num_key_value_heads=8, + ) + assert not _cudnn_varlen.eligible( + base, + base_k.cpu(), + num_query_heads=8, + num_key_value_heads=8, + ) + # Grouped-query attention stays on the boolean-mask path. + assert not _cudnn_varlen.eligible( + base, base_k, num_query_heads=8, num_key_value_heads=2 + ) + # Head dims must be multiples of eight (and at most 128). + assert not _cudnn_varlen.eligible( + base[..., :36], + base_k[..., :36], + num_query_heads=8, + num_key_value_heads=8, + ) + # Head dim 128 is the inclusive ceiling (136 is a + # multiple of eight, isolating the ceiling gate). + wide = torch.empty( + 2, 4, 8, 136, device="cuda", dtype=torch.bfloat16 + ) + assert _cudnn_varlen.eligible( + wide[..., :128], + wide[..., :128], + num_query_heads=8, + num_key_value_heads=8, + ) + assert not _cudnn_varlen.eligible( + wide, + wide, + num_query_heads=8, + num_key_value_heads=8, + ) + # Batch 65535 is the inclusive ceiling. + flat = torch.empty( + 65536, 1, 1, 8, device="cuda", dtype=torch.bfloat16 + ) + assert _cudnn_varlen.eligible( + flat[:65535], + flat[:65535], + num_query_heads=1, + num_key_value_heads=1, + ) + assert not _cudnn_varlen.eligible( + flat, + flat, + num_query_heads=1, + num_key_value_heads=1, + ) + # Full precision stays on the boolean-mask path. + assert not _cudnn_varlen.eligible( + base.float(), + base_k.float(), + num_query_heads=8, + num_key_value_heads=8, + ) + # Gradient-enabled calls stay on the boolean-mask path. + with torch.enable_grad(): + assert not _cudnn_varlen.eligible( + base.clone().requires_grad_(True), + base_k, + num_query_heads=8, + num_key_value_heads=8, + ) + finally: + _cudnn_varlen.enable_cudnn_varlen(False) + + +@pytest.mark.cuda +def test_cudnn_varlen_build_failure_degrades( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # A graph-build failure must degrade to the masked fallback inside + # the op (probed once, negatively cached), never raise mid-serving. + # Monkeypatching `_Graph` itself (rather than the cuDNN frontend) + # keeps the simulated failure reachable on every machine: the raise + # happens inside `_get_graph_or_none`'s probe on CPU and CUDA alike, + # and the warning match below pins the mock's message so the test + # fails if the simulated failure ever stops being exercised. + class _FailingGraph: + def __init__(self, *args: object, **kwargs: object) -> None: + raise RuntimeError("No execution plans support the graph.") + + monkeypatch.setattr(_cudnn_varlen, "_cudnn_fe", object()) + monkeypatch.setattr(_cudnn_varlen, "_enabled", True) + monkeypatch.setattr(_cudnn_varlen, "_Graph", _FailingGraph) + # Start from a clean cache so the public stats assertion below sees + # exactly this test's probe. + _cudnn_varlen._graph_cache.clear() + + generator = torch.Generator().manual_seed(0) + query = torch.randn(2, 16, 8, 64, generator=generator) + key = torch.randn(2, 32, 8, 64, generator=generator) + value = torch.randn(2, 32, 8, 64, generator=generator) + seqused = torch.tensor([20, 32], dtype=torch.int32) + try: + # Pin the math backend: the fused CPU kernels happen to return + # contiguous outputs even without the fallback's stride fix, + # which would make the contract assertions below vacuous. + with sdpa_kernel([SDPBackend.MATH]): + with pytest.warns(RuntimeWarning, match="No execution plans"): + out = _cudnn_varlen.cudnn_varlen_sdpa( + query=query, + key=key, + value=value, + seqused_key_value=seqused, + ) + expected = _cudnn_varlen._masked_fallback( + query, key, value, seqused + ) + torch.testing.assert_close(out, expected) + # Negatively cached: the second call neither warns nor rebuilds. + with warnings.catch_warnings(record=True) as record: + warnings.simplefilter("always") + _cudnn_varlen.cudnn_varlen_sdpa(query, key, value, seqused) + assert not record + # The public stats accessor reports the negatively cached probe + # as a degrade (0 built, 1 degraded). + assert _cudnn_varlen.cudnn_varlen_stats() == (0, 1) + # The fallback must honor the op's fake stride contract: a + # compiled caller receives the fake's layout and crashes if the + # real output is a transposed view. + assert out.is_contiguous() + fake = torch.empty_like(query) + assert out.stride() == fake.stride() + finally: + # The monkeypatch restores the module globals; drop the + # negatively cached probe so later tests see a clean cache. + _cudnn_varlen._graph_cache.clear() + + +@onlyCUDA +@pytest.mark.parametrize("mismatch", ["dtype", "device"]) +def test_cudnn_varlen_rejects_mismatched_key(mismatch: str) -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + query = torch.ones(2, 3, 4, 32, device="cuda", dtype=torch.bfloat16) + value = torch.ones(2, 5, 4, 32, device="cuda", dtype=torch.bfloat16) + key = value.float() if mismatch == "dtype" else value.cpu() + counts = torch.tensor([5, 3], dtype=torch.int32, device="cuda") + _cudnn_varlen.enable_cudnn_varlen(False) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + # Direct callers bypass eligible(): an incompatible key must never + # be bound as a raw pointer to a graph built for query's metadata. + with torch.no_grad(), pytest.raises(RuntimeError): + _cudnn_varlen.cudnn_varlen_sdpa( + query=query, + key=key, + value=value, + seqused_key_value=counts, + ) + assert _cudnn_varlen.cudnn_varlen_stats() == (0, 0) + finally: + _cudnn_varlen.enable_cudnn_varlen(False) + + +@onlyCUDA +def test_cudnn_varlen_disable_during_inflight_execution() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(0) + query = torch.randn( + 4, 64, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + key = torch.randn( + 4, 130, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + value = torch.randn( + 4, 130, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + counts = torch.tensor([130, 77, 1, 100], dtype=torch.int32, device=device) + # Disabling clears the graph cache, so the stats below see only this + # test's graph. + _cudnn_varlen.enable_cudnn_varlen(False) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + # Build the graph on the default stream, then execute on a second + # stream and drop the graph while those executions may still be + # in flight. + with torch.no_grad(): + expected = _cudnn_varlen.cudnn_varlen_sdpa( + query=query, + key=key, + value=value, + seqused_key_value=counts, + ) + torch.cuda.synchronize() + # The comparison below is only meaningful on the cuDNN graph, + # not on the masked fallback a failed build would degrade to. + assert _cudnn_varlen.cudnn_varlen_stats() == (1, 0) + stream = torch.cuda.Stream(device=device) + with torch.cuda.stream(stream): + outs = [ + _cudnn_varlen.cudnn_varlen_sdpa(query, key, value, counts) + for _ in range(8) + ] + _cudnn_varlen.enable_cudnn_varlen(False) + # Churn the allocator so memory released by the drop would be + # reused before the queued kernels have finished. + junk = [ + torch.full((4, 1, 1, 1), 7, dtype=torch.int32, device=device) + for _ in range(64) + ] + stream.synchronize() + for out in outs: + assert torch.equal(out, expected) + del junk + finally: + _cudnn_varlen.enable_cudnn_varlen(False) + + +@onlyCUDA +def test_cudnn_varlen_non_contiguous_inputs_take_fallback() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(1) + # [B, H, Q, D] storage viewed as [B, Q, H, D]: the layout the graph + # declares, but with the strides it does not. + query = torch.randn( + 4, 8, 64, 64, device=device, dtype=torch.bfloat16, generator=generator + ).transpose(1, 2) + key = torch.randn( + 4, 8, 130, 64, device=device, dtype=torch.bfloat16, generator=generator + ).transpose(1, 2) + value = torch.randn( + 4, 8, 130, 64, device=device, dtype=torch.bfloat16, generator=generator + ).transpose(1, 2) + assert not query.is_contiguous() + counts = torch.tensor([130, 77, 1, 100], dtype=torch.int32, device=device) + _cudnn_varlen.enable_cudnn_varlen(False) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + with torch.no_grad(): + out = _cudnn_varlen.cudnn_varlen_sdpa(query, key, value, counts) + expected = _cudnn_varlen._masked_fallback( + query.contiguous(), + key.contiguous(), + value.contiguous(), + counts, + ) + torch.testing.assert_close(out, expected) + # Routed to the fallback before any graph was built. + assert _cudnn_varlen.cudnn_varlen_stats() == (0, 0) + finally: + _cudnn_varlen.enable_cudnn_varlen(False) + + +@onlyCUDA +def test_cudnn_varlen_compiled_non_contiguous_query() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(2) + query = torch.randn( + 4, 8, 64, 64, device=device, dtype=torch.bfloat16, generator=generator + ).transpose(1, 2) + key = torch.randn( + 4, 130, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + value = torch.randn( + 4, 130, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + counts = torch.tensor([130, 77, 1, 100], dtype=torch.int32, device=device) + _cudnn_varlen.enable_cudnn_varlen(False) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + # The fake must describe the contiguous output the fallback + # returns for a non-contiguous query, or the compiled graph would + # be laid out for the query's strides instead. + compiled = torch.compile( + _cudnn_varlen.cudnn_varlen_sdpa, + fullgraph=True, + backend="aot_eager", + ) + with torch.no_grad(): + out = compiled(query, key, value, counts) + expected = _cudnn_varlen.cudnn_varlen_sdpa( + query=query, + key=key, + value=value, + seqused_key_value=counts, + ) + assert out.is_contiguous() + torch.testing.assert_close(out, expected) + finally: + torch._dynamo.reset() + _cudnn_varlen.enable_cudnn_varlen(False) + + +@onlyCUDA +def test_cudnn_varlen_over_range_count_saturates() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(3) + module = SDPA(num_query_heads=8) + query = torch.randn( + 3, 4, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + key = torch.randn( + 3, 33, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + value = torch.randn( + 3, 33, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + exact = torch.full((3,), 33, dtype=torch.int32, device=device) + _cudnn_varlen.enable_cudnn_varlen(False) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + with torch.no_grad(): + expected = module( + query=query, key=key, value=value, seqused_key_value=exact + ) + # cuDNN binds ``seq_len_kv`` as a raw length: a count past the + # key length reads beyond the valid keys and returns garbage or + # NaN instead of saturating like the boolean mask does. The op + # bounds the count, so over-range counts must reproduce the + # exact-count output bit for bit, through `SDPA` and directly. + for count in (34, 40, 1000, 2**31 - 1): + over = exact.new_full((3,), count) + out = module( + query=query, key=key, value=value, seqused_key_value=over + ) + assert torch.equal(out, expected), count + out = _cudnn_varlen.cudnn_varlen_sdpa(query, key, value, over) + assert torch.equal(out, expected), count + # The comparison is only meaningful on the cuDNN graph, not on the + # masked fallback a failed build would degrade to. + assert _cudnn_varlen.cudnn_varlen_stats() == (1, 0) + finally: + _cudnn_varlen.enable_cudnn_varlen(False) + + +@onlyCUDA +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_cudnn_varlen_reduce_overhead_compile(dtype: torch.dtype) -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(4) + channels = 64 + block = TransformerBlock( + channels=channels, + num_query_heads=4, + mlp=torch.nn.Linear(channels, channels, device=device, dtype=dtype), + device=device, + dtype=dtype, + ) + with torch.no_grad(): + # The residual exit is zero-initialized; randomize it so attention + # (and thus the masking) reaches the output. + block.attn.out_lin.weight.normal_(std=0.5, generator=generator) + eager = copy.deepcopy(block) + + def inputs(kv_len: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + query = torch.randn( + 4, 32, channels, device=device, dtype=dtype, generator=generator + ) + key_value = torch.randn( + 4, + kv_len, + channels, + device=device, + dtype=dtype, + generator=generator, + ) + counts = torch.tensor( + [kv_len, kv_len // 2, 1, kv_len - 3], + dtype=torch.int32, + device=device, + ) + return query, key_value, counts + + # Three calls at one shape walk CUDA graph trees through warmup, + # recording, and replay; the second shape arrives after warmup. + plan = [inputs(kv_len) for kv_len in (48, 48, 48, 80, 80)] + _cudnn_varlen.enable_cudnn_varlen(False) + with torch.inference_mode(): + # Boolean-mask references, taken before the path is enabled. + expected = [ + eager(query, key_value, seqused_key_value=counts) + for query, key_value, counts in plan + ] + _cudnn_varlen.enable_cudnn_varlen(True) + try: + block.compile(fullgraph=True, dynamic=True, mode="reduce-overhead") + with ( + torch.inference_mode(), + warnings.catch_warnings(record=True) as record, + ): + warnings.simplefilter("always") + for (query, key_value, counts), reference in zip( + plan, expected, strict=True + ): + torch.compiler.cudagraph_mark_step_begin() + out = block(query, key_value, seqused_key_value=counts) + torch.testing.assert_close( + out.clone(), reference, atol=1e-2, rtol=1e-2 + ) + # The first shape ran on a cuDNN graph built during warmup; no + # shape degraded to the fallback or was negatively cached, whether + # the second shape was warmed up or first met while recording. + built, degraded = _cudnn_varlen.cudnn_varlen_stats() + assert built >= 1 + assert degraded == 0 + assert not [w for w in record if "variable-length" in str(w.message)] + finally: + torch._dynamo.reset() + _cudnn_varlen.enable_cudnn_varlen(False) + + +@pytest.mark.cuda +def test_cudnn_varlen_masked_fallback_matches_sdpa() -> None: + # The degrade tests above assert the op against `_masked_fallback` + # itself, so a masking bug in the fallback would pass them. Pin the + # fallback to independent references: the boolean-mask path of + # `SDPA` with the same counts, and a per-batch slice of the valid + # keys. Counts sit at both ends of the range so an off-by-one in + # either direction is visible. + generator = torch.Generator().manual_seed(5) + query = torch.randn(4, 3, 2, 8, generator=generator) + key = torch.randn(4, 9, 2, 8, generator=generator) + value = torch.randn(4, 9, 2, 8, generator=generator) + counts = torch.tensor([9, 1, 8, 5], dtype=torch.int32) + module = SDPA(num_query_heads=2) + _cudnn_varlen.enable_cudnn_varlen(False) + with sdpa_kernel([SDPBackend.MATH]): + out = _cudnn_varlen._masked_fallback(query, key, value, counts) + expected = module( + query=query, key=key, value=value, seqused_key_value=counts + ) + torch.testing.assert_close(out, expected) + for i, count in enumerate(counts.tolist()): + sliced = module( + query=query[i], key=key[i, :count], value=value[i, :count] + ) + torch.testing.assert_close(out[i], sliced) + + +@onlyCUDA +def test_cudnn_varlen_custom_scale_stays_on_mask_path() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(6) + query = torch.randn( + 3, 4, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + key = torch.randn( + 3, 33, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + value = torch.randn( + 3, 33, 8, 64, device=device, dtype=torch.bfloat16, generator=generator + ) + counts = torch.tensor([33, 7, 20], dtype=torch.int32, device=device) + # The op bakes in the default `1 / sqrt(C)` scale; a module with a + # custom scale must stay on the boolean-mask path instead of silently + # attending with the wrong temperature. + module = SDPA(num_query_heads=8, scale=1.0) + _cudnn_varlen.enable_cudnn_varlen(False) + with torch.no_grad(): + expected = module( + query=query, key=key, value=value, seqused_key_value=counts + ) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + with torch.no_grad(): + out = module( + query=query, key=key, value=value, seqused_key_value=counts + ) + assert _cudnn_varlen.cudnn_varlen_stats() == (0, 0) + assert torch.equal(out, expected) + finally: + _cudnn_varlen.enable_cudnn_varlen(False) + + +@onlyCUDA +def test_cudnn_varlen_unprobed_shape_inside_capture_takes_fallback() -> None: + if not _cudnn_varlen.is_available(): + pytest.skip("requires nvidia-cudnn-frontend") + device = torch.device("cuda", 0) + generator = torch.Generator(device=device).manual_seed(7) + query = torch.randn( + 2, 3, 4, 32, device=device, dtype=torch.bfloat16, generator=generator + ) + key = torch.randn( + 2, 17, 4, 32, device=device, dtype=torch.bfloat16, generator=generator + ) + value = torch.randn( + 2, 17, 4, 32, device=device, dtype=torch.bfloat16, generator=generator + ) + counts = torch.tensor([17, 5], dtype=torch.int32, device=device) + # Disabling clears the graph cache, so this shape is unprobed. + _cudnn_varlen.enable_cudnn_varlen(False) + _cudnn_varlen.enable_cudnn_varlen(True) + try: + # Building a graph (handle creation, plan building, the build + # sync) under stream capture would invalidate the capture: a shape + # first met while capturing takes the masked fallback without + # touching the cache, and builds on its next eager call. + graph = torch.cuda.CUDAGraph() + stream = torch.cuda.Stream(device=device) + stream.wait_stream(torch.cuda.current_stream(device)) + with torch.no_grad(), torch.cuda.graph(graph, stream=stream): + out = _cudnn_varlen.cudnn_varlen_sdpa(query, key, value, counts) + torch.cuda.synchronize() + assert _cudnn_varlen.cudnn_varlen_stats() == (0, 0) + assert not _cudnn_varlen._graph_cache + # The replay is the captured fallback on the new data and counts. + query.copy_( + torch.randn( + query.shape, + device=device, + dtype=query.dtype, + generator=generator, + ) + ) + counts.copy_(torch.tensor([9, 1], dtype=torch.int32, device=device)) + graph.replay() + torch.cuda.synchronize() + expected = _cudnn_varlen._masked_fallback(query, key, value, counts) + assert torch.equal(out, expected) + with torch.no_grad(): + eager = _cudnn_varlen.cudnn_varlen_sdpa(query, key, value, counts) + assert _cudnn_varlen.cudnn_varlen_stats() == (1, 0) + torch.testing.assert_close(eager, expected, atol=1e-2, rtol=1e-2) + finally: + _cudnn_varlen.enable_cudnn_varlen(False) diff --git a/uv.lock b/uv.lock index 068c1b964..bbb3d079e 100644 --- a/uv.lock +++ b/uv.lock @@ -1095,6 +1095,23 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/5c/ba/791cffd048fe5b044e620df55267e3e95c0e6e07d50b41e377c03dfc910f/nvidia_cudnn_cu13-9.24.0.43-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:71f181cd810e90f9b6023b01186fe82d13d65f0ec098581ee201d39fad769e4b", size = 553099438, upload-time = "2026-07-02T16:27:42.58Z" }, ] +[[package]] +name = "nvidia-cudnn-frontend" +version = "1.28.0" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b2/55/880914d0c16bedce1c5ed185ce0847eee8c6fe6c5104a2698ace8ce93a2b/nvidia_cudnn_frontend-1.28.0-cp311-cp311-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5173437c88b3accb7cc6855b42a64d583d7962ce5a90b148257e0a0d3c9c1eaa", size = 5509420, upload-time = "2026-09-02T06:06:00.48Z" }, + { url = "https://files.pythonhosted.org/packages/25/9e/f65832a9bffec31ac9f55edf18c25c3c8a396ddafd89c108fe21389552c4/nvidia_cudnn_frontend-1.28.0-cp311-cp311-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:248e2d81fb670e7b591246eeb626fc336e2bad29551afa84daf062d63aaf76b3", size = 5671671, upload-time = "2026-09-02T06:06:19.418Z" }, + { url = "https://files.pythonhosted.org/packages/d1/9e/33b746b800c36a8aae8432605c5b2af6cc5fc683cc1a6954a084c3853690/nvidia_cudnn_frontend-1.28.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:060b0c021f6841312ad26dd837a51b4941d4e0d2a547d87455bd974c173da075", size = 5512380, upload-time = "2026-09-02T06:07:11.74Z" }, + { url = "https://files.pythonhosted.org/packages/fd/0d/b6670b5d2d193322e0765d01889998df11d7f57522f210dc808b456530d9/nvidia_cudnn_frontend-1.28.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:335c916be758a8d5533104e04fc4a15f9c30a419b103ecd2b6fd0d23e1d532e4", size = 5677342, upload-time = "2026-09-02T06:07:31.949Z" }, + { url = "https://files.pythonhosted.org/packages/1b/37/ce841fc013bc01de87d33eba58ba386477e411f31e3e41337d3f745ea1f3/nvidia_cudnn_frontend-1.28.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:30e20178e46a5a5049c2358f4147d1c5414b45b2753a455da05bab914807a44d", size = 5512441, upload-time = "2026-09-02T06:08:14.921Z" }, + { url = "https://files.pythonhosted.org/packages/3d/1a/51fa7f3d4fd8377d7acdf99e92d33a38be83979897360ca36b6f3521e92d/nvidia_cudnn_frontend-1.28.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:26cd299d202e832a1d562510d6844dc7c6f256cb1b4f0f29d5b2fa0b34f05e98", size = 5677523, upload-time = "2026-09-02T06:08:37.287Z" }, + { url = "https://files.pythonhosted.org/packages/79/21/dc09de787243c3333d6b958eca44697bfd56951465182373d37d88dec0e8/nvidia_cudnn_frontend-1.28.0-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:851fde83b08026f33fc83b2c656232833e4c79ad886c937ed5a035473ad02e1d", size = 5514494, upload-time = "2026-09-02T06:09:23.829Z" }, + { url = "https://files.pythonhosted.org/packages/e5/31/f6e3c1983cae0da1759c7bb6a57f4b2208b042a24b7ccaad775dd14601e0/nvidia_cudnn_frontend-1.28.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:56c596054d7ff929ba95ae560bcf1f41d4f0113a9f8b693401c0d508262a4717", size = 5679104, upload-time = "2026-09-02T06:09:44.703Z" }, + { url = "https://files.pythonhosted.org/packages/32/4a/73a453003eb527565b9a5b0b117ec04a75b4eddd0a0796027064d87b77f2/nvidia_cudnn_frontend-1.28.0-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:92d52770600d9b175e3098faed570477cf31c78646dc224cd92b2149d62afe0e", size = 5517650, upload-time = "2026-09-02T06:10:30.932Z" }, + { url = "https://files.pythonhosted.org/packages/a5/29/742fd93110343da4451eb4d1a713f6291bbb4e80e7b91fae91c826cc150c/nvidia_cudnn_frontend-1.28.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:04b209c40cf5295ea6141be8c74c2e65449b619a0911ceb2ae27cd914124d376", size = 5680701, upload-time = "2026-09-02T06:10:55.184Z" }, +] + [[package]] name = "nvidia-cufft" version = "12.0.0.61" @@ -2122,6 +2139,11 @@ dependencies = [ { name = "typing-extensions" }, ] +[package.optional-dependencies] +cudnn = [ + { name = "nvidia-cudnn-frontend", marker = "sys_platform == 'linux'" }, +] + [package.dev-dependencies] dev = [ { name = "cudf-cu13", marker = "sys_platform == 'linux'" }, @@ -2159,11 +2181,13 @@ test = [ requires-dist = [ { name = "huggingface-hub" }, { name = "numpy" }, + { name = "nvidia-cudnn-frontend", marker = "sys_platform == 'linux' and extra == 'cudnn'", specifier = ">=1.23,<1.29" }, { name = "pyarrow" }, { name = "safetensors" }, { name = "torch", specifier = ">=2.7" }, { name = "typing-extensions" }, ] +provides-extras = ["cudnn"] [package.metadata.requires-dev] dev = [