From 1e4052eb826cbbe03774ce8fd40c942b33a1684e Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 26 Aug 2026 21:20:31 +0000 Subject: [PATCH 01/36] Port JAX fused attention to cuDNN Python frontend Construct and plan standard fused-attention forward and backward graphs through the cuDNN frontend Python API. Share graph serialization and binding infrastructure with flex attention, including packed and ragged tensor bindings. Replace the JAX attention-specific TE common bridge with a graph-agnostic XLA FFI executor and a JAX-local capture-safe RNG state helper. Keep TE common unchanged for other frameworks. Signed-off-by: Vladimir Cherepanov --- build_tools/jax.py | 9 +- .../jax/cpp_extensions/attention.py | 330 +--- .../jax/cpp_extensions/cudnn_attention.py | 1323 +++++++++++++++++ .../jax/cpp_extensions/cudnn_graph.py | 259 ++++ .../jax/cpp_extensions/flex_attention.py | 220 +-- transformer_engine/jax/csrc/extensions.h | 24 +- .../jax/csrc/extensions/attention.cpp | 1029 +++---------- .../jax/csrc/extensions/attention_kernels.cu | 29 + .../jax/csrc/extensions/pybind.cpp | 3 - 9 files changed, 1987 insertions(+), 1239 deletions(-) create mode 100644 transformer_engine/jax/cpp_extensions/cudnn_attention.py create mode 100644 transformer_engine/jax/cpp_extensions/cudnn_graph.py create mode 100644 transformer_engine/jax/csrc/extensions/attention_kernels.cu diff --git a/build_tools/jax.py b/build_tools/jax.py index d61c26c128b..44297200ef3 100644 --- a/build_tools/jax.py +++ b/build_tools/jax.py @@ -6,19 +6,19 @@ import os from pathlib import Path -from packaging import version +from typing import List import setuptools +from packaging import version from .utils import ( - get_cuda_include_dirs, all_files_in_dir, cudnn_frontend_include_path, debug_build_enabled, - setup_mpi_flags, + get_cuda_include_dirs, nccl_ep_enabled, + setup_mpi_flags, ) -from typing import List def install_requirements() -> List[str]: @@ -87,6 +87,7 @@ def setup_jax_extension( csrc_source_files = Path(csrc_source_files) extensions_dir = csrc_source_files / "extensions" sources = all_files_in_dir(extensions_dir, name_extension="cpp") + sources += all_files_in_dir(extensions_dir, name_extension="cu") # Header files include_dirs = get_cuda_include_dirs() diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index 6ea54195ab3..bb9c82500ae 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -2,53 +2,51 @@ # # See LICENSE for license information. """JAX/TE custom ops for attention""" -import operator import os import warnings from dataclasses import dataclass, replace -from functools import partial, reduce +from functools import partial from typing import Optional, Tuple import jax import jax.numpy as jnp -from jax import dtypes, lax, ffi -from jax.sharding import PartitionSpec, NamedSharding +from jax import dtypes, ffi, lax from jax.experimental.custom_partitioning import SdyShardingRule - -import transformer_engine_jax +from jax.sharding import NamedSharding, PartitionSpec from transformer_engine_jax import NVTE_Fused_Attn_Backend + from transformer_engine.jax.attention import ( AttnBiasType, AttnMaskType, AttnSoftmaxType, - QKVLayout, - QKVFormat, CPStrategy, + QKVFormat, + QKVLayout, SequenceDescriptor, ) -from ..sharding import with_sharding_constraint_by_logical_axes, HEAD_AXES, is_mesh_available -from .base import BasePrimitive, register_primitive -from .misc import ( - check_valid_batch_dims, - jax_dtype_to_te_dtype, - te_dtype_to_jax_dtype, - get_padded_spec, - get_cudnn_version, - get_all_device_compute_capability, -) from ..sharding import ( - global_mesh_resource, - lax_paral_op, + HEAD_AXES, all_reduce_sum_along_dp_fsdp, - get_mesh_axis_size, + get_all_mesh_axes, get_mesh_axis_rank, get_mesh_axis_rank_host, - get_all_mesh_axes, + get_mesh_axis_size, + global_mesh_resource, + is_mesh_available, + lax_paral_op, num_of_devices, with_sharding_constraint, + with_sharding_constraint_by_logical_axes, +) +from .base import BasePrimitive, register_primitive +from .cudnn_attention import build_bwd_graph, build_fwd_graph, is_fused_attn_supported +from .misc import ( + check_valid_batch_dims, + get_all_device_compute_capability, + get_cudnn_version, + get_padded_spec, ) - __all__ = [ "FusedAttnHelper", @@ -132,25 +130,10 @@ def is_fused_attn_kernel_available(self): def get_fused_attn_backend(self): """Get the fused attention kernel backend""" - return transformer_engine_jax.get_fused_attn_backend( - self.is_training, - jax_dtype_to_te_dtype(self.q_dtype), - jax_dtype_to_te_dtype(self.kv_dtype), - self.qkv_layout.value, - self.attn_bias_type.value, - self.attn_mask_type.value, - self.softmax_type.value, - self.dropout_probability, - self.q_num_heads, - self.kv_num_heads, - self.q_max_seqlen, - self.kv_max_seqlen, - self.head_dim_qk, - self.head_dim_v, - self.window_size[0], - self.window_size[1], - self.return_max_logit, - not self.is_non_deterministic_allowed(), + return ( + NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen + if is_fused_attn_supported(self) + else NVTE_Fused_Attn_Backend.NVTE_No_Backend ) @staticmethod @@ -271,15 +254,6 @@ def check_seed(self, seed, dropout_probability, is_training): return seed -def generate_cu_seqlen(actual_seqlen): - """ - Generating cumsum seqlen for a batch - """ - actual_seqlen = jnp.where(actual_seqlen < 0, 0, actual_seqlen) - cu_seqlen = jnp.cumulative_sum(actual_seqlen, include_initial=True) - return cu_seqlen - - class FusedAttnFwdPrimitive(BasePrimitive): """ Fused Attention Forward Primitive @@ -338,7 +312,6 @@ def abstract( output_shape = (*batch_shape, q_max_seqlen, attn_heads, v_head_dim) out_aval = q_aval.update(shape=output_shape, dtype=q_dtype) - # backend determines the softmax buffer shape/dtype backend = FusedAttnHelper( config.is_training, q_dtype, @@ -357,35 +330,13 @@ def abstract( config.window_size, config.return_max_logit, ).get_fused_attn_backend() - - if backend == NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: - # cuDNN 9.6 reduces the required softmax shape - if get_cudnn_version() >= (9, 6, 0): - if config.qkv_layout.is_thd(): - softmax_shape = (*batch_shape, q_max_seqlen, attn_heads, 1) - else: - softmax_shape = (*batch_shape, attn_heads, q_max_seqlen, 1) - else: - softmax_shape = ( - *batch_shape, - attn_heads, - q_max_seqlen, - config.max_segments_per_seq, - ) - softmax_dtype = dtypes.canonicalize_dtype(jnp.float32) - else: + if backend != NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: raise ValueError(f"Unsupported {backend=}") - softmax_aux_aval = q_aval.update(shape=softmax_shape, dtype=softmax_dtype) - if config.return_max_logit: - # cuDNN Max is row-wise over S_kv. Dense and SM120 THD use - # [..., H, S_q, 1]; cuDNN >= 9.6 non-SM120 THD uses [..., S_q, H, 1]. - # Both raw layouts are reduced to the public per-head [H] result below. - if FusedAttnFwdPrimitive._uses_thd_ragged_max_tensor(config): - max_tensor_shape = (*batch_shape, q_max_seqlen, attn_heads, 1) - else: - max_tensor_shape = (*batch_shape, attn_heads, q_max_seqlen, 1) - else: - max_tensor_shape = (0,) + + graph_info = build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) + softmax_dtype = dtypes.canonicalize_dtype(jnp.float32) + softmax_aux_aval = q_aval.update(shape=graph_info.stats_shape, dtype=softmax_dtype) + max_tensor_shape = graph_info.max_shape max_tensor_aval = q_aval.update(shape=max_tensor_shape, dtype=softmax_dtype) # JAX does not enable 64-bit int by default so we get XLA to allocate x8 memory with @@ -398,46 +349,8 @@ def abstract( rng_state_shape = (seed_aval.shape[0], checker.rng_state_size) rng_state_aval = seed_aval.update(shape=rng_state_shape, dtype=checker.rng_state_dtype) - if config.attn_bias_type == AttnBiasType.NO_BIAS: - bias_batch = bias_heads = 0 - else: - *bias_batch_shape, bias_heads, _, _ = bias_aval.shape - bias_batch = reduce(operator.mul, bias_batch_shape) - - bottom_right_diagonal = config.attn_mask_type in [ - AttnMaskType.CAUSAL_BOTTOM_RIGHT_MASK, - AttnMaskType.PADDING_CAUSAL_BOTTOM_RIGHT_MASK, - ] - - # do a dummy kernel call here to get workspace buffer shapes/dtypes that XLA needs to - # prepare for the active fused-attn backend - input_batch = reduce(operator.mul, batch_shape) - wkspace_info = transformer_engine_jax.get_fused_attn_fwd_workspace_sizes( - input_batch, - bias_batch, - q_max_seqlen, - kv_max_seqlen, - attn_heads, - num_gqa_groups, - bias_heads, - q_head_dim, - v_head_dim, - config.scaling_factor, - config.dropout_probability, - config.attn_bias_type.value, - config.attn_mask_type.value, - config.softmax_type.value, - config.qkv_layout.value, - jax_dtype_to_te_dtype(q_aval.dtype), - config.is_training, - config.max_segments_per_seq, - config.window_size[0], - config.window_size[1], - config.return_max_logit, - bottom_right_diagonal, - ) wkspace_aval = q_aval.update( - shape=wkspace_info[0], dtype=te_dtype_to_jax_dtype(wkspace_info[1]) + shape=(graph_info.graph.workspace_size,), dtype=jnp.uint8 ) assert ( @@ -493,30 +406,7 @@ def lowering( """ q_aval, k_aval, v_aval, bias_aval, *_ = ctx.avals_in - ( - batch_shape, - q_max_seqlen, - kv_max_seqlen, - attn_heads, - num_gqa_groups, - q_head_dim, - v_head_dim, - ) = FusedAttnHelper.parse_qkv_aval(q_aval, k_aval, v_aval, config.qkv_layout) - - input_batch = reduce(operator.mul, batch_shape) - - if config.attn_bias_type == AttnBiasType.NO_BIAS: - bias_batch = bias_heads = 0 - else: - *bias_batch_shape, bias_heads, _, _ = bias_aval.shape - bias_batch = reduce(operator.mul, bias_batch_shape) - - if config.cp_striped_window_size is not None: - window_size_left = config.cp_striped_window_size[0] - window_size_right = config.cp_striped_window_size[1] - else: - window_size_left = config.window_size[0] - window_size_right = config.window_size[1] + graph = build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config).graph return ffi.ffi_lowering(FusedAttnFwdPrimitive.name)( ctx, @@ -534,28 +424,9 @@ def lowering( _kv_segment_ids, _q_segment_pos, _kv_segment_pos, # ffi_lowering needs number of parameters meets primitive.lowering - input_batch=input_batch, - bias_batch=bias_batch, - q_max_seqlen=q_max_seqlen, - kv_max_seqlen=kv_max_seqlen, - attn_heads=attn_heads, - num_gqa_groups=num_gqa_groups, - bias_heads=bias_heads, - qk_head_dim=q_head_dim, - v_head_dim=v_head_dim, - max_segments_per_seq=config.max_segments_per_seq, - scaling_factor=float(config.scaling_factor), - dropout_probability=float(config.dropout_probability), - bias_type=int(config.attn_bias_type.value), - mask_type=int(config.attn_mask_type.value), - qkv_layout=int(config.qkv_layout.value), - is_training=config.is_training, - return_max_logit=config.return_max_logit, - deterministic=not FusedAttnHelper.is_non_deterministic_allowed(), - window_size_left=window_size_left, - window_size_right=window_size_right, - bottom_right_diagonal=config.bottom_right_diagonal, - softmax_type=int(config.softmax_type.value), + is_ragged=config.qkv_layout.is_thd(), + rng_offset_increment=16, + **graph.ffi_attrs(), ) @staticmethod @@ -650,9 +521,6 @@ def convert_to_2d(offsets, batch, max_seqlen): k_seq_offsets, k_seq_offsets >= 0, fill_value=kv_batch * kv_max_seqlen ) - q_cu_seqlen = generate_cu_seqlen(q_seqlen.flatten()) - kv_cu_seqlen = generate_cu_seqlen(kv_seqlen.flatten()) - output, softmax_aux, max_tensor, rng_state, _ = FusedAttnFwdPrimitive.inner_primitive.bind( q, k, @@ -660,8 +528,8 @@ def convert_to_2d(offsets, batch, max_seqlen): bias, softmax_offset, seed, - q_cu_seqlen, - kv_cu_seqlen, + q_seqlen.flatten(), + kv_seqlen.flatten(), q_seq_offsets, k_seq_offsets, _q_segment_ids, @@ -929,7 +797,7 @@ def abstract( """ Fused attention bwd abstract """ - del softmax_aux_aval, rng_state_aval, output_aval + del rng_state_aval q_dtype = dtypes.canonicalize_dtype(q_aval.dtype) k_dtype = dtypes.canonicalize_dtype(k_aval.dtype) @@ -956,38 +824,36 @@ def abstract( v_head_dim, ) = FusedAttnHelper.parse_qkv_aval(q_aval, k_aval, v_aval, config.qkv_layout) - if config.attn_bias_type == AttnBiasType.NO_BIAS: - bias_batch = bias_heads = 0 - else: - *bias_batch_shape, bias_heads, _, _ = bias_aval.shape - bias_batch = reduce(operator.mul, bias_batch_shape) - - deterministic = not FusedAttnHelper.is_non_deterministic_allowed() - - input_batch = reduce(operator.mul, batch_shape) - wkspace_shape, wkspace_dtype = transformer_engine_jax.get_fused_attn_bwd_workspace_sizes( - input_batch, - bias_batch, - q_max_seqlen, - kv_max_seqlen, + backend = FusedAttnHelper( + config.is_training, + q_dtype, + k_dtype, + config.qkv_layout, + config.attn_bias_type, + config.attn_mask_type, + config.softmax_type, + config.dropout_probability, attn_heads, num_gqa_groups, - bias_heads, + q_max_seqlen, + kv_max_seqlen, qk_head_dim, v_head_dim, - config.scaling_factor, - config.dropout_probability, - config.attn_bias_type.value, - config.attn_mask_type.value, - config.softmax_type.value, - config.qkv_layout.value, - jax_dtype_to_te_dtype(q_aval.dtype), - config.is_training, - deterministic, - config.max_segments_per_seq, - config.window_size[0], - config.window_size[1], - config.bottom_right_diagonal, + config.window_size, + False, + ).get_fused_attn_backend() + if backend != NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: + raise ValueError(f"Unsupported {backend=}") + + graph_info = build_bwd_graph( + q_aval, + k_aval, + v_aval, + bias_aval, + softmax_aux_aval, + output_aval, + doutput_aval, + config, ) dq_aval = q_aval.update(shape=q_aval.shape, dtype=q_dtype) @@ -995,7 +861,7 @@ def abstract( dv_aval = v_aval.update(shape=v_aval.shape, dtype=v_dtype) dbias_aval = bias_aval.update(shape=bias_aval.shape, dtype=bias_dtype) wkspace_aval = q_aval.update( - shape=wkspace_shape, dtype=te_dtype_to_jax_dtype(wkspace_dtype) + shape=(graph_info.graph.workspace_size,), dtype=jnp.uint8 ) # Validate incoming softmax_offset shape and dtype @@ -1060,30 +926,16 @@ def lowering( """ q_aval, k_aval, v_aval, bias_aval, *_ = ctx.avals_in - ( - batch_shape, - q_max_seqlen, - kv_max_seqlen, - attn_heads, - num_gqa_groups, - qk_head_dim, - v_head_dim, - ) = FusedAttnHelper.parse_qkv_aval(q_aval, k_aval, v_aval, config.qkv_layout) - - input_batch = reduce(operator.mul, batch_shape) - - if config.attn_bias_type == AttnBiasType.NO_BIAS: - bias_batch = bias_heads = 0 - else: - *bias_batch_shape, bias_heads, _, _ = bias_aval.shape - bias_batch = reduce(operator.mul, bias_batch_shape) - - if config.cp_striped_window_size is not None: - window_size_left = config.cp_striped_window_size[0] - window_size_right = config.cp_striped_window_size[1] - else: - window_size_left = config.window_size[0] - window_size_right = config.window_size[1] + graph = build_bwd_graph( + q_aval, + k_aval, + v_aval, + bias_aval, + ctx.avals_in[5], + ctx.avals_in[7], + ctx.avals_in[8], + config, + ).graph return ffi.ffi_lowering(FusedAttnBwdPrimitive.name)( ctx, @@ -1104,27 +956,8 @@ def lowering( kv_segment_ids, q_segment_pos, kv_segment_pos, # ffi_lowering needs number of parameters meets primitive.lowering - input_batch=input_batch, - bias_batch=bias_batch, - q_max_seqlen=q_max_seqlen, - kv_max_seqlen=kv_max_seqlen, - attn_heads=attn_heads, - num_gqa_groups=num_gqa_groups, - bias_heads=bias_heads, - qk_head_dim=qk_head_dim, - v_head_dim=v_head_dim, - max_segments_per_seq=config.max_segments_per_seq, - scaling_factor=float(config.scaling_factor), - dropout_probability=float(config.dropout_probability), - bias_type=int(config.attn_bias_type.value), - mask_type=int(config.attn_mask_type.value), - qkv_layout=int(config.qkv_layout.value), - is_training=config.is_training, - deterministic=not FusedAttnHelper.is_non_deterministic_allowed(), - window_size_left=window_size_left, - window_size_right=window_size_right, - bottom_right_diagonal=config.bottom_right_diagonal, - softmax_type=int(config.softmax_type.value), + is_ragged=config.qkv_layout.is_thd(), + **graph.ffi_attrs(), ) @staticmethod @@ -1223,9 +1056,6 @@ def convert_to_2d(offsets, batch, max_seqlen): k_seq_offsets, k_seq_offsets >= 0, fill_value=kv_batch * kv_max_seqlen ) - q_cu_seqlen = generate_cu_seqlen(q_seqlen.flatten()) - kv_cu_seqlen = generate_cu_seqlen(kv_seqlen.flatten()) - dq, dk, dv, dbias, dsoftmax_offset, _ = FusedAttnBwdPrimitive.inner_primitive.bind( q, k, @@ -1236,8 +1066,8 @@ def convert_to_2d(offsets, batch, max_seqlen): rng_state, output, doutput, - q_cu_seqlen, - kv_cu_seqlen, + q_seqlen.flatten(), + kv_seqlen.flatten(), q_seq_offsets, k_seq_offsets, _q_segment_ids, diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py new file mode 100644 index 00000000000..a9a32c89284 --- /dev/null +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -0,0 +1,1323 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Python cuDNN frontend graphs for the standard JAX fused-attention path.""" + +from __future__ import annotations + +import operator +import os +from dataclasses import dataclass +from functools import reduce +from typing import Any + +import jax.numpy as jnp +import numpy as np + +from .cudnn_graph import ( + GraphBinding, + SerializedGraph, + cudnn_data_type, + dtype_name, + finalize_graph, + import_cudnn, + serialized_graph, +) +from .misc import get_all_device_compute_capability, get_cudnn_version + +# Stable graph tensor UIDs. They are deliberately shared by forward and backward +# so serialized graphs and variant packs remain easy to inspect. +_UID_Q = 1 +_UID_K = 2 +_UID_V = 3 +_UID_O = 4 +_UID_STATS = 5 +_UID_MAX = 6 +_UID_BIAS = 7 +_UID_SINK = 8 +_UID_SEQ_Q = 9 +_UID_SEQ_KV = 10 +_UID_OFFSET_Q = 11 +_UID_OFFSET_K = 12 +_UID_OFFSET_V = 13 +_UID_OFFSET_O = 14 +_UID_OFFSET_STATS = 15 +_UID_DROPOUT_SEED = 16 +_UID_DROPOUT_OFFSET = 17 +_UID_DO = 18 +_UID_DQ = 19 +_UID_DK = 20 +_UID_DV = 21 +_UID_DBIAS = 22 +_UID_DSINK = 23 +_UID_ATTN_SCALE = 24 +_UID_OFFSET_MULT_Q = 30 +_UID_OFFSET_MULT_K = 31 +_UID_OFFSET_MULT_V = 32 +_UID_OFFSET_MULT_O = 33 +_UID_OFFSET_MULT_STATS = 34 +_UID_OFFSET_MULT_MAX = 35 + + +@dataclass(frozen=True) +class AttentionGraphInfo: + """Serialized graph plus abstract-output information used by JAX lowering.""" + + graph: SerializedGraph + stats_shape: tuple[int, ...] + max_shape: tuple[int, ...] + + +@dataclass(frozen=True) +class _LayoutInfo: + batch_shape: tuple[int, ...] + input_batch: int + q_max_seqlen: int + kv_max_seqlen: int + q_heads: int + kv_heads: int + qk_dim: int + v_dim: int + + +_graph_cache: dict[tuple[Any, ...], AttentionGraphInfo] = {} + + +def _layout_info(q_aval, k_aval, v_aval, layout) -> _LayoutInfo: + """Parse TE's three supported JAX QKV layout groups.""" + if layout.is_qkvpacked(): + *batch_shape, q_seqlen, packed, q_heads, qk_dim = q_aval.shape + if packed != 3: + raise ValueError( + f"QKV-packed fused attention expects dimension 3, got {q_aval.shape}." + ) + kv_seqlen = q_seqlen + kv_heads = q_heads + v_dim = qk_dim + elif layout.is_kvpacked(): + *batch_shape, q_seqlen, q_heads, qk_dim = q_aval.shape + *kv_batch_shape, kv_seqlen, packed, kv_heads, v_dim = k_aval.shape + if tuple(batch_shape) != tuple(kv_batch_shape) or packed != 2: + raise ValueError( + f"Invalid KV-packed fused-attention shapes: {q_aval}, {k_aval}." + ) + if qk_dim != v_dim: + raise ValueError( + "KV-packed fused attention requires equal QK and V head dimensions." + ) + elif layout.is_separate(): + *batch_shape, q_seqlen, q_heads, qk_dim = q_aval.shape + *k_batch_shape, kv_seqlen, kv_heads, k_dim = k_aval.shape + *v_batch_shape, v_seqlen, v_heads, v_dim = v_aval.shape + if tuple(batch_shape) != tuple(k_batch_shape) or tuple(batch_shape) != tuple( + v_batch_shape + ): + raise ValueError( + "Separate Q, K and V tensors must have matching batch shapes." + ) + if qk_dim != k_dim or kv_seqlen != v_seqlen or kv_heads != v_heads: + raise ValueError( + "Separate fused-attention K and V shapes are inconsistent." + ) + else: + raise ValueError(f"Unsupported JAX fused-attention layout: {layout}.") + return _LayoutInfo( + batch_shape=tuple(int(dim) for dim in batch_shape), + input_batch=reduce(operator.mul, batch_shape, 1), + q_max_seqlen=int(q_seqlen), + kv_max_seqlen=int(kv_seqlen), + q_heads=int(q_heads), + kv_heads=int(kv_heads), + qk_dim=int(qk_dim), + v_dim=int(v_dim), + ) + + +def _matrix_stride( + info: _LayoutInfo, layout, matrix: str, graph_sq: int, graph_skv: int +): + """Port generateMatrixStrides for JAX's BSHD/THD layout subset.""" + if matrix in ("q", "o"): + heads = info.q_heads + dim = info.qk_dim if matrix == "q" else info.v_dim + seqlen = graph_sq + else: + heads = info.kv_heads + dim = info.qk_dim if matrix == "k" else info.v_dim + seqlen = graph_skv + + if matrix in ("q", "k", "v") and layout.is_qkvpacked(): + return ( + graph_sq * 3 * info.q_heads * info.qk_dim, + dim, + 3 * info.q_heads * info.qk_dim, + 1, + ) + if matrix in ("k", "v") and layout.is_kvpacked(): + return ( + graph_skv * 2 * info.kv_heads * info.qk_dim, + dim, + 2 * info.kv_heads * info.qk_dim, + 1, + ) + return (seqlen * heads * dim, dim, heads * dim, 1) + + +def _qkv_bindings(info: _LayoutInfo, layout, itemsize: int, *, outputs: bool = False): + """Return UID bindings for separate or physically packed QKV buffers.""" + if outputs: + uids = (_UID_DQ, _UID_DK, _UID_DV) + else: + uids = (_UID_Q, _UID_K, _UID_V) + if layout.is_qkvpacked(): + stride = info.q_heads * info.qk_dim * itemsize + return ( + GraphBinding(uids[0], 0, 0), + GraphBinding(uids[1], 0, stride), + GraphBinding(uids[2], 0, 2 * stride), + ) + if layout.is_kvpacked(): + stride = info.kv_heads * info.qk_dim * itemsize + return ( + GraphBinding(uids[0], 0, 0), + GraphBinding(uids[1], 1, 0), + GraphBinding(uids[2], 1, stride), + ) + return tuple(GraphBinding(uid, index, 0) for index, uid in enumerate(uids)) + + +def _is_bias(config) -> bool: + return getattr(config.attn_bias_type, "name", "") == "POST_SCALE_BIAS" + + +def _is_padding(config) -> bool: + return bool(config.attn_mask_type.is_padding()) + + +def _is_causal(config) -> bool: + name = getattr(config.attn_mask_type, "name", "") + return name in ("CAUSAL_MASK", "PADDING_CAUSAL_MASK") + + +def _is_bottom_right(config) -> bool: + return bool(config.attn_mask_type.is_bottom_right()) + + +def _has_sink(config) -> bool: + return getattr(config.softmax_type, "name", "") != "VANILLA_SOFTMAX" + + +def _is_dropout(config) -> bool: + return bool(config.is_training and config.dropout_probability != 0.0) + + +def _device_arch() -> int: + capabilities = get_all_device_compute_capability() + return int(capabilities[0]) if capabilities else 0 + + +def _graph_dimensions(info: _LayoutInfo, config): + """Return logical cuDNN B/H/S dimensions and physical auxiliary shapes.""" + is_ragged = config.qkv_layout.is_thd() + cudnn_version = get_cudnn_version() + arch = _device_arch() + use_ragged_stats = is_ragged and cudnn_version >= (9, 6, 0) and arch != 120 + if is_ragged: + graph_batch = info.input_batch * int(config.max_segments_per_seq) + if arch == 120: + graph_sq = info.q_max_seqlen + graph_skv = info.kv_max_seqlen + else: + # Static token extents replace common's quantized max-token buckets. JAX already + # compiles per local shape, so exact extents avoid unnecessary graph variants. + graph_sq = info.input_batch * info.q_max_seqlen + graph_skv = info.input_batch * info.kv_max_seqlen + else: + graph_batch = info.input_batch + graph_sq = info.q_max_seqlen + graph_skv = info.kv_max_seqlen + + if is_ragged and cudnn_version >= (9, 6, 0): + stats_shape = (*info.batch_shape, info.q_max_seqlen, info.q_heads, 1) + elif cudnn_version >= (9, 6, 0): + stats_shape = (*info.batch_shape, info.q_heads, info.q_max_seqlen, 1) + else: + stats_shape = ( + *info.batch_shape, + info.q_heads, + info.q_max_seqlen, + int(config.max_segments_per_seq), + ) + if config.return_max_logit: + max_shape = ( + (*info.batch_shape, info.q_max_seqlen, info.q_heads, 1) + if use_ragged_stats + else (*info.batch_shape, info.q_heads, info.q_max_seqlen, 1) + ) + else: + max_shape = (0,) + return graph_batch, graph_sq, graph_skv, use_ragged_stats, stats_shape, max_shape + + +def _tensor(graph, cudnn, *, name, dim, stride, dtype, uid): + return graph.tensor( + name=name, + dim=tuple(int(x) for x in dim), + stride=tuple(int(x) for x in stride), + data_type=dtype, + uid=uid, + ) + + +def _ragged_offset(graph, cudnn, name: str, uid: int, graph_batch: int): + return _tensor( + graph, + cudnn, + name=name, + dim=(graph_batch + 1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=uid, + ) + + +def _set_ragged( + tensor, + offset, + multiplier: int, + *, + graph=None, + cudnn=None, + multiplier_uid=None, + scalar_uids=None, + scalar_values=None, + materialize_multiplier: bool = False, +): + """Attach token-unit offsets, optionally materializing element offsets in the graph. + + The unified SDPA engine understands ``ragged_offset_multiplier`` directly. cuDNN's + dropout forward path currently selects the composite engine, so for that case an + INT64 pointwise node performs the conversion that TE common used to launch as a + separate CUDA kernel. + """ + if materialize_multiplier: + use_int64 = get_cudnn_version() >= (9, 5, 0) + offset_data_type = cudnn.data_type.INT64 if use_int64 else cudnn.data_type.INT32 + offset_numpy_type = np.int64 if use_int64 else np.int32 + scale = _scalar_tensor( + graph, + cudnn, + f"ragged_multiplier_{multiplier_uid}", + multiplier_uid, + offset_data_type, + ) + effective_offset = graph.mul( + offset, + scale, + compute_data_type=offset_data_type, + name=f"scale_ragged_offset_{multiplier_uid}", + ) + effective_offset.set_data_type(offset_data_type) + effective_offset.set_dim((offset.get_dim()[0], 1, 1, 1)).set_stride( + (1, 1, 1, 1) + ) + tensor.set_ragged_offset(effective_offset) + scalar_uids.append(multiplier_uid) + scalar_values.append(np.asarray(multiplier, dtype=offset_numpy_type).tobytes()) + return effective_offset + else: + tensor.set_ragged_offset(offset) + tensor.set_ragged_offset_multiplier(int(multiplier)) + return offset + + +def _mask_options(cudnn, info: _LayoutInfo, config): + is_padding = _is_padding(config) + causal = _is_causal(config) + bottom_right = _is_bottom_right(config) + bottom_right_diagonal = bool(config.bottom_right_diagonal) + if bottom_right and info.q_max_seqlen == info.kv_max_seqlen and not is_padding: + causal = True + bottom_right = False + bottom_right_diagonal = False + window_left, window_right = ( + config.cp_striped_window_size + if config.cp_striped_window_size is not None + else config.window_size + ) + cudnn_version = get_cudnn_version() + options = { + "diagonal_alignment": ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if bottom_right_diagonal or bottom_right + else cudnn.diagonal_alignment.TOP_LEFT + ), + } + # Before cuDNN 9.6 the preferred right-band API was unavailable, so preserve + # the legacy causal flags used by the C++ frontend graph. + if cudnn_version < (9, 6, 0): + options["use_causal_mask"] = causal + options["use_causal_mask_bottom_right"] = bottom_right + if cudnn_version >= (9, 2, 0) and window_left != -1: + options["diagonal_band_left_bound"] = int(window_left) + 1 + if cudnn_version >= (9, 6, 0): + if window_right != -1: + options["diagonal_band_right_bound"] = int(window_right) + elif causal or bottom_right: + options["diagonal_band_right_bound"] = 0 + return options + + +def _scalar_tensor(graph, cudnn, name: str, uid: int, dtype): + return graph.tensor( + name=name, + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=dtype, + is_pass_by_value=True, + uid=uid, + ) + + +def _cache_key(direction: str, q_aval, k_aval, v_aval, bias_aval, config, *extra_avals): + avals = (q_aval, k_aval, v_aval, bias_aval, *extra_avals) + return ( + direction, + config, + tuple((tuple(aval.shape), dtype_name(aval.dtype)) for aval in avals), + get_cudnn_version(), + _device_arch(), + ) + + +def build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGraphInfo: + """Build or retrieve the standard fused-attention forward graph.""" + key = _cache_key("fwd", q_aval, k_aval, v_aval, bias_aval, config) + if key not in _graph_cache: + _graph_cache[key] = _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) + return _graph_cache[key] + + +def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGraphInfo: + cudnn = import_cudnn() + info = _layout_info(q_aval, k_aval, v_aval, config.qkv_layout) + graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = ( + _graph_dimensions(info, config) + ) + io_dtype = cudnn_data_type(cudnn, q_aval.dtype) + graph = cudnn.pygraph( + io_data_type=io_dtype, + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + ) + + q = _tensor( + graph, + cudnn, + name="q", + dim=(graph_batch, info.q_heads, graph_sq, info.qk_dim), + stride=_matrix_stride(info, config.qkv_layout, "q", graph_sq, graph_skv), + dtype=io_dtype, + uid=_UID_Q, + ) + k = _tensor( + graph, + cudnn, + name="k", + dim=(graph_batch, info.kv_heads, graph_skv, info.qk_dim), + stride=_matrix_stride(info, config.qkv_layout, "k", graph_sq, graph_skv), + dtype=io_dtype, + uid=_UID_K, + ) + v = _tensor( + graph, + cudnn, + name="v", + dim=(graph_batch, info.kv_heads, graph_skv, info.v_dim), + stride=_matrix_stride(info, config.qkv_layout, "v", graph_sq, graph_skv), + dtype=io_dtype, + uid=_UID_V, + ) + + input_bindings = list( + _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) + ) + output_bindings = [GraphBinding(_UID_O, 0), GraphBinding(_UID_STATS, 1)] + scale = _scalar_tensor( + graph, cudnn, "attn_scale", _UID_ATTN_SCALE, cudnn.data_type.FLOAT + ) + scalar_uids = [_UID_ATTN_SCALE] + scalar_values = [np.asarray(config.scaling_factor, dtype=np.float32).tobytes()] + + kwargs = { + "name": "te_fused_attention", + "q": q, + "k": k, + "v": v, + "generate_stats": True, + "attn_scale": scale, + **_mask_options(cudnn, info, config), + } + + if _is_bias(config): + *bias_batch_shape, bias_heads, bias_sq, bias_skv = bias_aval.shape + bias_batch = reduce(operator.mul, bias_batch_shape, 1) + bias = _tensor( + graph, + cudnn, + name="bias", + dim=(bias_batch, bias_heads, bias_sq, bias_skv), + stride=(bias_heads * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1), + dtype=io_dtype, + uid=_UID_BIAS, + ) + kwargs["bias"] = bias + input_bindings.append(GraphBinding(_UID_BIAS, 3)) + + if _has_sink(config): + sink = _tensor( + graph, + cudnn, + name="softmax_offset", + dim=(1, info.q_heads, 1, 1), + stride=(info.q_heads, 1, 1, 1), + dtype=cudnn.data_type.FLOAT, + uid=_UID_SINK, + ) + kwargs["sink_token"] = sink + input_bindings.append(GraphBinding(_UID_SINK, 4)) + + if _is_padding(config): + seq_q = _tensor( + graph, + cudnn, + name="seq_len_q", + dim=(graph_batch, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=_UID_SEQ_Q, + ) + seq_kv = _tensor( + graph, + cudnn, + name="seq_len_kv", + dim=(graph_batch, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=_UID_SEQ_KV, + ) + kwargs.update(use_padding_mask=True, seq_len_q=seq_q, seq_len_kv=seq_kv) + input_bindings.extend( + (GraphBinding(_UID_SEQ_Q, 6), GraphBinding(_UID_SEQ_KV, 7)) + ) + + offset_q = offset_k = offset_v = offset_o = offset_stats = None + if config.qkv_layout.is_thd(): + offset_q = _ragged_offset(graph, cudnn, "offset_q", _UID_OFFSET_Q, graph_batch) + offset_k = _ragged_offset(graph, cudnn, "offset_k", _UID_OFFSET_K, graph_batch) + offset_v = _ragged_offset(graph, cudnn, "offset_v", _UID_OFFSET_V, graph_batch) + offset_o = _ragged_offset(graph, cudnn, "offset_o", _UID_OFFSET_O, graph_batch) + materialize_offsets = True + _set_ragged( + q, + offset_q, + info.q_heads * info.qk_dim * (3 if config.qkv_layout.is_qkvpacked() else 1), + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_Q, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=materialize_offsets, + ) + if config.qkv_layout.is_qkvpacked(): + kv_offset_operand = 8 + kv_multiplier = 3 * info.q_heads * info.qk_dim + else: + kv_offset_operand = 9 + kv_multiplier = ( + 2 * info.kv_heads * info.qk_dim + if config.qkv_layout.is_kvpacked() + else info.kv_heads * info.qk_dim + ) + _set_ragged( + k, + offset_k, + kv_multiplier, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_K, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=materialize_offsets, + ) + _set_ragged( + v, + offset_v, + kv_multiplier + if config.qkv_layout.is_kvpacked() + else info.kv_heads * info.v_dim, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_V, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=materialize_offsets, + ) + input_bindings.extend( + ( + GraphBinding(_UID_OFFSET_Q, 8), + GraphBinding(_UID_OFFSET_K, kv_offset_operand), + GraphBinding(_UID_OFFSET_V, kv_offset_operand), + GraphBinding(_UID_OFFSET_O, 8), + ) + ) + + if _is_dropout(config): + seed = _tensor( + graph, + cudnn, + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT64, + uid=_UID_DROPOUT_SEED, + ) + offset = _tensor( + graph, + cudnn, + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT64, + uid=_UID_DROPOUT_OFFSET, + ) + kwargs["dropout"] = (float(config.dropout_probability), seed, offset) + # rng_state is forward result #3: two int64 values stored as four uint32 values. + output_bindings.extend( + ( + GraphBinding(_UID_DROPOUT_SEED, 3, 0), + GraphBinding(_UID_DROPOUT_OFFSET, 3, 8), + ) + ) + + if config.return_max_logit: + max_tensor = _tensor( + graph, + cudnn, + name="max_logit", + dim=(graph_batch, info.q_heads, graph_sq, 1), + stride=( + (info.q_heads * graph_sq, 1, info.q_heads, 1) + if ragged_stats + else (info.q_heads * graph_sq, graph_sq, 1, 1) + ), + dtype=cudnn.data_type.FLOAT, + uid=_UID_MAX, + ) + if ragged_stats: + offset_stats = _ragged_offset( + graph, cudnn, "offset_stats", _UID_OFFSET_STATS, graph_batch + ) + _set_ragged( + max_tensor, + offset_stats, + info.q_heads, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_MAX, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 8)) + max_tensor.set_output(True) + kwargs["score_max"] = max_tensor + output_bindings.append(GraphBinding(_UID_MAX, 2)) + + output, stats = graph.sdpa(**kwargs) + output.set_output(True).set_uid(_UID_O).set_dim( + (graph_batch, info.q_heads, graph_sq, info.v_dim) + ).set_stride(_matrix_stride(info, config.qkv_layout, "o", graph_sq, graph_skv)) + output.set_data_type(io_dtype) + if config.qkv_layout.is_thd(): + _set_ragged( + output, + offset_o, + info.q_heads * info.v_dim, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_O, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + + stats.set_output(True).set_uid(_UID_STATS).set_data_type(cudnn.data_type.FLOAT) + stats.set_dim((graph_batch, info.q_heads, graph_sq, 1)) + if ragged_stats: + if offset_stats is None: + offset_stats = _ragged_offset( + graph, cudnn, "offset_stats", _UID_OFFSET_STATS, graph_batch + ) + input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 8)) + stats.set_stride((info.q_heads * graph_sq, 1, info.q_heads, 1)) + _set_ragged( + stats, + offset_stats, + info.q_heads, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_STATS, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + else: + stats.set_stride((info.q_heads * graph_sq, graph_sq, 1, 1)) + + workspace, data, version = finalize_graph( + cudnn, graph, description="fused-attention forward" + ) + result = serialized_graph( + serialized_graph_data=data, + cudnn_frontend_version=version, + workspace_size=workspace, + input_bindings=input_bindings, + output_bindings=output_bindings, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + ) + return AttentionGraphInfo(result, stats_shape, max_shape) + + +def build_bwd_graph( + q_aval, + k_aval, + v_aval, + bias_aval, + stats_aval, + output_aval, + doutput_aval, + config, +) -> AttentionGraphInfo: + """Build or retrieve the standard fused-attention backward graph.""" + key = _cache_key( + "bwd", + q_aval, + k_aval, + v_aval, + bias_aval, + config, + stats_aval, + output_aval, + doutput_aval, + ) + if key not in _graph_cache: + _graph_cache[key] = _build_bwd_graph( + q_aval, + k_aval, + v_aval, + bias_aval, + stats_aval, + output_aval, + doutput_aval, + config, + ) + return _graph_cache[key] + + +def _build_bwd_graph( + q_aval, + k_aval, + v_aval, + bias_aval, + stats_aval, + output_aval, + doutput_aval, + config, +) -> AttentionGraphInfo: + cudnn = import_cudnn() + info = _layout_info(q_aval, k_aval, v_aval, config.qkv_layout) + graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = ( + _graph_dimensions(info, config) + ) + io_dtype = cudnn_data_type(cudnn, q_aval.dtype) + graph = cudnn.pygraph( + io_data_type=io_dtype, + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + ) + + def io_tensor(name, dim, stride, uid, dtype=io_dtype): + return _tensor( + graph, cudnn, name=name, dim=dim, stride=stride, dtype=dtype, uid=uid + ) + + q = io_tensor( + "q", + (graph_batch, info.q_heads, graph_sq, info.qk_dim), + _matrix_stride(info, config.qkv_layout, "q", graph_sq, graph_skv), + _UID_Q, + ) + k = io_tensor( + "k", + (graph_batch, info.kv_heads, graph_skv, info.qk_dim), + _matrix_stride(info, config.qkv_layout, "k", graph_sq, graph_skv), + _UID_K, + ) + v = io_tensor( + "v", + (graph_batch, info.kv_heads, graph_skv, info.v_dim), + _matrix_stride(info, config.qkv_layout, "v", graph_sq, graph_skv), + _UID_V, + ) + output = io_tensor( + "o", + (graph_batch, info.q_heads, graph_sq, info.v_dim), + _matrix_stride(info, config.qkv_layout, "o", graph_sq, graph_skv), + _UID_O, + ) + doutput = io_tensor( + "dO", + (graph_batch, info.q_heads, graph_sq, info.v_dim), + _matrix_stride(info, config.qkv_layout, "o", graph_sq, graph_skv), + _UID_DO, + ) + stats = io_tensor( + "stats", + (graph_batch, info.q_heads, graph_sq, 1), + ( + (info.q_heads * graph_sq, 1, info.q_heads, 1) + if ragged_stats + else (info.q_heads * graph_sq, graph_sq, 1, 1) + ), + _UID_STATS, + cudnn.data_type.FLOAT, + ) + + itemsize = jnp.dtype(q_aval.dtype).itemsize + input_bindings = list(_qkv_bindings(info, config.qkv_layout, itemsize)) + input_bindings.extend( + ( + GraphBinding(_UID_STATS, 5), + GraphBinding(_UID_O, 7), + GraphBinding(_UID_DO, 8), + ) + ) + output_bindings = list( + _qkv_bindings(info, config.qkv_layout, itemsize, outputs=True) + ) + scale = _scalar_tensor( + graph, cudnn, "attn_scale", _UID_ATTN_SCALE, cudnn.data_type.FLOAT + ) + scalar_uids = [_UID_ATTN_SCALE] + scalar_values = [np.asarray(config.scaling_factor, dtype=np.float32).tobytes()] + + kwargs = { + "name": "te_fused_attention_backward", + "q": q, + "k": k, + "v": v, + "o": output, + "dO": doutput, + "stats": stats, + "attn_scale": scale, + **_mask_options(cudnn, info, config), + } + if get_cudnn_version() >= (9, 0, 0): + kwargs["use_deterministic_algorithm"] = not bool( + int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) + ) + if ragged_stats: + kwargs["max_total_seq_len_q"] = graph_sq + if ( + config.qkv_layout.is_thd() + and get_cudnn_version() >= (9, 6, 0) + and _device_arch() != 120 + ): + kwargs["max_total_seq_len_kv"] = graph_skv + + if _is_bias(config): + *bias_batch_shape, bias_heads, bias_sq, bias_skv = bias_aval.shape + bias_batch = reduce(operator.mul, bias_batch_shape, 1) + bias = io_tensor( + "bias", + (bias_batch, bias_heads, bias_sq, bias_skv), + (bias_heads * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1), + _UID_BIAS, + ) + kwargs["bias"] = bias + input_bindings.append(GraphBinding(_UID_BIAS, 3)) + if not (bias_batch == 1 and bias_heads == 1 and bias_sq == 1): + dbias = io_tensor( + "dBias", + (bias_batch, bias_heads, bias_sq, bias_skv), + (bias_heads * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1), + _UID_DBIAS, + ) + dbias.set_output(True) + kwargs["dBias"] = dbias + output_bindings.append(GraphBinding(_UID_DBIAS, 3)) + + if _has_sink(config): + sink = io_tensor( + "softmax_offset", + (1, info.q_heads, 1, 1), + (info.q_heads, 1, 1, 1), + _UID_SINK, + cudnn.data_type.FLOAT, + ) + dsink = io_tensor( + "dsoftmax_offset", + (1, info.q_heads, 1, 1), + (info.q_heads, 1, 1, 1), + _UID_DSINK, + cudnn.data_type.FLOAT, + ) + dsink.set_output(True) + kwargs.update(sink_token=sink, dSink_token=dsink) + input_bindings.append(GraphBinding(_UID_SINK, 4)) + output_bindings.append(GraphBinding(_UID_DSINK, 4)) + + if _is_padding(config): + seq_q = io_tensor( + "seq_len_q", + (graph_batch, 1, 1, 1), + (1, 1, 1, 1), + _UID_SEQ_Q, + cudnn.data_type.INT32, + ) + seq_kv = io_tensor( + "seq_len_kv", + (graph_batch, 1, 1, 1), + (1, 1, 1, 1), + _UID_SEQ_KV, + cudnn.data_type.INT32, + ) + kwargs.update(use_padding_mask=True, seq_len_q=seq_q, seq_len_kv=seq_kv) + input_bindings.extend( + (GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10)) + ) + + if config.qkv_layout.is_thd(): + offset_q = _ragged_offset(graph, cudnn, "offset_q", _UID_OFFSET_Q, graph_batch) + offset_k = _ragged_offset(graph, cudnn, "offset_k", _UID_OFFSET_K, graph_batch) + offset_v = _ragged_offset(graph, cudnn, "offset_v", _UID_OFFSET_V, graph_batch) + offset_o = _ragged_offset(graph, cudnn, "offset_o", _UID_OFFSET_O, graph_batch) + q_mult = ( + info.q_heads * info.qk_dim * (3 if config.qkv_layout.is_qkvpacked() else 1) + ) + kv_mult = ( + 3 * info.q_heads * info.qk_dim + if config.qkv_layout.is_qkvpacked() + else 2 * info.kv_heads * info.qk_dim + if config.qkv_layout.is_kvpacked() + else info.kv_heads * info.qk_dim + ) + v_mult = ( + kv_mult + if not config.qkv_layout.is_separate() + else info.kv_heads * info.v_dim + ) + effective_q_offset = _set_ragged( + q, + offset_q, + q_mult, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_Q, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + effective_k_offset = _set_ragged( + k, + offset_k, + kv_mult, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_K, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + effective_v_offset = _set_ragged( + v, + offset_v, + v_mult, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_V, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + effective_o_offset = _set_ragged( + output, + offset_o, + info.q_heads * info.v_dim, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_O, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + doutput.set_ragged_offset(effective_o_offset) + kv_operand = 11 if config.qkv_layout.is_qkvpacked() else 12 + input_bindings.extend( + ( + GraphBinding(_UID_OFFSET_Q, 11), + GraphBinding(_UID_OFFSET_K, kv_operand), + GraphBinding(_UID_OFFSET_V, kv_operand), + GraphBinding(_UID_OFFSET_O, 11), + ) + ) + if ragged_stats: + offset_stats = _ragged_offset( + graph, cudnn, "offset_stats", _UID_OFFSET_STATS, graph_batch + ) + _set_ragged( + stats, + offset_stats, + info.q_heads, + graph=graph, + cudnn=cudnn, + multiplier_uid=_UID_OFFSET_MULT_STATS, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + materialize_multiplier=True, + ) + input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 11)) + + if _is_dropout(config): + seed = io_tensor( + "dropout_seed", + (1, 1, 1, 1), + (1, 1, 1, 1), + _UID_DROPOUT_SEED, + cudnn.data_type.INT64, + ) + offset = io_tensor( + "dropout_offset", + (1, 1, 1, 1), + (1, 1, 1, 1), + _UID_DROPOUT_OFFSET, + cudnn.data_type.INT64, + ) + kwargs["dropout"] = (float(config.dropout_probability), seed, offset) + input_bindings.extend( + ( + GraphBinding(_UID_DROPOUT_SEED, 6, 0), + GraphBinding(_UID_DROPOUT_OFFSET, 6, 8), + ) + ) + + dq, dk, dv = graph.sdpa_backward(**kwargs) + q_stride = _matrix_stride(info, config.qkv_layout, "q", graph_sq, graph_skv) + k_stride = _matrix_stride(info, config.qkv_layout, "k", graph_sq, graph_skv) + v_stride = _matrix_stride(info, config.qkv_layout, "v", graph_sq, graph_skv) + dq.set_output(True).set_uid(_UID_DQ).set_dim( + (graph_batch, info.q_heads, graph_sq, info.qk_dim) + ).set_stride(q_stride) + dk.set_output(True).set_uid(_UID_DK).set_dim( + (graph_batch, info.kv_heads, graph_skv, info.qk_dim) + ).set_stride(k_stride) + dv.set_output(True).set_uid(_UID_DV).set_dim( + (graph_batch, info.kv_heads, graph_skv, info.v_dim) + ).set_stride(v_stride) + if config.qkv_layout.is_thd(): + dq.set_ragged_offset(effective_q_offset) + dk.set_ragged_offset(effective_k_offset) + dv.set_ragged_offset(effective_v_offset) + + workspace, data, version = finalize_graph( + cudnn, graph, description="fused-attention backward" + ) + result = serialized_graph( + serialized_graph_data=data, + cudnn_frontend_version=version, + workspace_size=workspace, + input_bindings=input_bindings, + output_bindings=output_bindings, + scalar_uids=scalar_uids, + scalar_values=scalar_values, + ) + return AttentionGraphInfo(result, stats_shape, max_shape) + + +def clear_graph_cache(): + """Clear the process-local serialized graph cache (primarily for tests).""" + _graph_cache.clear() + + +def _encoded_cudnn_version() -> int: + major, minor, patch = get_cudnn_version() + magnitude = 1000 if major < 9 else 10000 + return major * magnitude + minor * 100 + patch + + +def is_fused_attn_supported(helper) -> bool: + """JAX-local port of the F16/BF16 cuDNN attention compatibility policy. + + cuDNN frontend ``check_support`` remains authoritative when a concrete graph is + built. This early policy preserves the public fallback behavior for callers that + ask about availability before Q/K/V abstract values exist. + """ + if jnp.dtype(helper.q_dtype) not in ( + jnp.dtype(jnp.float16), + jnp.dtype(jnp.bfloat16), + ): + return False + if jnp.dtype(helper.q_dtype) != jnp.dtype(helper.kv_dtype): + return False + + version = _encoded_cudnn_version() + arch = _device_arch() + is_thd = helper.qkv_layout.is_thd() + is_training = bool(helper.is_training) + sq, skv = int(helper.q_max_seqlen), int(helper.kv_max_seqlen) + h, hg = int(helper.q_num_heads), int(helper.kv_num_heads) + dqk, dv = int(helper.head_dim_qk), int(helper.head_dim_v) + dropout = float(helper.dropout_probability) + bias_name = helper.attn_bias_type.name + mask_name = helper.attn_mask_type.name + softmax_name = helper.softmax_type.name + left, right = helper.window_size + deterministic = not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))) + + architecture_ok = ( + (version < 8903 and arch in (80, 90)) + or (version >= 8903 and 80 <= arch < 100) + or (version >= 90700 and arch >= 100) + ) + if version < 8900 or not architecture_ok: + return False + if version < 90000 and (sq % 64 or skv % 64): + return False + if version < 8907 and h != hg: + return False + if dqk % 8 or dv % 8: + return False + + standard_dim = dqk <= 128 and dv <= 128 + hopper_large_dim = ( + dqk <= 256 + and dv <= 256 + and ( + (not is_training and arch == 90 and version >= 90100) + or (is_training and arch == 90 and version >= 90500) + ) + ) + blackwell_fwd_any_dim = ( + not is_training and arch >= 100 and version >= 90900 and sq > 1 + ) + generic_fwd_any_dim = ( + not is_training + and version >= 91002 + and ( + sq > 1 + or (sq == 1 and mask_name not in ("CAUSAL_MASK", "PADDING_CAUSAL_MASK")) + ) + ) + blackwell_mla_bwd = ( + dqk == 192 and dv == 128 and is_training and arch >= 100 and version >= 91100 + ) + blackwell_d256_bwd = ( + dqk == 256 + and dv == 256 + and is_training + and 100 <= arch < 110 + and version >= (92500 if is_thd else 92300) + and bias_name == "NO_BIAS" + and dropout == 0.0 + and softmax_name == "VANILLA_SOFTMAX" + and ( + (left == -1 and right == -1) + or ( + mask_name + in ( + "CAUSAL_MASK", + "PADDING_CAUSAL_MASK", + "CAUSAL_BOTTOM_RIGHT_MASK", + "PADDING_CAUSAL_BOTTOM_RIGHT_MASK", + ) + and right in (-1, 0) + ) + ) + ) + if not ( + standard_dim + or hopper_large_dim + or blackwell_fwd_any_dim + or generic_fwd_any_dim + or blackwell_mla_bwd + or blackwell_d256_bwd + ): + return False + if ( + version >= 91100 + and is_training + and arch == 90 + and dqk >= 128 + and dv >= 128 + and (dqk, dv) != (192, 128) + and dqk != dv + ): + return False + + post_scale_bias_supported = bias_name == "POST_SCALE_BIAS" and ( + (version >= 8906 and arch >= 90) or (version >= 90000 and arch >= 80) + ) + if bias_name != "NO_BIAS" and not post_scale_bias_supported: + return False + + dense_basic_masks = mask_name in ( + "NO_MASK", + "CAUSAL_MASK", + "PADDING_MASK", + "PADDING_CAUSAL_MASK", + ) + thd_basic_masks = mask_name in ("PADDING_MASK", "PADDING_CAUSAL_MASK") + br_mask = mask_name == "CAUSAL_BOTTOM_RIGHT_MASK" + padding_br_mask = mask_name == "PADDING_CAUSAL_BOTTOM_RIGHT_MASK" + mask_ok = False + if version < 8906: + mask_ok = mask_name == "CAUSAL_MASK" and not is_thd + elif not is_thd and dense_basic_masks: + mask_ok = True + if version >= 90100 and is_thd and thd_basic_masks: + mask_ok = True + if ( + version >= 90300 + and not is_thd + and br_mask + and sq % 64 == 0 + and skv % 64 == 0 + and sq <= skv + and bias_name == "NO_BIAS" + and dropout == 0.0 + ): + mask_ok = True + if ( + version >= 90600 + and padding_br_mask + and sq % 64 == 0 + and skv % 64 == 0 + and sq <= skv + and bias_name == "NO_BIAS" + and dropout == 0.0 + ): + mask_ok = True + if version >= 90700: + mask_ok = ( + mask_name in ("NO_MASK", "CAUSAL_MASK") + or ( + mask_name + in ( + "PADDING_MASK", + "PADDING_CAUSAL_MASK", + "PADDING_CAUSAL_BOTTOM_RIGHT_MASK", + ) + and bias_name == "NO_BIAS" + and dropout == 0.0 + ) + or ( + mask_name + in ("CAUSAL_BOTTOM_RIGHT_MASK", "PADDING_CAUSAL_BOTTOM_RIGHT_MASK") + and sq <= skv + ) + ) + if not mask_ok: + return False + if ( + mask_name in ("PADDING_MASK", "PADDING_CAUSAL_MASK") + and bias_name == "POST_SCALE_BIAS" + ): + return False + + if is_thd and not ( + arch >= 90 and ((version >= 90100 and h == hg) or version >= 90600) + ): + return False + + full_window = left == -1 and right == -1 + if version < 90200: + window_ok = left == -1 and right in (-1, 0) + elif version < 90600: + window_ok = (full_window and mask_name == "NO_MASK") or ( + left >= -1 + and right == 0 + and mask_name in ("NO_MASK", "CAUSAL_MASK", "CAUSAL_BOTTOM_RIGHT_MASK") + and (mask_name != "CAUSAL_BOTTOM_RIGHT_MASK" or sq == skv) + and sq <= skv + and dropout == 0.0 + and bias_name == "NO_BIAS" + and not is_thd + ) + else: + bottom_right_swa_supported = ( + mask_name + not in ("CAUSAL_BOTTOM_RIGHT_MASK", "PADDING_CAUSAL_BOTTOM_RIGHT_MASK") + or arch < 100 + or sq == skv + or version > 90700 + ) + window_ok = ( + left == -1 + and right in (-1, 0) + or ( + left >= -1 + and right >= -1 + and mask_name + in ( + "NO_MASK", + "CAUSAL_MASK", + "PADDING_MASK", + "PADDING_CAUSAL_MASK", + "CAUSAL_BOTTOM_RIGHT_MASK", + "PADDING_CAUSAL_BOTTOM_RIGHT_MASK", + ) + and sq <= skv + and bias_name == "NO_BIAS" + and dropout == 0.0 + and bottom_right_swa_supported + ) + ) + if not window_ok: + return False + + if is_thd: + if helper.qkv_layout.is_qkvpacked(): + max_offset = 3 * h * dqk * sq + elif helper.qkv_layout.is_kvpacked(): + max_offset = max(h * dqk * sq, 2 * hg * dqk * skv) + else: + max_offset = max(h * dqk * sq, hg * dqk * skv, hg * dv * skv) + if max_offset > np.iinfo(np.int32).max and version < 90500: + return False + + if version in (91000, 91001): + return False + if version < 91301 and softmax_name != "VANILLA_SOFTMAX": + return False + if helper.return_max_logit and version < 92100: + return False + if arch >= 100 and is_training: + if deterministic: + if version < 91801 or dropout != 0.0 or bias_name != "NO_BIAS": + return False + elif dropout != 0.0 and bias_name != "NO_BIAS": + return False + if arch == 120 and ( + version < 91801 + or (deterministic and is_training) + or (is_thd and helper.qkv_layout.is_qkvpacked()) + ): + return False + return not ( + version == 91400 + and skv > 1024 + and left != -1 + and mask_name not in ("CAUSAL_MASK", "CAUSAL_BOTTOM_RIGHT_MASK") + ) diff --git a/transformer_engine/jax/cpp_extensions/cudnn_graph.py b/transformer_engine/jax/cpp_extensions/cudnn_graph.py new file mode 100644 index 00000000000..057150750ec --- /dev/null +++ b/transformer_engine/jax/cpp_extensions/cudnn_graph.py @@ -0,0 +1,259 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Shared cuDNN frontend graph serialization support for JAX custom calls. + +cuDNN frontend graphs are constructed and planned in Python while XLA executes a +serialized graph through a small, graph-agnostic FFI runtime. A binding maps a +cuDNN tensor UID to an XLA operand/result and an optional byte offset. Offsets +are required for TE's packed QKV layouts, where several logical cuDNN tensors +share one JAX buffer. +""" + +from __future__ import annotations + +import hashlib +import importlib +from collections.abc import Sequence +from dataclasses import dataclass +from typing import Any + +import jax +import jax.numpy as jnp +import numpy as np +import transformer_engine_jax + + +@dataclass(frozen=True) +class GraphBinding: + """Bind a cuDNN tensor UID to a JAX buffer and byte offset.""" + + uid: int + buffer_index: int + byte_offset: int = 0 + + +@dataclass(frozen=True) +class SerializedGraph: + """Serialized cuDNN graph and static metadata for the generic FFI executor.""" + + serialized_graph: bytes + graph_hash: tuple[int, int] + cudnn_frontend_version: int + workspace_size: int + input_uids: np.ndarray + input_buffer_indices: np.ndarray + input_byte_offsets: np.ndarray + output_uids: np.ndarray + output_buffer_indices: np.ndarray + output_byte_offsets: np.ndarray + scalar_uids: np.ndarray + scalar_sizes: np.ndarray + scalar_values: np.ndarray + + def ffi_attrs(self) -> dict[str, Any]: + """Return the static attributes consumed by the generic C++ executor.""" + return { + "serialized_graph": self.serialized_graph, + "graph_hash0": self.graph_hash[0], + "graph_hash1": self.graph_hash[1], + "cudnn_frontend_version": self.cudnn_frontend_version, + "input_uids": self.input_uids, + "input_buffer_indices": self.input_buffer_indices, + "input_byte_offsets": self.input_byte_offsets, + "output_uids": self.output_uids, + "output_buffer_indices": self.output_buffer_indices, + "output_byte_offsets": self.output_byte_offsets, + "scalar_uids": self.scalar_uids, + "scalar_sizes": self.scalar_sizes, + "scalar_values": self.scalar_values, + } + + +def row_major_stride(shape: Sequence[int]) -> tuple[int, ...]: + """Return element strides for a contiguous row-major tensor.""" + stride = [] + running = 1 + for dim in reversed(tuple(shape)): + stride.append(running) + running *= int(dim) + return tuple(reversed(stride)) + + +def bshd_as_bhsd_dim_stride( + shape: Sequence[int], +) -> tuple[tuple[int, ...], tuple[int, ...]]: + """Describe a contiguous BSHD buffer as cuDNN's logical BHSD tensor.""" + if len(shape) != 4: + raise ValueError(f"Expected a rank-4 BSHD tensor, got shape={shape}.") + batch, seqlen, heads, head_dim = (int(dim) for dim in shape) + return ( + (batch, heads, seqlen, head_dim), + (seqlen * heads * head_dim, head_dim, heads * head_dim, 1), + ) + + +def dtype_name(dtype) -> str: + """Stable dtype name for graph cache keys.""" + return str(jnp.dtype(dtype)) + + +def cudnn_data_type(cudnn, dtype): + """Convert a NumPy/JAX dtype to a cuDNN frontend data type.""" + dtype = jnp.dtype(dtype) + if dtype == jnp.float16: + return cudnn.data_type.HALF + if dtype == jnp.bfloat16: + return cudnn.data_type.BFLOAT16 + if dtype == jnp.float32: + return cudnn.data_type.FLOAT + if dtype == jnp.float64: + return cudnn.data_type.DOUBLE + if dtype == jnp.int32: + return cudnn.data_type.INT32 + if dtype == jnp.int64: + return cudnn.data_type.INT64 + if dtype == jnp.uint8: + return cudnn.data_type.UINT8 + if dtype == jnp.bool_: + return cudnn.data_type.BOOLEAN + raise ValueError(f"Unsupported cuDNN graph tensor dtype: {dtype}.") + + +def cudnn_data_type_from_name(cudnn, dtype_name_: str): + """Convert a serialized NumPy dtype name to a cuDNN frontend dtype.""" + if dtype_name_ == "bfloat16": + return cudnn.data_type.BFLOAT16 + return cudnn_data_type(cudnn, np.dtype(dtype_name_)) + + +def graph_tensor_from_aval(cudnn, graph, name: str, aval, uid: int): + """Create a contiguous graph tensor from a JAX abstract value.""" + shape = tuple(int(dim) for dim in aval.shape) + return graph.tensor( + name=name, + dim=shape, + stride=row_major_stride(shape), + data_type=cudnn_data_type(cudnn, aval.dtype), + uid=uid, + ) + + +def encode_cudnn_frontend_version(version: str) -> int: + """Encode a PEP-440 cuDNN frontend version as MMmmpp.""" + public_version = version.split("+", 1)[0].split("-", 1)[0] + parts = public_version.split(".") + if len(parts) < 3: + raise RuntimeError( + f"Could not parse cuDNN frontend Python version: {version!r}." + ) + major, minor, patch = (int(part) for part in parts[:3]) + return major * 10000 + minor * 100 + patch + + +def check_cudnn_frontend_version_match(cudnn) -> int: + """Ensure Python and C++ frontend versions use a compatible wire format.""" + python_version_string = getattr(cudnn, "__version__", None) + if python_version_string is None: + raise RuntimeError("cuDNN frontend Python package does not expose __version__.") + python_version = encode_cudnn_frontend_version(python_version_string) + cpp_version = int(transformer_engine_jax.get_cudnn_frontend_version()) + if python_version != cpp_version: + raise RuntimeError( + "cuDNN frontend Python/C++ version mismatch for graph serialization: " + f"Python cudnn.__version__={python_version_string!r} encodes to {python_version}, " + f"but Transformer Engine C++ was built with CUDNN_FRONTEND_VERSION={cpp_version}. " + "Use matching cuDNN frontend Python package and C++ headers." + ) + return python_version + + +def import_cudnn(): + """Import and validate the cuDNN frontend Python binding.""" + try: + cudnn = importlib.import_module("cudnn") + except ImportError as exc: + raise ImportError( + "JAX fused_attn requires the cuDNN frontend Python package (`cudnn`)." + ) from exc + check_cudnn_frontend_version_match(cudnn) + return cudnn + + +def graph_hash(serialized_graph: bytes) -> tuple[int, int]: + """Return two signed int64 values used as the C++ graph-cache key.""" + digest = hashlib.sha256(serialized_graph).digest() + return ( + int.from_bytes(digest[0:8], byteorder="little", signed=True), + int.from_bytes(digest[8:16], byteorder="little", signed=True), + ) + + +def pack_scalar_values(scalar_values: Sequence[bytes]) -> tuple[np.ndarray, np.ndarray]: + """Pack pass-by-value scalars into fixed, aligned 16-byte records.""" + scalar_sizes = np.asarray([len(value) for value in scalar_values], dtype=np.int64) + packed_values = np.zeros((len(scalar_values), 16), dtype=np.uint8) + for index, value in enumerate(scalar_values): + if len(value) > 16: + raise ValueError("cuDNN pass-by-value scalars must be at most 16 bytes.") + packed_values[index, : len(value)] = np.frombuffer(value, dtype=np.uint8) + return scalar_sizes, packed_values.reshape(-1) + + +def serialized_graph( + *, + serialized_graph_data: bytes, + cudnn_frontend_version: int, + workspace_size: int, + input_bindings: Sequence[GraphBinding], + output_bindings: Sequence[GraphBinding], + scalar_uids: Sequence[int] = (), + scalar_values: Sequence[bytes] = (), +) -> SerializedGraph: + """Construct normalized, NumPy-backed metadata for an FFI graph call.""" + scalar_sizes, packed_scalar_values = pack_scalar_values(scalar_values) + + def binding_array(bindings, field): + return np.asarray( + [getattr(binding, field) for binding in bindings], dtype=np.int64 + ) + + return SerializedGraph( + serialized_graph=serialized_graph_data, + graph_hash=graph_hash(serialized_graph_data), + cudnn_frontend_version=int(cudnn_frontend_version), + workspace_size=max(int(workspace_size), 1), + input_uids=binding_array(input_bindings, "uid"), + input_buffer_indices=binding_array(input_bindings, "buffer_index"), + input_byte_offsets=binding_array(input_bindings, "byte_offset"), + output_uids=binding_array(output_bindings, "uid"), + output_buffer_indices=binding_array(output_bindings, "buffer_index"), + output_byte_offsets=binding_array(output_bindings, "byte_offset"), + scalar_uids=np.asarray(scalar_uids, dtype=np.int64), + scalar_sizes=scalar_sizes, + scalar_values=packed_scalar_values, + ) + + +def finalize_graph(cudnn, graph, *, description: str) -> tuple[int, bytes, int]: + """Validate, plan and serialize a cuDNN frontend graph.""" + graph.validate() + graph.build_operation_graph() + try: + graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError( + f"cuDNN {description} graph is not supported: {exc}" + ) from exc + graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) + return ( + max(int(graph.get_workspace_size()), 1), + bytes(graph.serialize()), + check_cudnn_frontend_version_match(cudnn), + ) + + +def shape_dtype(value) -> jax.ShapeDtypeStruct: + """Return a hashable-enough static shape/dtype descriptor for graph construction.""" + return jax.ShapeDtypeStruct(tuple(value.shape), value.dtype) diff --git a/transformer_engine/jax/cpp_extensions/flex_attention.py b/transformer_engine/jax/cpp_extensions/flex_attention.py index ff9f5fb1460..2ddc9afef14 100644 --- a/transformer_engine/jax/cpp_extensions/flex_attention.py +++ b/transformer_engine/jax/cpp_extensions/flex_attention.py @@ -3,8 +3,6 @@ # See LICENSE for license information. """cuDNN frontend score_mod fused attention helpers.""" -import hashlib -import importlib import inspect import os from dataclasses import dataclass @@ -15,7 +13,36 @@ import numpy as np from jax import ffi -import transformer_engine_jax +from .cudnn_graph import ( + GraphBinding, + SerializedGraph, + finalize_graph, + import_cudnn, +) +from .cudnn_graph import ( + bshd_as_bhsd_dim_stride as _bshd_as_bhsd_dim_stride, +) +from .cudnn_graph import ( + cudnn_data_type as _cudnn_data_type, +) +from .cudnn_graph import ( + cudnn_data_type_from_name as _cudnn_data_type_from_name, +) +from .cudnn_graph import ( + dtype_name as _dtype_name, +) +from .cudnn_graph import ( + graph_tensor_from_aval as _graph_tensor_from_aval, +) +from .cudnn_graph import ( + row_major_stride as _row_major_stride, +) +from .cudnn_graph import ( + serialized_graph as make_serialized_graph, +) +from .cudnn_graph import ( + shape_dtype as _shape_dtype, +) __all__ = [ "FusedAttnScoreModHelper", @@ -254,19 +281,7 @@ def __eq__(self, other): ) -@dataclass(frozen=True) -class _SerializedScoreModGraph: - """Serialized cuDNN frontend graph and static metadata for C++ execution.""" - - serialized_graph: bytes - graph_hash: Tuple[int, int] - cudnn_frontend_version: int - workspace_size: int - input_uids: np.ndarray - output_uids: np.ndarray - scalar_uids: np.ndarray - scalar_sizes: np.ndarray - scalar_values: np.ndarray +_SerializedScoreModGraph = SerializedGraph # cuDNN frontend tensor UIDs are arbitrary, but assigning stable values makes serialized @@ -288,29 +303,6 @@ class _SerializedScoreModGraph: _score_mod_graph_cache: Dict[Tuple[Any, ...], _SerializedScoreModGraph] = {} -def _row_major_stride(shape: Sequence[int]) -> Tuple[int, ...]: - stride = [] - running = 1 - for dim in reversed(tuple(shape)): - stride.append(running) - running *= dim - return tuple(reversed(stride)) - - -def _bshd_as_bhsd_dim_stride(shape: Sequence[int]) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: - if len(shape) != 4: - raise ValueError(f"score_mod requires rank-4 BSHD tensors, got shape={shape}.") - batch, seqlen, heads, head_dim = tuple(shape) - return ( - (batch, heads, seqlen, head_dim), - (seqlen * heads * head_dim, head_dim, heads * head_dim, 1), - ) - - -def _dtype_name(dtype) -> str: - return str(jnp.dtype(dtype)) - - def _is_array_operand(value: Any) -> bool: return ( hasattr(value, "shape") @@ -405,44 +397,6 @@ def _make_fused_attn_score_mod_config( return config, tensor_operands, bprop_tensor_operands -def _cudnn_data_type(cudnn, dtype): - dtype = jnp.dtype(dtype) - if dtype == jnp.float16: - return cudnn.data_type.HALF - if dtype == jnp.bfloat16: - return cudnn.data_type.BFLOAT16 - if dtype == jnp.float32: - return cudnn.data_type.FLOAT - if dtype == jnp.float64: - return cudnn.data_type.DOUBLE - if dtype == jnp.int32: - return cudnn.data_type.INT32 - if dtype == jnp.int64: - return cudnn.data_type.INT64 - if dtype == jnp.uint8: - return cudnn.data_type.UINT8 - if dtype == jnp.bool_: - return cudnn.data_type.BOOLEAN - raise ValueError(f"Unsupported score_mod tensor dtype: {dtype}.") - - -def _cudnn_data_type_from_name(cudnn, dtype_name: str): - if dtype_name == "bfloat16": - return cudnn.data_type.BFLOAT16 - return _cudnn_data_type(cudnn, np.dtype(dtype_name)) - - -def _graph_tensor_from_aval(cudnn, graph, name: str, aval, uid: int): - shape = tuple(int(dim) for dim in aval.shape) - return graph.tensor( - name=name, - dim=shape, - stride=_row_major_stride(shape), - data_type=_cudnn_data_type(cudnn, aval.dtype), - uid=uid, - ) - - def _score_mod_graph_tensors( cudnn, graph, @@ -477,51 +431,6 @@ def _score_mod_graph_tensors( return graph_tensors, tuple(tensor_uids), tuple(scalar_uids), tuple(scalar_values) -def _encode_cudnn_frontend_version(version: str) -> int: - public_version = version.split("+", 1)[0].split("-", 1)[0] - parts = public_version.split(".") - if len(parts) < 3: - raise RuntimeError(f"Could not parse cuDNN frontend Python version: {version!r}.") - major, minor, patch = (int(part) for part in parts[:3]) - return major * 10000 + minor * 100 + patch - - -def _check_cudnn_frontend_version_match(cudnn) -> int: - python_version_string = getattr(cudnn, "__version__", None) - if python_version_string is None: - raise RuntimeError("cuDNN frontend Python package does not expose __version__.") - python_version = _encode_cudnn_frontend_version(python_version_string) - cpp_version = int(transformer_engine_jax.get_cudnn_frontend_version()) - if python_version != cpp_version: - raise RuntimeError( - "cuDNN frontend Python/C++ version mismatch for score_mod graph serialization: " - f"Python cudnn.__version__={python_version_string!r} encodes to {python_version}, " - f"but Transformer Engine C++ was built with CUDNN_FRONTEND_VERSION={cpp_version}. " - "Use matching cuDNN frontend Python package and C++ headers." - ) - return python_version - - -def _score_mod_graph_hash(serialized_graph: bytes) -> Tuple[int, int]: - digest = hashlib.sha256(serialized_graph).digest() - return ( - int.from_bytes(digest[0:8], byteorder="little", signed=True), - int.from_bytes(digest[8:16], byteorder="little", signed=True), - ) - - -def _pack_score_mod_scalar_values( - scalar_values: Sequence[bytes], -) -> Tuple[np.ndarray, np.ndarray]: - scalar_sizes = np.asarray([len(value) for value in scalar_values], dtype=np.int64) - packed_values = np.zeros((len(scalar_values), 16), dtype=np.uint8) - for index, value in enumerate(scalar_values): - if len(value) > 16: - raise ValueError("score_mod pass-by-value scalars must be at most 16 bytes.") - packed_values[index, : len(value)] = np.frombuffer(value, dtype=np.uint8) - return scalar_sizes, packed_values.reshape(-1) - - def _serialized_score_mod_graph( *, serialized_graph: bytes, @@ -532,17 +441,20 @@ def _serialized_score_mod_graph( scalar_uids: Sequence[int], scalar_values: Sequence[bytes], ) -> _SerializedScoreModGraph: - scalar_sizes, packed_scalar_values = _pack_score_mod_scalar_values(scalar_values) - return _SerializedScoreModGraph( - serialized_graph=serialized_graph, - graph_hash=_score_mod_graph_hash(serialized_graph), + return make_serialized_graph( + serialized_graph_data=serialized_graph, cudnn_frontend_version=int(cudnn_frontend_version), workspace_size=int(workspace_size), - input_uids=np.asarray(input_uids, dtype=np.int64), - output_uids=np.asarray(output_uids, dtype=np.int64), + input_bindings=[ + GraphBinding(uid=int(uid), buffer_index=index) + for index, uid in enumerate(input_uids) + ], + output_bindings=[ + GraphBinding(uid=int(uid), buffer_index=index) + for index, uid in enumerate(output_uids) + ], scalar_uids=np.asarray(scalar_uids, dtype=np.int64), - scalar_sizes=scalar_sizes, - scalar_values=packed_scalar_values, + scalar_values=scalar_values, ) @@ -557,20 +469,7 @@ def wrapped_score_mod(sdpa_graph, score_tensor): def _finalize_score_mod_graph(cudnn, graph) -> Tuple[int, bytes, int]: - graph.validate() - graph.build_operation_graph() - try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN score_mod SDPA graph is not supported: {exc}") from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) - serialized_graph = bytes(graph.serialize()) - return ( - max(int(graph.get_workspace_size()), 1), - serialized_graph, - _check_cudnn_frontend_version_match(cudnn), - ) + return finalize_graph(cudnn, graph, description="score_mod SDPA") def _graph_cache_key( @@ -589,19 +488,8 @@ def _graph_cache_key( ) -def _shape_dtype(value) -> jax.ShapeDtypeStruct: - return jax.ShapeDtypeStruct(tuple(value.shape), value.dtype) - - def _import_cudnn_for_score_mod(): - try: - cudnn = importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "score_mod fused_attn requires the cuDNN frontend Python package (`cudnn`)." - ) from exc - _check_cudnn_frontend_version_match(cudnn) - return cudnn + return import_cudnn() def _build_score_mod_fwd_graph(q_aval, k_aval, v_aval, score_mod_avals, config): @@ -820,15 +708,7 @@ def _fused_attn_score_mod_fwd( k, v, *score_mod_tensors, - serialized_graph=graph.serialized_graph, - graph_hash0=graph.graph_hash[0], - graph_hash1=graph.graph_hash[1], - cudnn_frontend_version=graph.cudnn_frontend_version, - input_uids=graph.input_uids, - output_uids=graph.output_uids, - scalar_uids=graph.scalar_uids, - scalar_sizes=graph.scalar_sizes, - scalar_values=graph.scalar_values, + **graph.ffi_attrs(), ) return output, softmax_stats @@ -884,15 +764,7 @@ def _fused_attn_score_mod_bwd( softmax_stats, *score_mod_tensors, *score_mod_bprop_tensors, - serialized_graph=graph.serialized_graph, - graph_hash0=graph.graph_hash[0], - graph_hash1=graph.graph_hash[1], - cudnn_frontend_version=graph.cudnn_frontend_version, - input_uids=graph.input_uids, - output_uids=graph.output_uids, - scalar_uids=graph.scalar_uids, - scalar_sizes=graph.scalar_sizes, - scalar_values=graph.scalar_values, + **graph.ffi_attrs(), ) return dq, dk, dv diff --git a/transformer_engine/jax/csrc/extensions.h b/transformer_engine/jax/csrc/extensions.h index 580219baf20..c3e80b4202b 100644 --- a/transformer_engine/jax/csrc/extensions.h +++ b/transformer_engine/jax/csrc/extensions.h @@ -151,28 +151,8 @@ XLA_FFI_DECLARE_HANDLER_SYMBOL(FusedAttnScoreModForwardHandler); XLA_FFI_DECLARE_HANDLER_SYMBOL(FusedAttnScoreModBackwardHandler); -NVTE_Fused_Attn_Backend GetFusedAttnBackend( - bool is_training, DType q_dtype, DType kv_dtype, NVTE_QKV_Layout qkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - float dropout_probability, size_t q_attn_heads, size_t kv_attn_heads, size_t q_max_seqlen, - size_t kv_max_seqlen, size_t qk_head_dim, size_t v_head_dim, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool deterministic); - -pybind11::tuple GetFusedAttnForwardWorkspaceSizes( - size_t input_batch, size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, - size_t attn_heads, size_t num_gqa_groups, size_t bias_heads, size_t qk_head_dim, - size_t v_head_dim, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, - DType dtype, bool is_training, size_t max_segments_per_seq, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool bottom_right_diagonal); - -pybind11::tuple GetFusedAttnBackwardWorkspaceSizes( - size_t input_batch, size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, - size_t attn_heads, size_t num_gqa_groups, size_t bias_heads, size_t qk_head_dim, - size_t v_head_dim, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, - DType dtype, bool is_training, bool deterministic, size_t max_segments_per_seq, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal); +void PopulateFusedAttnRngState(void *rng_state, const void *seed, uint64_t offset, + cudaStream_t stream); // GEMM XLA_FFI_DECLARE_HANDLER_SYMBOL(GemmHandler); diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 41f87ae71bc..bf2fc542f0d 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -18,739 +18,31 @@ #include #include "../extensions.h" -#include "transformer_engine/fused_attn.h" -#include "transformer_engine/transformer_engine.h" namespace transformer_engine { namespace jax { - -NVTE_Fused_Attn_Backend GetFusedAttnBackend( - bool is_training, DType q_dtype, DType kv_dtype, NVTE_QKV_Layout qkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - float dropout_probability, size_t q_attn_heads, size_t kv_attn_heads, size_t q_max_seqlen, - size_t kv_max_seqlen, size_t qk_head_dim, size_t v_head_dim, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool deterministic) { - auto backend = nvte_get_fused_attn_backend( - is_training, static_cast(q_dtype), static_cast(kv_dtype), qkv_layout, - bias_type, mask_type, softmax_type, dropout_probability, q_attn_heads, kv_attn_heads, - q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, - return_max_logit, false, deterministic); - return backend; -} - -/* - NOTE: PrepareFusedAttnForwardAuxTensors unifies the auxiliary tensor pack logic from the fused - attention forward kernels in: - - common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu lines 1270-1281 and 1348-1359 -*/ -void PrepareFusedAttnForwardAuxTensors(NVTETensorPack *tensor_pack, const size_t input_batch, - const size_t bias_batch, const size_t attn_heads, - const size_t bias_heads, const size_t q_max_seqlen, - const size_t kv_max_seqlen, DType dtype, - NVTE_Bias_Type bias_type, NVTE_Fused_Attn_Backend backend, - void *softmax_buf, void *max_logits_buf = nullptr, - void *rng_state_buf = nullptr, void *bias_buf = nullptr, - void *softmax_offset_buf = nullptr) { - // all backends need softmax but expect different shapes/dtypes - tensor_pack->size = 1; - NVTETensor &softmax_aux = tensor_pack->tensors[0]; - NVTEBasicTensor softmax_aux_data; - softmax_aux_data.data_ptr = softmax_buf; - softmax_aux_data.shape.ndim = 4; - softmax_aux_data.shape.data[0] = input_batch; - softmax_aux_data.shape.data[1] = attn_heads; - softmax_aux_data.shape.data[2] = q_max_seqlen; - softmax_aux_data.shape.data[3] = kv_max_seqlen; - softmax_aux_data.dtype = static_cast(dtype); - - // arbitrary sequence length backend needs the RNG state and a different shape/dtype softmax - if (backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) { - int size = 1; // Start after softmax. - auto next_aux_tensor = [&]() -> NVTETensor & { - NVTE_CHECK(size < NVTETensorPack::MAX_SIZE, - "Fused attention auxiliary tensor pack capacity exceeded."); - return tensor_pack->tensors[size++]; - }; - - if (max_logits_buf != nullptr) { - NVTETensor &max_aux = next_aux_tensor(); - NVTEBasicTensor max_aux_data; - max_aux_data.data_ptr = max_logits_buf; - max_aux_data.shape = {}; - max_aux_data.shape.ndim = 4; - max_aux_data.shape.data[0] = input_batch; - max_aux_data.shape.data[1] = attn_heads; - max_aux_data.shape.data[2] = q_max_seqlen; - max_aux_data.shape.data[3] = 1; - max_aux_data.dtype = static_cast(DType::kFloat32); - nvte_set_tensor_param(&max_aux, kNVTERowwiseData, &max_aux_data); - } - - NVTETensor &rng_state_aux = next_aux_tensor(); - NVTEBasicTensor rng_state_aux_data; - rng_state_aux_data.data_ptr = rng_state_buf; - rng_state_aux_data.shape = {}; - rng_state_aux_data.shape.ndim = 2; - rng_state_aux_data.dtype = static_cast(DType::kInt64); - nvte_set_tensor_param(&rng_state_aux, kNVTERowwiseData, &rng_state_aux_data); - // correct softmax shape/dtype - softmax_aux_data.shape.data[3] = 1; // {B,H,Qs,Ks} -> {B,H,Qs,1} - softmax_aux_data.dtype = static_cast(DType::kFloat32); - - // include bias if enabled - if (bias_type != NVTE_Bias_Type::NVTE_NO_BIAS && bias_type != NVTE_Bias_Type::NVTE_ALIBI) { - NVTETensor &bias_aux = next_aux_tensor(); - NVTEBasicTensor bias_aux_data; - bias_aux_data.data_ptr = bias_buf; - bias_aux_data.shape.ndim = 4; - bias_aux_data.shape.data[0] = bias_batch; - bias_aux_data.shape.data[1] = bias_heads; - bias_aux_data.shape.data[2] = q_max_seqlen; - bias_aux_data.shape.data[3] = kv_max_seqlen; - bias_aux_data.dtype = static_cast(dtype); - nvte_set_tensor_param(&bias_aux, kNVTERowwiseData, &bias_aux_data); - } - - // include softmax_offset if provided - if (softmax_offset_buf != nullptr) { - NVTETensor &softmax_offset_aux = next_aux_tensor(); - NVTEBasicTensor softmax_offset_aux_data; - softmax_offset_aux_data.data_ptr = softmax_offset_buf; - softmax_offset_aux_data.shape.ndim = 4; - softmax_offset_aux_data.shape.data[0] = 1; - softmax_offset_aux_data.shape.data[1] = attn_heads; - softmax_offset_aux_data.shape.data[2] = 1; - softmax_offset_aux_data.shape.data[3] = 1; - softmax_offset_aux_data.dtype = static_cast(DType::kFloat32); - nvte_set_tensor_param(&softmax_offset_aux, kNVTERowwiseData, &softmax_offset_aux_data); - } - - // Set final size - tensor_pack->size = size; - } - nvte_set_tensor_param(&softmax_aux, kNVTERowwiseData, &softmax_aux_data); -} - -/* - NOTE: Backward fused attention kernels accept auxiliary tensors as explicit function arguments - instead of an NVTETensorPack and nvte_fused_attn_bwd() API does all the logic for pulling the - necessary tensors out of the tensor pack for the active kernel. That means we can just dump - everything we got into the tensor pack and not worry about its sizing for the backward pass. - - TODO(Alp): Refactor the nvte_fused_attn_fwd() to work like nvte_fused_attn_bwd()? -*/ -void PrepareFusedAttnBackwardAuxTensors(NVTETensorPack *tensor_pack, const size_t input_batch, - const size_t bias_batch, const size_t attn_heads, - const size_t bias_heads, const size_t q_max_seqlen, - const size_t kv_max_seqlen, DType dtype, - NVTE_Fused_Attn_Backend backend, void *softmax_buf, - void *rng_state_buf, void *bias_buf, - void *softmax_offset_buf = nullptr) { - // Backward calls put everything into the tensor pack for every backend - // so we set dummy bias_type and backend choices here to follow the correct code path - auto dummy_bias_type = NVTE_Bias_Type::NVTE_POST_SCALE_BIAS; - auto dummy_backend = NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen; - PrepareFusedAttnForwardAuxTensors(tensor_pack, input_batch, bias_batch, attn_heads, bias_heads, - q_max_seqlen, kv_max_seqlen, dtype, dummy_bias_type, - dummy_backend, softmax_buf, nullptr, rng_state_buf, bias_buf, - softmax_offset_buf); -} - -pybind11::tuple GetFusedAttnForwardWorkspaceSizes( - size_t input_batch, size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, - size_t attn_heads, size_t num_gqa_groups, size_t bias_heads, size_t qk_head_dim, - size_t v_head_dim, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, - DType dtype, bool is_training, size_t max_segments_per_seq, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool bottom_right_diagonal) { - auto is_ragged = nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD; - auto q_shape = is_ragged - ? std::vector{input_batch * q_max_seqlen, attn_heads, qk_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, qk_head_dim}; - auto q_tensor = TensorWrapper(nullptr, q_shape, dtype); - auto k_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, qk_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, qk_head_dim}; - auto k_tensor = TensorWrapper(nullptr, k_shape, dtype); - auto v_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, v_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, v_head_dim}; - auto v_tensor = TensorWrapper(nullptr, v_shape, dtype); - auto o_shape = is_ragged ? std::vector{input_batch * q_max_seqlen, attn_heads, v_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, v_head_dim}; - - auto bias_shape = std::vector{bias_batch, bias_heads, q_max_seqlen, kv_max_seqlen}; - auto bias_tensor = TensorWrapper(nullptr, bias_shape, dtype); - - // F16 doesn't use this tensor - auto s_tensor = TensorWrapper(nullptr, std::vector{1}, dtype); - auto o_tensor = TensorWrapper(nullptr, o_shape, dtype); - - auto dummy_rng_state_tensor = TensorWrapper(nullptr, std::vector{2}, DType::kInt64); - auto dummy_page_table_tensor = TensorWrapper(nullptr, std::vector{1}, DType::kInt32); - auto dummy_softmax_offset_tensor = - TensorWrapper(nullptr, std::vector{1}, DType::kFloat32); - - NVTETensorPack aux_output_tensors; - nvte_tensor_pack_create(&aux_output_tensors); - - TensorWrapper query_workspace_tensor; - // It is a WAR to pre-create all possible cuDNN graph at the JIT compile time - size_t max_num_segments = is_ragged ? input_batch * max_segments_per_seq : input_batch; - size_t min_num_segments = input_batch; - auto cudnn_runtime_version = cudnnGetVersion(); - if (is_ragged && cudnn_runtime_version >= 90300) { - // For cuDNN < 9.3.0, it requires to run all possible seqlens to address act_seqlen = 0 - min_num_segments = input_batch * max_segments_per_seq; - } - for (auto num_segments = min_num_segments; num_segments <= max_num_segments; ++num_segments) { - // the last one is the largest which will be the returned workspace size - auto q_cu_seqlens_tensor = - TensorWrapper(nullptr, std::vector{num_segments + 1}, DType::kInt32); - auto kv_cu_seqlens_tensor = - TensorWrapper(nullptr, std::vector{num_segments + 1}, DType::kInt32); - auto ragged_offset_tensor = - TensorWrapper(nullptr, std::vector{num_segments + 1}, DType::kInt32); - nvte_fused_attn_fwd( - q_tensor.data(), k_tensor.data(), v_tensor.data(), bias_tensor.data(), - dummy_softmax_offset_tensor.data(), s_tensor.data(), o_tensor.data(), &aux_output_tensors, - q_cu_seqlens_tensor.data(), kv_cu_seqlens_tensor.data(), ragged_offset_tensor.data(), - ragged_offset_tensor.data(), dummy_page_table_tensor.data(), dummy_page_table_tensor.data(), - dummy_rng_state_tensor.data(), q_max_seqlen, kv_max_seqlen, is_training, return_max_logit, - false, scaling_factor, dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), - NVTE_QKV_Format_NOT_SET, bias_type, mask_type, softmax_type, window_size_left, - window_size_right, bottom_right_diagonal, query_workspace_tensor.data(), nullptr); - } - - nvte_tensor_pack_destroy(&aux_output_tensors); - - auto workspace_shape = MakeShapeVector(query_workspace_tensor.shape()); - return pybind11::make_tuple(workspace_shape, query_workspace_tensor.dtype()); -} - -#define FUSED_ATTN_IMPL_COMMON_BLOCK \ - auto is_ragged = nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD; \ - auto bias_shape = std::vector{bias_batch, bias_heads, q_max_seqlen, kv_max_seqlen}; \ - size_t num_segments = input_batch; \ - if (is_ragged) { \ - auto cudnn_runtime_version = cudnnGetVersion(); \ - if (cudnn_runtime_version >= 90300) { \ - num_segments = input_batch * max_segments_per_seq; \ - } else { \ - size_t runtime_num_segments_q = nvte_get_runtime_num_segments( \ - q_cu_seqlens, workspace, input_batch * q_max_seqlen, stream); \ - size_t runtime_num_segments_kv = nvte_get_runtime_num_segments( \ - kv_cu_seqlens, workspace, input_batch * kv_max_seqlen, stream); \ - NVTE_CHECK(runtime_num_segments_q == runtime_num_segments_kv); \ - NVTE_CHECK(runtime_num_segments_q <= input_batch * max_segments_per_seq); \ - num_segments = runtime_num_segments_q; \ - } \ - } \ - std::vector seq_shape{num_segments + 1}; \ - auto q_cu_seqlens_tensor = TensorWrapper(q_cu_seqlens, seq_shape, DType::kInt32); \ - auto kv_cu_seqlens_tensor = TensorWrapper(kv_cu_seqlens, seq_shape, DType::kInt32); \ - auto q_seq_offsets_tensor = TensorWrapper(q_seq_offsets, seq_shape, DType::kInt32); \ - auto k_seq_offsets_tensor = TensorWrapper(k_seq_offsets, seq_shape, DType::kInt32); \ - auto workspace_tensor = \ - TensorWrapper(workspace, std::vector{wkspace_size}, wkspace_dtype); \ - auto layout_group = nvte_get_qkv_layout_group(qkv_layout); - -static void FusedAttnForwardImpl( - cudaStream_t stream, void *q, void *k, void *v, void *bias, void *softmax_offset, void *seed, - void *q_cu_seqlens, void *kv_cu_seqlens, void *q_seq_offsets, void *k_seq_offsets, void *output, - void *softmax_aux, void *max_tensor, void *rng_state, void *workspace, size_t input_batch, - size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, size_t attn_heads, - size_t num_gqa_groups, size_t bias_heads, size_t qk_head_dim, size_t v_head_dim, - size_t max_segments_per_seq, size_t wkspace_size, float scaling_factor, - float dropout_probability, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, - NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, DType dtype, DType wkspace_dtype, - bool is_training, bool return_max_logit, bool deterministic, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal) { - FUSED_ATTN_IMPL_COMMON_BLOCK; - - /* Input tensors */ - auto bias_tensor = TensorWrapper(bias, bias_shape, dtype); - auto softmax_offset_tensor = - TensorWrapper(softmax_offset, std::vector{1, attn_heads, 1, 1}, DType::kFloat32); - - if (is_ragged) { - auto output_size = input_batch * q_max_seqlen * attn_heads * v_head_dim; - cudaMemsetAsync(output, 0, output_size * typeToSize(dtype), stream); - - // Memset to 0xF0 for filling large negative numbers - auto softmax_aux_size = input_batch * q_max_seqlen * attn_heads; - cudaMemsetAsync(softmax_aux, 0xF0, softmax_aux_size * sizeof(float), stream); - if (return_max_logit) { - cudaMemsetAsync(max_tensor, 0xF0, softmax_aux_size * sizeof(float), stream); - } - } - - /* Output tensors */ - auto s_tensor = TensorWrapper(nullptr, std::vector{1}, dtype); // not used in F16 - auto o_shape = is_ragged ? std::vector{input_batch * q_max_seqlen, attn_heads, v_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, v_head_dim}; - auto o_tensor = TensorWrapper(output, o_shape, dtype); - - /* Prepare RNG state */ - auto rng_state_tensor = TensorWrapper(rng_state, std::vector{2}, DType::kInt64); - - auto backend = nvte_get_fused_attn_backend( - is_training, static_cast(dtype), static_cast(dtype), qkv_layout, - bias_type, mask_type, softmax_type, dropout_probability, attn_heads, num_gqa_groups, - q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, - return_max_logit, false, deterministic); - nvte_populate_rng_state_async(rng_state, seed, q_max_seqlen, kv_max_seqlen, backend, stream); - - /* Auxiliary tensors (to be propagated to the backward pass later) */ - NVTETensorPack aux_output_tensors; - nvte_tensor_pack_create(&aux_output_tensors); - PrepareFusedAttnForwardAuxTensors(&aux_output_tensors, input_batch, bias_batch, attn_heads, - bias_heads, q_max_seqlen, kv_max_seqlen, dtype, bias_type, - backend, softmax_aux, return_max_logit ? max_tensor : nullptr, - rng_state, bias, softmax_offset); - - /* Call the underlying NVTE API */ - auto dummy_page_table_tensor = TensorWrapper(nullptr, std::vector{1}, DType::kInt32); - - // Prepare Q, K, V pointers and shapes based on layout - // Python passes dummy tensors for unused slots, so we extract from the actual packed data - void *q_ptr = q; - void *k_ptr = k; - void *v_ptr = v; - auto q_shape = is_ragged - ? std::vector{input_batch * q_max_seqlen, attn_heads, qk_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, qk_head_dim}; - auto k_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, qk_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, qk_head_dim}; - auto v_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, v_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, v_head_dim}; - - if (layout_group == NVTE_QKV_Layout_Group::NVTE_3HD) { - // QKV packed in q: [batch*seqlen, 3, heads, dim] - // Python passes: q=packed_qkv, k=dummy, v=dummy - // Extract K and V pointers from the packed q data - NVTE_CHECK(q_max_seqlen == kv_max_seqlen, "q_max_seqlen must equal kv_max_seqlen"); - NVTE_CHECK(qk_head_dim == v_head_dim, - "For QKV packed layout, qk_head_dim must equal v_head_dim"); - size_t stride = (typeToSize(dtype) * attn_heads * qk_head_dim); - q_ptr = q; - k_ptr = static_cast(static_cast(q) + stride); - v_ptr = static_cast(static_cast(q) + 2 * stride); - // For packed QKV, all have same shape since they're views into the same packed tensor - k_shape = q_shape; - v_shape = q_shape; - } else if (layout_group == NVTE_QKV_Layout_Group::NVTE_HD_2HD) { - // Q separate, KV packed in k: [batch*seqlen, 2, num_gqa_groups, dim] - // Python passes: q=query, k=packed_kv, v=dummy - // Extract V pointer from the packed k data - NVTE_CHECK(qk_head_dim == v_head_dim, - "For KV packed layout, qk_head_dim must equal v_head_dim"); - size_t stride = (typeToSize(dtype) * num_gqa_groups * qk_head_dim); - q_ptr = q; - k_ptr = k; - v_ptr = static_cast(static_cast(k) + stride); - // V has same shape as K since they're packed together - v_shape = k_shape; - } - // else NVTE_HD_HD_HD: pointers and shapes already correct - - auto q_tensor = TensorWrapper(q_ptr, q_shape, dtype); - auto k_tensor = TensorWrapper(k_ptr, k_shape, dtype); - auto v_tensor = TensorWrapper(v_ptr, v_shape, dtype); - - nvte_fused_attn_fwd( - q_tensor.data(), k_tensor.data(), v_tensor.data(), bias_tensor.data(), - softmax_offset_tensor.data(), s_tensor.data(), o_tensor.data(), &aux_output_tensors, - q_cu_seqlens_tensor.data(), kv_cu_seqlens_tensor.data(), q_seq_offsets_tensor.data(), - k_seq_offsets_tensor.data(), dummy_page_table_tensor.data(), dummy_page_table_tensor.data(), - rng_state_tensor.data(), q_max_seqlen, kv_max_seqlen, is_training, return_max_logit, false, - scaling_factor, dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), - NVTE_QKV_Format_NOT_SET, bias_type, mask_type, softmax_type, window_size_left, - window_size_right, bottom_right_diagonal, workspace_tensor.data(), stream); - - nvte_tensor_pack_destroy(&aux_output_tensors); -} - -#define FUSED_ATTN_FFI_GET_ATTRS \ - size_t input_batch = get_attr_value(attrs, "input_batch"); \ - size_t bias_batch = get_attr_value(attrs, "bias_batch"); \ - size_t q_max_seqlen = get_attr_value(attrs, "q_max_seqlen"); \ - size_t kv_max_seqlen = get_attr_value(attrs, "kv_max_seqlen"); \ - size_t attn_heads = get_attr_value(attrs, "attn_heads"); \ - size_t num_gqa_groups = get_attr_value(attrs, "num_gqa_groups"); \ - size_t bias_heads = get_attr_value(attrs, "bias_heads"); \ - size_t qk_head_dim = get_attr_value(attrs, "qk_head_dim"); \ - size_t v_head_dim = get_attr_value(attrs, "v_head_dim"); \ - size_t max_segments_per_seq = get_attr_value(attrs, "max_segments_per_seq"); \ - auto window_size_left = get_attr_value(attrs, "window_size_left"); \ - auto window_size_right = get_attr_value(attrs, "window_size_right"); \ - bool bottom_right_diagonal = get_attr_value(attrs, "bottom_right_diagonal"); \ - float scaling_factor = get_attr_value(attrs, "scaling_factor"); \ - float dropout_probability = get_attr_value(attrs, "dropout_probability"); \ - NVTE_Bias_Type bias_type = \ - static_cast(get_attr_value(attrs, "bias_type")); \ - NVTE_Mask_Type mask_type = \ - static_cast(get_attr_value(attrs, "mask_type")); \ - NVTE_Softmax_Type softmax_type = \ - static_cast(get_attr_value_or_default( \ - attrs, "softmax_type", static_cast(NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX))); \ - NVTE_QKV_Layout qkv_layout = \ - static_cast(get_attr_value(attrs, "qkv_layout")); \ - bool is_training = get_attr_value(attrs, "is_training"); \ - bool return_max_logit = get_attr_value_or_default(attrs, "return_max_logit", false); \ - bool deterministic = get_attr_value(attrs, "deterministic"); \ - auto is_ragged = nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD; \ - size_t wkspace_size = product(workspace_buf->dimensions()); \ - DType dtype = convert_ffi_datatype_to_te_dtype(q_buf.element_type()); \ - DType wkspace_dtype = convert_ffi_datatype_to_te_dtype(workspace_buf->element_type()); - -Error_Type FusedAttnForwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, - Buffer_Type v_buf, Buffer_Type bias_buf, - Buffer_Type softmax_offset_buf, Buffer_Type seed_buf, - Buffer_Type q_cu_seqlens_buf, Buffer_Type kv_cu_seqlens_buf, - Buffer_Type q_seq_offsets_buf, Buffer_Type k_seq_offsets_buf, - Variadic_Buffer_Type _unused_args, Result_Type output_buf, - Result_Type softmax_aux_buf, Result_Type max_tensor_buf, - Result_Type rng_state_buf, Result_Type workspace_buf, - Dictionary attrs) { - FUSED_ATTN_FFI_GET_ATTRS; - - FusedAttnForwardImpl( - stream, q_buf.untyped_data(), k_buf.untyped_data(), v_buf.untyped_data(), - bias_buf.untyped_data(), softmax_offset_buf.untyped_data(), seed_buf.untyped_data(), - q_cu_seqlens_buf.untyped_data(), kv_cu_seqlens_buf.untyped_data(), - is_ragged ? q_seq_offsets_buf.untyped_data() : nullptr, - is_ragged ? k_seq_offsets_buf.untyped_data() : nullptr, output_buf->untyped_data(), - softmax_aux_buf->untyped_data(), max_tensor_buf->untyped_data(), - rng_state_buf->untyped_data(), workspace_buf->untyped_data(), input_batch, bias_batch, - q_max_seqlen, kv_max_seqlen, attn_heads, num_gqa_groups, bias_heads, qk_head_dim, v_head_dim, - max_segments_per_seq, wkspace_size, scaling_factor, dropout_probability, bias_type, mask_type, - softmax_type, qkv_layout, dtype, wkspace_dtype, is_training, return_max_logit, deterministic, - window_size_left, window_size_right, bottom_right_diagonal); - return ffi_with_cuda_error_check(); -} - -XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnForwardHandler, FusedAttnForwardFFI, - FFI::Bind() - .Ctx() // stream - .Arg() // q - .Arg() // k - .Arg() // v - .Arg() // bias - .Arg() // softmax_offset - .Arg() // seed_buf - .Arg() // q_cu_seqlens - .Arg() // kv_cu_seqlens - .Arg() // q_seq_offsets - .Arg() // k_seq_offsets - .RemainingArgs() // _cp_aux_args unused - .Ret() // output - .Ret() // softmax_aux - .Ret() // max_tensor - .Ret() // rng_state - .Ret() // workspace - .Attrs(), - FFI_CudaGraph_Traits); - -pybind11::tuple GetFusedAttnBackwardWorkspaceSizes( - size_t input_batch, size_t bias_batch, size_t q_max_seqlen, size_t kv_max_seqlen, - size_t attn_heads, size_t num_gqa_groups, size_t bias_heads, size_t qk_head_dim, - size_t v_head_dim, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, - DType dtype, bool is_training, bool deterministic, size_t max_segments_per_seq, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal) { - auto is_ragged = nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD; - auto q_shape = is_ragged - ? std::vector{input_batch * q_max_seqlen, attn_heads, qk_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, qk_head_dim}; - auto q_tensor = TensorWrapper(nullptr, q_shape, dtype); - auto dq_tensor = TensorWrapper(nullptr, q_shape, dtype); - auto k_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, qk_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, qk_head_dim}; - auto k_tensor = TensorWrapper(nullptr, k_shape, dtype); - auto dk_tensor = TensorWrapper(nullptr, k_shape, dtype); - auto v_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, v_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, v_head_dim}; - auto v_tensor = TensorWrapper(nullptr, v_shape, dtype); - auto dv_tensor = TensorWrapper(nullptr, v_shape, dtype); - - auto output_shape = is_ragged - ? std::vector{input_batch * q_max_seqlen, attn_heads, v_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, v_head_dim}; - auto doutput_tensor = TensorWrapper(nullptr, output_shape, dtype); - auto output_tensor = TensorWrapper(nullptr, output_shape, dtype); - - // F16 doesn't use this tensor - auto s_tensor = TensorWrapper(nullptr, std::vector{1}, dtype); - - auto bias_shape = std::vector{bias_batch, bias_heads, q_max_seqlen, kv_max_seqlen}; - auto dbias_tensor = TensorWrapper(nullptr, bias_shape, dtype); - - NVTETensorPack aux_input_tensors; - nvte_tensor_pack_create(&aux_input_tensors); - - TensorWrapper query_workspace_tensor; - - // It is a WAR to pre-create all possible cuDNN graph at the JIT compile time - size_t max_num_segments = is_ragged ? input_batch * max_segments_per_seq : input_batch; - size_t min_num_segments = input_batch; - auto cudnn_runtime_version = cudnnGetVersion(); - if (is_ragged && cudnn_runtime_version >= 90300) { - // For cuDNN < 9.3.0, it requires to run all possible seqlens to address act_seqlen = 0 - min_num_segments = input_batch * max_segments_per_seq; - } - - TensorWrapper dummy_d_softmax_offset_tensor; - if (softmax_type == NVTE_Softmax_Type::NVTE_OFF_BY_ONE_SOFTMAX || - softmax_type == NVTE_Softmax_Type::NVTE_LEARNABLE_SOFTMAX) { - dummy_d_softmax_offset_tensor = - TensorWrapper(nullptr, std::vector{1, attn_heads, 1, 1}, DType::kFloat32); - } - - for (auto num_segments = min_num_segments; num_segments <= max_num_segments; ++num_segments) { - // the last one is the largest which will be the returned workspace size - auto q_cu_seqlens_tensor = - TensorWrapper(nullptr, std::vector{num_segments + 1}, DType::kInt32); - auto kv_cu_seqlens_tensor = - TensorWrapper(nullptr, std::vector{num_segments + 1}, DType::kInt32); - auto dummy_ragged_offset_tensor = - TensorWrapper(nullptr, std::vector{num_segments + 1}, DType::kInt32); - - nvte_fused_attn_bwd( - q_tensor.data(), k_tensor.data(), v_tensor.data(), output_tensor.data(), - doutput_tensor.data(), - s_tensor.data(), // not used for F16 - s_tensor.data(), // not used for F16 - &aux_input_tensors, dq_tensor.data(), dk_tensor.data(), dv_tensor.data(), - dbias_tensor.data(), dummy_d_softmax_offset_tensor.data(), q_cu_seqlens_tensor.data(), - kv_cu_seqlens_tensor.data(), dummy_ragged_offset_tensor.data(), - dummy_ragged_offset_tensor.data(), q_max_seqlen, kv_max_seqlen, scaling_factor, - dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), - nvte_get_q_format(qkv_layout), qkv_layout, NVTE_QKV_Format_NOT_SET, NVTE_QKV_Format_NOT_SET, - bias_type, mask_type, softmax_type, window_size_left, window_size_right, - bottom_right_diagonal, deterministic, false, query_workspace_tensor.data(), nullptr); - } - - nvte_tensor_pack_destroy(&aux_input_tensors); - - auto work_shape = MakeShapeVector(query_workspace_tensor.shape()); - return pybind11::make_tuple(work_shape, query_workspace_tensor.dtype()); -} - -static void FusedAttnBackwardImpl( - cudaStream_t stream, void *q, void *k, void *v, void *bias, void *softmax_offset, - void *softmax_aux, void *rng_state, void *output, void *doutput, void *q_cu_seqlens, - void *kv_cu_seqlens, void *q_seq_offsets, void *k_seq_offsets, void *dq, void *dk, void *dv, - void *dbias, void *dsoftmax_offset, void *workspace, size_t input_batch, size_t bias_batch, - size_t q_max_seqlen, size_t kv_max_seqlen, size_t attn_heads, size_t num_gqa_groups, - size_t bias_heads, size_t qk_head_dim, size_t v_head_dim, size_t max_segments_per_seq, - size_t wkspace_size, float scaling_factor, float dropout_probability, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, NVTE_QKV_Layout qkv_layout, - DType dtype, DType wkspace_dtype, bool is_training, bool deterministic, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal) { - FUSED_ATTN_IMPL_COMMON_BLOCK; - - /* Input tensors */ - auto output_shape = is_ragged - ? std::vector{input_batch * q_max_seqlen, attn_heads, v_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, v_head_dim}; - auto output_tensor = TensorWrapper(output, output_shape, dtype); - auto doutput_tensor = TensorWrapper(doutput, output_shape, dtype); - - /* Output tensors */ - auto s_tensor = TensorWrapper(nullptr, std::vector{1}, dtype); // not used in F16 - auto dbias_tensor = TensorWrapper(dbias, bias_shape, dtype); - - TensorWrapper dsoftmax_offset_tensor; - if (softmax_type == NVTE_Softmax_Type::NVTE_OFF_BY_ONE_SOFTMAX || - softmax_type == NVTE_Softmax_Type::NVTE_LEARNABLE_SOFTMAX) { - dsoftmax_offset_tensor = - TensorWrapper(dsoftmax_offset, std::vector{1, attn_heads, 1, 1}, DType::kFloat32); - } - - /* Auxiliary tensors (propagated from the forward pass) */ - NVTETensorPack aux_input_tensors; - nvte_tensor_pack_create(&aux_input_tensors); - auto backend = nvte_get_fused_attn_backend( - is_training, static_cast(dtype), static_cast(dtype), qkv_layout, - bias_type, mask_type, softmax_type, dropout_probability, attn_heads, num_gqa_groups, - q_max_seqlen, kv_max_seqlen, qk_head_dim, v_head_dim, window_size_left, window_size_right, - false, false, deterministic); - PrepareFusedAttnBackwardAuxTensors(&aux_input_tensors, input_batch, bias_batch, attn_heads, - bias_heads, q_max_seqlen, kv_max_seqlen, dtype, backend, - softmax_aux, rng_state, bias, softmax_offset); - - /* Call the underlying NVTE API */ - // Prepare Q, K, V pointers and shapes based on layout - void *q_ptr = q; - void *k_ptr = k; - void *v_ptr = v; - void *dq_ptr = dq; - void *dk_ptr = dk; - void *dv_ptr = dv; - auto q_shape = is_ragged - ? std::vector{input_batch * q_max_seqlen, attn_heads, qk_head_dim} - : std::vector{input_batch, q_max_seqlen, attn_heads, qk_head_dim}; - auto k_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, qk_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, qk_head_dim}; - auto v_shape = is_ragged - ? std::vector{input_batch * kv_max_seqlen, num_gqa_groups, v_head_dim} - : std::vector{input_batch, kv_max_seqlen, num_gqa_groups, v_head_dim}; - - if (layout_group == NVTE_QKV_Layout_Group::NVTE_3HD) { - // QKV packed in q: [batch*seqlen, 3, heads, dim] - NVTE_CHECK(q_max_seqlen == kv_max_seqlen, "q_max_seqlen must equal kv_max_seqlen"); - NVTE_CHECK(qk_head_dim == v_head_dim, - "For QKV packed layout, qk_head_dim must equal v_head_dim"); - size_t stride = (typeToSize(dtype) * attn_heads * qk_head_dim); - q_ptr = q; - k_ptr = static_cast(static_cast(q) + stride); - v_ptr = static_cast(static_cast(q) + 2 * stride); - dq_ptr = dq; - dk_ptr = static_cast(static_cast(dq) + stride); - dv_ptr = static_cast(static_cast(dq) + 2 * stride); - k_shape = q_shape; - v_shape = q_shape; - } else if (layout_group == NVTE_QKV_Layout_Group::NVTE_HD_2HD) { - // Q separate, KV packed in k: [batch*seqlen, 2, num_gqa_groups, dim] - NVTE_CHECK(qk_head_dim == v_head_dim, - "For KV packed layout, qk_head_dim must equal v_head_dim"); - size_t stride = (typeToSize(dtype) * num_gqa_groups * qk_head_dim); - q_ptr = q; - k_ptr = k; - v_ptr = static_cast(static_cast(k) + stride); - dq_ptr = dq; - dk_ptr = dk; - dv_ptr = static_cast(static_cast(dk) + stride); - // V has same shape as K since they're packed together - v_shape = k_shape; - } - - auto q_tensor = TensorWrapper(q_ptr, q_shape, dtype); - auto k_tensor = TensorWrapper(k_ptr, k_shape, dtype); - auto v_tensor = TensorWrapper(v_ptr, v_shape, dtype); - auto dq_tensor = TensorWrapper(dq_ptr, q_shape, dtype); - auto dk_tensor = TensorWrapper(dk_ptr, k_shape, dtype); - auto dv_tensor = TensorWrapper(dv_ptr, v_shape, dtype); - - if (is_ragged) { - size_t dtype_size = typeToSize(dtype); - if (layout_group == NVTE_QKV_Layout_Group::NVTE_3HD) { - // For packed QKV, dq contains all gradients (dq, dk, dv) - clear all at once - cudaMemsetAsync(dq, 0, 3 * transformer_engine::jax::product(q_shape) * dtype_size, stream); - } else if (layout_group == NVTE_QKV_Layout_Group::NVTE_HD_2HD) { - // Clear dq - cudaMemsetAsync(dq, 0, transformer_engine::jax::product(q_shape) * dtype_size, stream); - // For packed KV, dk contains both dk and dv - clear all at once - cudaMemsetAsync(dk, 0, 2 * transformer_engine::jax::product(k_shape) * dtype_size, stream); - } else { - // All separate - clear each individually - cudaMemsetAsync(dq, 0, transformer_engine::jax::product(q_shape) * dtype_size, stream); - cudaMemsetAsync(dk, 0, transformer_engine::jax::product(k_shape) * dtype_size, stream); - cudaMemsetAsync(dv, 0, transformer_engine::jax::product(v_shape) * dtype_size, stream); - } - } - - nvte_fused_attn_bwd( - q_tensor.data(), k_tensor.data(), v_tensor.data(), output_tensor.data(), - doutput_tensor.data(), - s_tensor.data(), // not used for F16 - s_tensor.data(), // not used for F16 - &aux_input_tensors, dq_tensor.data(), dk_tensor.data(), dv_tensor.data(), dbias_tensor.data(), - dsoftmax_offset_tensor.data(), q_cu_seqlens_tensor.data(), kv_cu_seqlens_tensor.data(), - q_seq_offsets_tensor.data(), k_seq_offsets_tensor.data(), q_max_seqlen, kv_max_seqlen, - scaling_factor, dropout_probability, qkv_layout, nvte_get_q_format(qkv_layout), - nvte_get_q_format(qkv_layout), qkv_layout, NVTE_QKV_Format_NOT_SET, NVTE_QKV_Format_NOT_SET, - bias_type, mask_type, softmax_type, window_size_left, window_size_right, - bottom_right_diagonal, deterministic, false, workspace_tensor.data(), stream); - - nvte_tensor_pack_destroy(&aux_input_tensors); -} - -Error_Type FusedAttnBackwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, - Buffer_Type v_buf, Buffer_Type bias_buf, - Buffer_Type softmax_offset_buf, Buffer_Type softmax_aux_buf, - Buffer_Type rng_state_buf, Buffer_Type output_buf, - Buffer_Type doutput_buf, Buffer_Type q_cu_seqlens_buf, - Buffer_Type kv_cu_seqlens_buf, Buffer_Type q_seq_offsets_buf, - Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type _unused_args, - Result_Type dq_buf, Result_Type dk_buf, Result_Type dv_buf, - Result_Type dbias_buf, Result_Type dsoftmax_offset_buf, - Result_Type workspace_buf, Dictionary attrs) { - FUSED_ATTN_FFI_GET_ATTRS; - - FusedAttnBackwardImpl( - stream, q_buf.untyped_data(), k_buf.untyped_data(), v_buf.untyped_data(), - bias_buf.untyped_data(), softmax_offset_buf.untyped_data(), softmax_aux_buf.untyped_data(), - rng_state_buf.untyped_data(), output_buf.untyped_data(), doutput_buf.untyped_data(), - q_cu_seqlens_buf.untyped_data(), kv_cu_seqlens_buf.untyped_data(), - is_ragged ? q_seq_offsets_buf.untyped_data() : nullptr, - is_ragged ? k_seq_offsets_buf.untyped_data() : nullptr, dq_buf->untyped_data(), - dk_buf->untyped_data(), dv_buf->untyped_data(), dbias_buf->untyped_data(), - dsoftmax_offset_buf->untyped_data(), workspace_buf->untyped_data(), input_batch, bias_batch, - q_max_seqlen, kv_max_seqlen, attn_heads, num_gqa_groups, bias_heads, qk_head_dim, v_head_dim, - max_segments_per_seq, wkspace_size, scaling_factor, dropout_probability, bias_type, mask_type, - softmax_type, qkv_layout, dtype, wkspace_dtype, is_training, deterministic, window_size_left, - window_size_right, bottom_right_diagonal); - - return ffi_with_cuda_error_check(); -} - -XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnBackwardHandler, FusedAttnBackwardFFI, - FFI::Bind() - .Ctx() // stream - .Arg() // q - .Arg() // k - .Arg() // v - .Arg() // bias - .Arg() // softmax_offset - .Arg() // softmax_aux - .Arg() // rng_state - .Arg() // output - .Arg() // doutput - .Arg() // q_cu_seqlens - .Arg() // kv_cu_seqlens - .Arg() // q_seq_offsets - .Arg() // k_seq_offsets - .RemainingArgs() // _cp_aux_args unused - .Ret() // dq - .Ret() // dk - .Ret() // dv - .Ret() // dbias - .Ret() // dsoftmax_offset - .Ret() // workspace - .Attrs(), - FFI_CudaGraph_Traits); - namespace { -struct ScoreModScalarStorage { +struct CudnnGraphScalarStorage { alignas(16) std::array data{}; - size_t size = 0; }; -struct ScoreModGraphCacheKey { +struct CudnnGraphCacheKey { int device_id = 0; int64_t hash0 = 0; int64_t hash1 = 0; int64_t frontend_version = 0; - bool operator==(const ScoreModGraphCacheKey &other) const { + bool operator==(const CudnnGraphCacheKey &other) const { return device_id == other.device_id && hash0 == other.hash0 && hash1 == other.hash1 && frontend_version == other.frontend_version; } }; -struct ScoreModGraphCacheKeyHash { - size_t operator()(const ScoreModGraphCacheKey &key) const { +struct CudnnGraphCacheKeyHash { + size_t operator()(const CudnnGraphCacheKey &key) const { size_t seed = std::hash{}(key.device_id); auto combine = [&seed](int64_t value) { - // 64-bit golden ratio constant from boost::hash_combine to spread mixed keys. seed ^= std::hash{}(value) + 0x9e3779b97f4a7c15ULL + (seed << 6) + (seed >> 2); }; combine(key.hash0); @@ -760,21 +52,20 @@ struct ScoreModGraphCacheKeyHash { } }; -using ScoreModGraphPtr = std::shared_ptr; +using CudnnGraphPtr = std::shared_ptr; -std::unordered_map & -getScoreModeGraphCache() { - static std::unordered_map - cache; +std::unordered_map & +GetCudnnGraphCache() { + static std::unordered_map cache; return cache; } -std::mutex &getScoreModGraphCacheMutex() { +std::mutex &GetCudnnGraphCacheMutex() { static std::mutex mutex; return mutex; } -struct ScoreModCudnnHandleCache { +struct CudnnHandleCache { std::unordered_map handles; cudnnHandle_t GetHandle() { @@ -789,30 +80,29 @@ struct ScoreModCudnnHandleCache { return it->second; } - ~ScoreModCudnnHandleCache() { + ~CudnnHandleCache() { for (auto &[_, handle] : handles) { cudnnDestroy(handle); } } }; -cudnnHandle_t GetScoreModCudnnHandle() { - static thread_local ScoreModCudnnHandleCache cache; +cudnnHandle_t GetCudnnHandle() { + static thread_local CudnnHandleCache cache; return cache.GetHandle(); } -ScoreModGraphCacheKey GetScoreModGraphCacheKey(Dictionary &attrs) { +CudnnGraphCacheKey GetCudnnGraphCacheKey(Dictionary &attrs) { const int64_t frontend_version = get_attr_value(attrs, "cudnn_frontend_version"); NVTE_CHECK(frontend_version == CUDNN_FRONTEND_VERSION, - "cuDNN frontend version mismatch for score_mod graph deserialization: graph was " - "serialized with Python cuDNN frontend version ", - frontend_version, - ", but Transformer Engine C++ was built with CUDNN_FRONTEND_VERSION ", + "cuDNN frontend version mismatch for graph deserialization: graph was serialized " + "with Python frontend version ", + frontend_version, ", but Transformer Engine C++ was built with version ", CUDNN_FRONTEND_VERSION, "."); int device_id = 0; NVTE_CHECK_CUDA(cudaGetDevice(&device_id)); - return ScoreModGraphCacheKey{ + return CudnnGraphCacheKey{ device_id, get_attr_value(attrs, "graph_hash0"), get_attr_value(attrs, "graph_hash1"), @@ -820,11 +110,11 @@ ScoreModGraphCacheKey GetScoreModGraphCacheKey(Dictionary &attrs) { }; } -ScoreModGraphPtr GetScoreModGraph(cudaStream_t stream, Dictionary &attrs) { - const auto key = GetScoreModGraphCacheKey(attrs); +CudnnGraphPtr GetCudnnGraph(cudaStream_t stream, Dictionary &attrs) { + const auto key = GetCudnnGraphCacheKey(attrs); { - std::lock_guard lock(getScoreModGraphCacheMutex()); - auto &cache = getScoreModeGraphCache(); + std::lock_guard lock(GetCudnnGraphCacheMutex()); + auto &cache = GetCudnnGraphCache(); auto it = cache.find(key); if (it != cache.end()) { return it->second; @@ -833,17 +123,16 @@ ScoreModGraphPtr GetScoreModGraph(cudaStream_t stream, Dictionary &attrs) { const auto serialized_graph = get_attr_value(attrs, "serialized_graph"); std::vector serialized_data(serialized_graph.begin(), serialized_graph.end()); - - auto handle = GetScoreModCudnnHandle(); + auto handle = GetCudnnHandle(); NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); auto graph = std::make_shared(); auto status = graph->deserialize(handle, serialized_data); - NVTE_CHECK(status.is_good(), - "Failed to deserialize cuDNN score_mod SDPA graph: ", status.get_message()); + NVTE_CHECK(status.is_good(), "Failed to deserialize cuDNN frontend graph: ", + status.get_message()); - std::lock_guard lock(getScoreModGraphCacheMutex()); - auto &cache = getScoreModeGraphCache(); + std::lock_guard lock(GetCudnnGraphCacheMutex()); + auto &cache = GetCudnnGraphCache(); auto it = cache.find(key); if (it != cache.end()) { return it->second; @@ -852,47 +141,68 @@ ScoreModGraphPtr GetScoreModGraph(cudaStream_t stream, Dictionary &attrs) { return graph; } -Error_Type ExecuteScoreModGraph(cudaStream_t stream, Dictionary &attrs, - const std::vector &input_ptrs, - const std::vector &output_ptrs, void *workspace) { - auto graph = GetScoreModGraph(stream, attrs); +Error_Type ExecuteCudnnGraph(cudaStream_t stream, Dictionary &attrs, + const std::vector &input_ptrs, + const std::vector &output_ptrs, void *workspace) { + auto graph = GetCudnnGraph(stream, attrs); auto input_uids = get_attr_value>(attrs, "input_uids"); + auto input_buffer_indices = + get_attr_value>(attrs, "input_buffer_indices"); + auto input_byte_offsets = + get_attr_value>(attrs, "input_byte_offsets"); auto output_uids = get_attr_value>(attrs, "output_uids"); + auto output_buffer_indices = + get_attr_value>(attrs, "output_buffer_indices"); + auto output_byte_offsets = + get_attr_value>(attrs, "output_byte_offsets"); auto scalar_uids = get_attr_value>(attrs, "scalar_uids"); auto scalar_sizes = get_attr_value>(attrs, "scalar_sizes"); auto scalar_values = get_attr_value>(attrs, "scalar_values"); - NVTE_CHECK(input_ptrs.size() == input_uids.size(), "cuDNN score_mod graph expected ", - input_uids.size(), " inputs but got ", input_ptrs.size()); - NVTE_CHECK(output_ptrs.size() >= output_uids.size(), "cuDNN score_mod graph expected at least ", - output_uids.size(), " outputs but got ", output_ptrs.size()); + NVTE_CHECK(input_uids.size() == input_buffer_indices.size() && + input_uids.size() == input_byte_offsets.size(), + "Mismatched cuDNN graph input binding metadata."); + NVTE_CHECK(output_uids.size() == output_buffer_indices.size() && + output_uids.size() == output_byte_offsets.size(), + "Mismatched cuDNN graph output binding metadata."); NVTE_CHECK(scalar_uids.size() == scalar_sizes.size(), - "Mismatched score_mod scalar uid/value-size counts."); + "Mismatched cuDNN graph scalar uid/value-size counts."); NVTE_CHECK(scalar_values.size() == scalar_uids.size() * 16, - "Mismatched score_mod packed scalar value size."); + "Mismatched cuDNN graph packed scalar value size."); std::unordered_map variant_pack; for (size_t i = 0; i < input_uids.size(); ++i) { - variant_pack.emplace(input_uids[i], input_ptrs[i]); + NVTE_CHECK(input_buffer_indices[i] >= 0 && + static_cast(input_buffer_indices[i]) < input_ptrs.size(), + "cuDNN graph input binding index is out of range."); + NVTE_CHECK(input_byte_offsets[i] >= 0, "cuDNN graph input byte offset must be non-negative."); + auto *ptr = static_cast(input_ptrs[input_buffer_indices[i]]) + + input_byte_offsets[i]; + variant_pack.emplace(input_uids[i], ptr); } for (size_t i = 0; i < output_uids.size(); ++i) { - variant_pack.emplace(output_uids[i], output_ptrs[i]); + NVTE_CHECK(output_buffer_indices[i] >= 0 && + static_cast(output_buffer_indices[i]) < output_ptrs.size(), + "cuDNN graph output binding index is out of range."); + NVTE_CHECK(output_byte_offsets[i] >= 0, + "cuDNN graph output byte offset must be non-negative."); + auto *ptr = static_cast(output_ptrs[output_buffer_indices[i]]) + + output_byte_offsets[i]; + variant_pack.emplace(output_uids[i], ptr); } - std::vector scalar_storage(scalar_uids.size()); + std::vector scalar_storage(scalar_uids.size()); for (size_t i = 0; i < scalar_uids.size(); ++i) { NVTE_CHECK(scalar_sizes[i] >= 0 && scalar_sizes[i] <= 16, - "score_mod pass-by-value scalars must be at most 16 bytes."); - scalar_storage[i].size = static_cast(scalar_sizes[i]); + "cuDNN graph pass-by-value scalars must be at most 16 bytes."); std::copy_n(scalar_values.begin() + i * 16, 16, scalar_storage[i].data.begin()); variant_pack.emplace(scalar_uids[i], scalar_storage[i].data.data()); } - auto handle = GetScoreModCudnnHandle(); + auto handle = GetCudnnHandle(); NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); auto status = graph->execute(handle, variant_pack, workspace); - NVTE_CHECK(status.is_good(), - "cuDNN score_mod SDPA graph execution failed: ", status.get_message()); + NVTE_CHECK(status.is_good(), "cuDNN frontend graph execution failed: ", status.get_message()); return ffi_with_cuda_error_check(); } @@ -900,13 +210,160 @@ void AppendRemainingBuffers(Variadic_Buffer_Type args, std::vector *ptrs ptrs->reserve(ptrs->size() + args.size()); for (size_t i = 0; i < args.size(); ++i) { auto maybe_buf = args.get(i); - NVTE_CHECK(!maybe_buf.has_error(), "Failed to decode variadic score_mod input buffer."); + NVTE_CHECK(!maybe_buf.has_error(), "Failed to decode variadic cuDNN graph input buffer."); ptrs->push_back(maybe_buf.value().untyped_data()); } } +size_t BufferBytes(const Buffer_Type &buffer) { + return product(buffer.dimensions()) * + typeToSize(convert_ffi_datatype_to_te_dtype(buffer.element_type())); +} + +void MemsetResultAsync(cudaStream_t stream, Result_Type result, int value) { + NVTE_CHECK_CUDA(cudaMemsetAsync(result->untyped_data(), value, BufferBytes(*result), stream)); +} + +class FusedAttnOffsetManager { + public: + static FusedAttnOffsetManager &Instance() { + static thread_local FusedAttnOffsetManager manager; + return manager; + } + + uint64_t GetAndUpdate(uint64_t increment) { + uint64_t current = offset_; + offset_ += increment; + return current; + } + + private: + uint64_t offset_ = 0; +}; + +void PopulateRngStateAsync(cudaStream_t stream, const Buffer_Type &seed, Result_Type rng_state, + uint64_t increment) { + NVTE_CHECK(BufferBytes(seed) >= sizeof(uint64_t), "Fused-attention seed buffer is too small."); + NVTE_CHECK(BufferBytes(*rng_state) >= 2 * sizeof(uint64_t), + "Fused-attention RNG-state buffer is too small."); + const uint64_t offset = FusedAttnOffsetManager::Instance().GetAndUpdate(increment); + PopulateFusedAttnRngState(rng_state->untyped_data(), seed.untyped_data(), offset, stream); +} + } // namespace +Error_Type FusedAttnForwardFFI( + cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, Buffer_Type v_buf, + Buffer_Type bias_buf, Buffer_Type softmax_offset_buf, Buffer_Type seed_buf, + Buffer_Type q_seqlens_buf, Buffer_Type kv_seqlens_buf, Buffer_Type q_seq_offsets_buf, + Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type remaining_args, Result_Type output_buf, + Result_Type stats_buf, Result_Type max_buf, Result_Type rng_state_buf, + Result_Type workspace_buf, Dictionary attrs) { + const bool is_ragged = get_attr_value(attrs, "is_ragged"); + const uint64_t rng_increment = + static_cast(get_attr_value(attrs, "rng_offset_increment")); + PopulateRngStateAsync(stream, seed_buf, rng_state_buf, rng_increment); + if (is_ragged) { + MemsetResultAsync(stream, output_buf, 0); + MemsetResultAsync(stream, stats_buf, 0xF0); + if (BufferBytes(*max_buf) != 0) { + MemsetResultAsync(stream, max_buf, 0xF0); + } + } + + std::vector input_ptrs = { + q_buf.untyped_data(), k_buf.untyped_data(), + v_buf.untyped_data(), bias_buf.untyped_data(), + softmax_offset_buf.untyped_data(), seed_buf.untyped_data(), + q_seqlens_buf.untyped_data(), kv_seqlens_buf.untyped_data(), + q_seq_offsets_buf.untyped_data(), k_seq_offsets_buf.untyped_data(), + }; + AppendRemainingBuffers(remaining_args, &input_ptrs); + std::vector output_ptrs = {output_buf->untyped_data(), stats_buf->untyped_data(), + max_buf->untyped_data(), rng_state_buf->untyped_data()}; + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, + workspace_buf->untyped_data()); +} + +XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnForwardHandler, FusedAttnForwardFFI, + FFI::Bind() + .Ctx() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .RemainingArgs() + .Ret() + .Ret() + .Ret() + .Ret() + .Ret() + .Attrs(), + FFI_CudaGraph_Traits); + +Error_Type FusedAttnBackwardFFI( + cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, Buffer_Type v_buf, + Buffer_Type bias_buf, Buffer_Type softmax_offset_buf, Buffer_Type stats_buf, + Buffer_Type rng_state_buf, Buffer_Type output_buf, Buffer_Type doutput_buf, + Buffer_Type q_seqlens_buf, Buffer_Type kv_seqlens_buf, Buffer_Type q_seq_offsets_buf, + Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type remaining_args, Result_Type dq_buf, + Result_Type dk_buf, Result_Type dv_buf, Result_Type dbias_buf, + Result_Type dsoftmax_offset_buf, Result_Type workspace_buf, Dictionary attrs) { + if (get_attr_value(attrs, "is_ragged")) { + MemsetResultAsync(stream, dq_buf, 0); + MemsetResultAsync(stream, dk_buf, 0); + MemsetResultAsync(stream, dv_buf, 0); + } + std::vector input_ptrs = { + q_buf.untyped_data(), k_buf.untyped_data(), + v_buf.untyped_data(), bias_buf.untyped_data(), + softmax_offset_buf.untyped_data(), stats_buf.untyped_data(), + rng_state_buf.untyped_data(), output_buf.untyped_data(), + doutput_buf.untyped_data(), q_seqlens_buf.untyped_data(), + kv_seqlens_buf.untyped_data(), q_seq_offsets_buf.untyped_data(), + k_seq_offsets_buf.untyped_data(), + }; + AppendRemainingBuffers(remaining_args, &input_ptrs); + std::vector output_ptrs = { + dq_buf->untyped_data(), dk_buf->untyped_data(), dv_buf->untyped_data(), + dbias_buf->untyped_data(), dsoftmax_offset_buf->untyped_data(), + }; + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, + workspace_buf->untyped_data()); +} + +XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnBackwardHandler, FusedAttnBackwardFFI, + FFI::Bind() + .Ctx() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .RemainingArgs() + .Ret() + .Ret() + .Ret() + .Ret() + .Ret() + .Ret() + .Attrs(), + FFI_CudaGraph_Traits); + Error_Type FusedAttnScoreModForwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, Buffer_Type v_buf, Variadic_Buffer_Type score_mod_args, Result_Type output_buf, Result_Type stats_buf, @@ -914,23 +371,23 @@ Error_Type FusedAttnScoreModForwardFFI(cudaStream_t stream, Buffer_Type q_buf, B std::vector input_ptrs = {q_buf.untyped_data(), k_buf.untyped_data(), v_buf.untyped_data()}; AppendRemainingBuffers(score_mod_args, &input_ptrs); - std::vector output_ptrs = {output_buf->untyped_data(), stats_buf->untyped_data()}; - return ExecuteScoreModGraph(stream, attrs, input_ptrs, output_ptrs, - workspace_buf->untyped_data()); + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, + workspace_buf->untyped_data()); } XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnScoreModForwardHandler, FusedAttnScoreModForwardFFI, FFI::Bind() - .Ctx() // stream - .Arg() // q - .Arg() // k - .Arg() // v - .RemainingArgs() // score_mod tensor operands - .Ret() // output - .Ret() // stats - .Ret() // workspace - .Attrs()); + .Ctx() + .Arg() + .Arg() + .Arg() + .RemainingArgs() + .Ret() + .Ret() + .Ret() + .Attrs(), + FFI_CudaGraph_Traits); Error_Type FusedAttnScoreModBackwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, Buffer_Type v_buf, Buffer_Type output_buf, @@ -942,28 +399,28 @@ Error_Type FusedAttnScoreModBackwardFFI(cudaStream_t stream, Buffer_Type q_buf, v_buf.untyped_data(), output_buf.untyped_data(), doutput_buf.untyped_data(), stats_buf.untyped_data()}; AppendRemainingBuffers(score_mod_args, &input_ptrs); - std::vector output_ptrs = {dq_buf->untyped_data(), dk_buf->untyped_data(), dv_buf->untyped_data()}; - return ExecuteScoreModGraph(stream, attrs, input_ptrs, output_ptrs, - workspace_buf->untyped_data()); + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, + workspace_buf->untyped_data()); } XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnScoreModBackwardHandler, FusedAttnScoreModBackwardFFI, FFI::Bind() - .Ctx() // stream - .Arg() // q - .Arg() // k - .Arg() // v - .Arg() // output - .Arg() // doutput - .Arg() // stats - .RemainingArgs() // score_mod tensor operands - .Ret() // dq - .Ret() // dk - .Ret() // dv - .Ret() // workspace - .Attrs()); + .Ctx() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .Arg() + .RemainingArgs() + .Ret() + .Ret() + .Ret() + .Ret() + .Attrs(), + FFI_CudaGraph_Traits); } // namespace jax } // namespace transformer_engine diff --git a/transformer_engine/jax/csrc/extensions/attention_kernels.cu b/transformer_engine/jax/csrc/extensions/attention_kernels.cu new file mode 100644 index 00000000000..a5d32fa8d9c --- /dev/null +++ b/transformer_engine/jax/csrc/extensions/attention_kernels.cu @@ -0,0 +1,29 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +#include "../extensions.h" + +namespace transformer_engine { +namespace jax { +namespace { + +__global__ void PopulateFusedAttnRngStateKernel(int64_t *rng_state, const int64_t *seed, + uint64_t offset) { + rng_state[0] = seed[0]; + rng_state[1] = static_cast(offset); +} + +} // namespace + +void PopulateFusedAttnRngState(void *rng_state, const void *seed, uint64_t offset, + cudaStream_t stream) { + PopulateFusedAttnRngStateKernel<<<1, 1, 0, stream>>>(static_cast(rng_state), + static_cast(seed), offset); + NVTE_CHECK_CUDA(cudaGetLastError()); +} + +} // namespace jax +} // namespace transformer_engine diff --git a/transformer_engine/jax/csrc/extensions/pybind.cpp b/transformer_engine/jax/csrc/extensions/pybind.cpp index 3927e2686ed..8a59c891c81 100644 --- a/transformer_engine/jax/csrc/extensions/pybind.cpp +++ b/transformer_engine/jax/csrc/extensions/pybind.cpp @@ -134,7 +134,6 @@ pybind11::dict Registrations() { PYBIND11_MODULE(transformer_engine_jax, m) { m.def("registrations", &Registrations); - m.def("get_fused_attn_backend", &GetFusedAttnBackend); m.def("get_cuda_version", &GetCudaRuntimeVersion); m.def("get_cudnn_version", &GetCudnnRuntimeVersion); m.def("get_cudnn_frontend_version", &GetCudnnFrontendVersion); @@ -145,8 +144,6 @@ PYBIND11_MODULE(transformer_engine_jax, m) { m.def("get_dbias_quantize_workspace_sizes", &GetDBiasQuantizeWorkspaceSizes); m.def("get_norm_fwd_workspace_sizes", &GetNormForwardWorkspaceSizes); m.def("get_norm_bwd_workspace_sizes", &GetNormBackwardWorkspaceSizes); - m.def("get_fused_attn_fwd_workspace_sizes", &GetFusedAttnForwardWorkspaceSizes); - m.def("get_fused_attn_bwd_workspace_sizes", &GetFusedAttnBackwardWorkspaceSizes); m.def("get_topk_workspace_sizes", &GetTopkWorkspaceSizes); m.def("nvte_get_qkv_format", &nvte_get_qkv_format); m.def("is_non_nt_fp8_gemm_supported", &nvte_is_non_tn_fp8_gemm_supported); From df6416352f0f1f5ba9186ec9ae7de483a5d3c43d Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 26 Aug 2026 21:34:06 +0000 Subject: [PATCH 02/36] Test shared cuDNN frontend version check Update the experimental flex-attention version test to exercise the cuDNN frontend compatibility check in the shared JAX graph module instead of retaining aliases for the old private helpers. Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fused_attn_score_mod.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/tests/jax/test_fused_attn_score_mod.py b/tests/jax/test_fused_attn_score_mod.py index 7d1137b21f2..4a00862d2d3 100644 --- a/tests/jax/test_fused_attn_score_mod.py +++ b/tests/jax/test_fused_attn_score_mod.py @@ -8,7 +8,10 @@ import jax.numpy as jnp import numpy as np import pytest +from test_fused_attn import FusedAttnRunner, SeqDescFormat +from transformer_engine_jax import get_device_compute_capability +import transformer_engine.jax.cpp_extensions.cudnn_graph as cudnn_graph import transformer_engine.jax.cpp_extensions.flex_attention as tex_attention from transformer_engine.jax.attention import ( AttnBiasType, @@ -18,9 +21,6 @@ ) from transformer_engine.jax.cpp_extensions import make_fused_attn_score_mod_config from transformer_engine.jax.flax import transformer as flax_transformer -from transformer_engine_jax import get_device_compute_capability -from test_fused_attn import FusedAttnRunner, SeqDescFormat - _CONFIG_TEST_HEAD_DIM = 128 _CONFIG_TEST_SCALING_FACTOR = 1.0 / sqrt(_CONFIG_TEST_HEAD_DIM) @@ -655,19 +655,19 @@ class FakeCudnn: __version__ = "1.22.0" monkeypatch.setattr( - tex_attention.transformer_engine_jax, + cudnn_graph.transformer_engine_jax, "get_cudnn_frontend_version", lambda: 12200, ) - assert tex_attention._check_cudnn_frontend_version_match(FakeCudnn) == 12200 + assert cudnn_graph.check_cudnn_frontend_version_match(FakeCudnn) == 12200 monkeypatch.setattr( - tex_attention.transformer_engine_jax, + cudnn_graph.transformer_engine_jax, "get_cudnn_frontend_version", lambda: 12100, ) with pytest.raises(RuntimeError, match="Python/C\\+\\+ version mismatch"): - tex_attention._check_cudnn_frontend_version_match(FakeCudnn) + cudnn_graph.check_cudnn_frontend_version_match(FakeCudnn) def test_fused_attn_score_mod_config_stabilizes_bound_method_cache_keys(): From e42c3ac09c52bfabc418eab657e52e87ca4cd38c Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 26 Aug 2026 22:05:38 +0000 Subject: [PATCH 03/36] Port PyTorch cuDNN attention to Python frontend Build and execute FusedAttention graphs through the cuDNN Frontend Python API for FP16, BF16, FP8, and MXFP8 configurations. Share graph runtime utilities with FlexAttention and mirror backend selection in Python. Remove the legacy PyTorch fused-attention bindings while retaining TE common for other frameworks. Add a graph-safe Philox reservation helper and require cuDNN Frontend 1.27.0. Signed-off-by: Vladimir Cherepanov --- README.rst | 2 +- build_tools/pytorch.py | 4 +- build_tools/wheel_utils/build_wheels.sh | 2 +- pyproject.toml | 2 +- tests/pytorch/test_torch_compile.py | 6 +- .../dot_product_attention/_cudnn_backend.py | 547 ++++ .../dot_product_attention/_cudnn_graph.py | 185 ++ .../dot_product_attention/backends.py | 2 +- .../dot_product_attention/context_parallel.py | 2 +- .../dot_product_attention/cudnn_attention.py | 2833 +++++++++++++++++ .../dot_product_attention/flex_attention.py | 98 +- .../attention/dot_product_attention/utils.py | 21 +- .../pytorch/cpp_extensions/fused_attn.py | 624 +--- transformer_engine/pytorch/csrc/extensions.h | 36 +- .../pytorch/csrc/extensions/attention.cpp | 598 ---- .../pytorch/csrc/extensions/attention_rng.cpp | 26 + .../pytorch/csrc/extensions/pybind.cpp | 11 +- 17 files changed, 3684 insertions(+), 1315 deletions(-) create mode 100644 transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py create mode 100644 transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py create mode 100644 transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py create mode 100644 transformer_engine/pytorch/csrc/extensions/attention_rng.cpp diff --git a/README.rst b/README.rst index 144bae37141..e5ed6ac971e 100644 --- a/README.rst +++ b/README.rst @@ -164,7 +164,7 @@ System Requirements * Compiler: GCC 9+ or Clang 10+ with C++17 support * Python: 3.12 recommended -* **Source Build Requirements:** CMake 3.18+, Ninja, Git 2.17+, pybind11 2.6.0+, nvidia-cudnn-frontend 1.25.0+ +* **Source Build Requirements:** CMake 3.18+, Ninja, Git 2.17+, pybind11 2.6.0+, nvidia-cudnn-frontend 1.27.0+ * **Notes:** FP8 features require Compute Capability 8.9+ (Ada/Hopper/Blackwell) diff --git a/build_tools/pytorch.py b/build_tools/pytorch.py index 98331ccbd86..3fb1148ffad 100644 --- a/build_tools/pytorch.py +++ b/build_tools/pytorch.py @@ -31,7 +31,9 @@ def install_requirements() -> List[str]: "packaging", "pydantic", "nvdlfw-inspect", - "nvidia-cudnn-frontend>=1.25.0", + # PyTorch cuDNN attention is built and executed with the Python graph + # API; FP8/MXFP8 graph capture requires Frontend 1.27 or newer. + "nvidia-cudnn-frontend>=1.27.0", ] diff --git a/build_tools/wheel_utils/build_wheels.sh b/build_tools/wheel_utils/build_wheels.sh index 8fe2b629bed..42174ebe688 100644 --- a/build_tools/wheel_utils/build_wheels.sh +++ b/build_tools/wheel_utils/build_wheels.sh @@ -23,7 +23,7 @@ git checkout $TARGET_BRANCH git submodule update --init --recursive # Install deps -/opt/python/cp310-cp310/bin/pip install cmake pybind11[global] ninja setuptools wheel 'nvidia-cudnn-frontend>=1.25.0' +/opt/python/cp310-cp310/bin/pip install cmake pybind11[global] ninja setuptools wheel 'nvidia-cudnn-frontend>=1.27.0' if $BUILD_METAPACKAGE ; then cd /TransformerEngine diff --git a/pyproject.toml b/pyproject.toml index 2c9f224c14c..3b3649aabe1 100755 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,7 +3,7 @@ # See LICENSE for license information. [build-system] -requires = ["setuptools>=61.0", "cmake>=3.21", "wheel", "pybind11[global]", "ninja", "pip", "torch>=2.1", "jax>=0.5.0", "flax>=0.7.1", "nvidia-cudnn-frontend>=1.25.0"] +requires = ["setuptools>=61.0", "cmake>=3.21", "wheel", "pybind11[global]", "ninja", "pip", "torch>=2.1", "jax>=0.5.0", "flax>=0.7.1", "nvidia-cudnn-frontend>=1.27.0"] # Use legacy backend to import local packages in setup.py build-backend = "setuptools.build_meta:__legacy__" diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index 6ebc408dcb8..807e1b05f0d 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -1187,7 +1187,7 @@ def test_get_attention_backend_traceable(monkeypatch): """get_attention_backend must trace under torch.compile(fullgraph=True) without graph breaks. The compiled selection must stay consistent with eager when NVTE_* env vars flip (dynamo guards on os.environ) and when - attention params change, and the baked tex.get_fused_attn_backend result + attention params change, and the baked Python cuDNN backend result must drive the selection.""" from transformer_engine.pytorch.attention.dot_product_attention import utils as dpa_utils @@ -1267,8 +1267,8 @@ def fn(x, params): # guard on the wrapped function). monkeypatch.setenv("NVTE_FLASH_ATTN", "0") monkeypatch.setattr( - dpa_utils.tex, - "get_fused_attn_backend", + dpa_utils, + "get_cudnn_fused_attn_backend", lambda *args: dpa_utils.FusedAttnBackend["No_Backend"], ) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py new file mode 100644 index 00000000000..2fb78a856a3 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py @@ -0,0 +1,547 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""cuDNN SDPA capability selection for the PyTorch frontend. + +This is intentionally kept in Python alongside the Python cuDNN graph builder. +The conditions mirror ``nvte_get_fused_attn_backend`` in TE common, which is +still used by the other framework frontends. +""" + +from __future__ import annotations + +import warnings + +from transformer_engine.pytorch.constants import DType +from transformer_engine.pytorch.utils import ( + get_cudnn_version, + get_device_compute_capability, +) + + +def _layout_info(qkv_layout: str): + paged = qkv_layout.startswith("paged_kv_") + layout = qkv_layout.removeprefix("paged_kv_") + components = layout.split("_") + + def tensor_format(component: str) -> str: + return "".join(char for char in component if char.isalpha()) + + q_format = tensor_format(components[0]) + kv_format = tensor_format(components[-1]) if len(components) > 1 else q_format + qkv_format = q_format if q_format == kv_format else f"{q_format}_2{kv_format}" + if paged: + layout_group = "paged_separate" + elif len(components) == 1 and "3" in components[0]: + layout_group = "h3d" if "h3d" in components[0] else "3hd" + elif len(components) == 2: + layout_group = "hd_h2d" if "h2d" in components[1] else "hd_2hd" + elif q_format == "bhsd": + layout_group = "sd_sd_sd" + else: + layout_group = "separate" + return qkv_format, q_format, kv_format, layout_group + + +def _requires_64bit_ragged_offset( + qkv_format: str, + layout_group: str, + num_attn_heads: int, + num_gqa_groups: int, + max_seqlen_q: int, + max_seqlen_kv: int, + head_dim_qk: int, + head_dim_v: int, +) -> bool: + if qkv_format != "thd": + return False + if layout_group in ("3hd", "h3d"): + q = k = v = 3 * num_attn_heads * head_dim_qk * max_seqlen_q + elif layout_group in ("hd_2hd", "hd_h2d"): + q = num_attn_heads * head_dim_qk * max_seqlen_q + k = v = 2 * num_gqa_groups * head_dim_qk * max_seqlen_kv + else: + q = num_attn_heads * head_dim_qk * max_seqlen_q + k = num_gqa_groups * head_dim_qk * max_seqlen_kv + v = num_gqa_groups * head_dim_v * max_seqlen_kv + output = num_attn_heads * head_dim_qk * max_seqlen_q + return max(q, k, v, output) > 2**31 - 1 + + +def _version_number() -> int: + major, minor, patch = get_cudnn_version() + return major * 10000 + minor * 100 + patch + + +def get_fused_attn_backend( + is_training, + q_dtype, + kv_dtype, + qkv_layout, + bias_type, + attn_mask_type, + softmax_type, + dropout, + num_attn_heads, + num_gqa_groups, + max_seqlen_q, + max_seqlen_kv, + head_dim_qk, + head_dim_v, + window_size_left, + window_size_right, + return_max_logit, + cuda_graph, + deterministic, +): + """Return the Python cuDNN SDPA backend value for an attention configuration.""" + + # Import lazily to keep this module independent of graph construction and + # avoid a circular import through dot_product_attention.utils. + from .cudnn_attention import FusedAttnBackend + + q_dtype = DType.cast(q_dtype) + kv_dtype = DType.cast(kv_dtype) + if q_dtype != kv_dtype: + raise ValueError("Q and KV must have the same data type") + + major, minor = get_device_compute_capability() + sm_arch = major * 10 + minor + cudnn_version = _version_number() + qkv_format, q_format, kv_format, layout_group = _layout_info(qkv_layout) + is_thd_layout = q_format == "thd" or kv_format == "thd" + requires_i64 = _requires_64bit_ragged_offset( + qkv_format, + layout_group, + num_attn_heads, + num_gqa_groups, + max_seqlen_q, + max_seqlen_kv, + head_dim_qk, + head_dim_v, + ) + supported_ragged_offset_size = not requires_i64 or cudnn_version >= 90500 + + fp8_dtype = q_dtype in (DType.kFloat8E4M3, DType.kFloat8E5M2) + fp8_shape_mask = ( + ( + cudnn_version >= 90201 + and sm_arch < 100 + and max_seqlen_q % 128 == 0 + and max_seqlen_kv % 128 == 0 + and head_dim_qk == 128 + and head_dim_v == 128 + and attn_mask_type in ("causal", "no_mask") + ) + or ( + cudnn_version >= 90700 + and ( + ( + sm_arch < 100 + and not is_training + and head_dim_qk <= 256 + and head_dim_v <= 256 + ) + or ( + sm_arch < 100 + and is_training + and head_dim_qk == 128 + and head_dim_v == 128 + ) + or (sm_arch >= 100 and head_dim_qk <= 128 and head_dim_v <= 128) + ) + and head_dim_qk % 16 == 0 + and head_dim_v % 16 == 0 + and attn_mask_type in ("no_mask", "causal", "padding", "padding_causal") + ) + or ( + cudnn_version >= 92100 + and sm_arch >= 100 + and head_dim_qk <= 192 + and head_dim_v <= 128 + and head_dim_qk % 16 == 0 + and head_dim_v % 16 == 0 + and attn_mask_type in ("no_mask", "causal", "causal_bottom_right") + ) + ) + fp8_format_softmax = ( + cudnn_version < 92100 + and qkv_format in ("bshd", "sbhd") + and softmax_type == "vanilla" + ) or (cudnn_version >= 92100 and qkv_format in ("bshd", "sbhd", "bhsd")) + if ( + fp8_dtype + and sm_arch >= 90 + and bias_type == "no_bias" + and fp8_shape_mask + and fp8_format_softmax + and not requires_i64 + and cudnn_version != 91000 + and not return_max_logit + ): + return FusedAttnBackend.FP8 + + if q_dtype not in (DType.kFloat16, DType.kBFloat16): + return FusedAttnBackend.No_Backend + + arch_supported = ( + (cudnn_version < 8903 and sm_arch in (80, 90)) + or (cudnn_version >= 8903 and 80 <= sm_arch < 100) + or (cudnn_version >= 90700 and sm_arch >= 100) + ) + seq_supported = cudnn_version >= 90000 or ( + max_seqlen_q % 64 == 0 and max_seqlen_kv % 64 == 0 + ) + heads_supported = cudnn_version >= 8907 or num_attn_heads == num_gqa_groups + + dim_supported = ( + head_dim_qk % 8 == 0 + and head_dim_v % 8 == 0 + and ( + (head_dim_qk <= 128 and head_dim_v <= 128) + or ( + head_dim_qk <= 256 + and head_dim_v <= 256 + and ( + (not is_training and sm_arch == 90 and cudnn_version >= 90100) + or (is_training and sm_arch == 90 and cudnn_version >= 90500) + ) + ) + or ( + not is_training + and sm_arch >= 100 + and cudnn_version >= 90900 + and max_seqlen_q > 1 + and layout_group != "paged_separate" + ) + or ( + not is_training + and cudnn_version >= 91002 + and ( + layout_group == "paged_separate" + or max_seqlen_q > 1 + or ( + max_seqlen_q == 1 + and attn_mask_type not in ("causal", "padding_causal") + ) + ) + ) + or ( + head_dim_qk == 192 + and head_dim_v == 128 + and is_training + and sm_arch >= 100 + and cudnn_version >= 91100 + ) + or ( + head_dim_qk == 256 + and head_dim_v == 256 + and is_training + and 100 <= sm_arch < 110 + and cudnn_version >= (92500 if is_thd_layout else 92300) + and layout_group != "paged_separate" + and bias_type == "no_bias" + and dropout == 0.0 + and softmax_type == "vanilla" + and ( + (window_size_left == -1 and window_size_right == -1) + or ( + attn_mask_type + in ( + "causal", + "padding_causal", + "causal_bottom_right", + "padding_causal_bottom_right", + ) + and window_size_right in (-1, 0) + ) + ) + ) + ) + ) + dim_supported = dim_supported and not ( + cudnn_version >= 91100 + and is_training + and sm_arch == 90 + and head_dim_qk >= 128 + and head_dim_v >= 128 + and (head_dim_qk, head_dim_v) != (192, 128) + and head_dim_qk != head_dim_v + ) + + bias_supported = ( + (cudnn_version < 8906 and bias_type == "no_bias") + or ( + cudnn_version >= 8906 + and ( + bias_type == "no_bias" + or ( + bias_type == "alibi" + and attn_mask_type + not in ( + "no_mask", + "padding", + "padding_causal", + "padding_causal_bottom_right", + ) + and sm_arch >= 90 + ) + or (bias_type == "post_scale_bias" and sm_arch >= 90) + ) + ) + or (cudnn_version >= 90000 and bias_type == "post_scale_bias" and sm_arch >= 80) + ) + + standard_format = qkv_format in ("sbhd", "bshd") + mask_supported = ( + (cudnn_version < 8906 and attn_mask_type == "causal") + or ( + cudnn_version >= 8906 + and standard_format + and attn_mask_type in ("causal", "padding", "padding_causal", "no_mask") + ) + or ( + cudnn_version >= 90100 + and qkv_format == "thd" + and attn_mask_type in ("padding", "padding_causal") + ) + or ( + cudnn_version >= 90300 + and standard_format + and attn_mask_type == "causal_bottom_right" + and max_seqlen_q % 64 == 0 + and max_seqlen_kv % 64 == 0 + and max_seqlen_q <= max_seqlen_kv + and bias_type == "no_bias" + and dropout == 0.0 + ) + or ( + cudnn_version >= 90500 + and layout_group == "paged_separate" + and ( + attn_mask_type in ("padding", "padding_causal") + or ( + attn_mask_type == "padding_causal_bottom_right" + and max_seqlen_q % 64 == 0 + and max_seqlen_kv % 64 == 0 + and max_seqlen_q <= max_seqlen_kv + ) + ) + and bias_type == "no_bias" + and dropout == 0.0 + ) + or ( + cudnn_version >= 90600 + and attn_mask_type == "padding_causal_bottom_right" + and max_seqlen_q % 64 == 0 + and max_seqlen_kv % 64 == 0 + and max_seqlen_q <= max_seqlen_kv + and bias_type == "no_bias" + and dropout == 0.0 + ) + or ( + cudnn_version >= 90700 + and ( + attn_mask_type in ("no_mask", "causal") + or ( + attn_mask_type + in ("padding", "padding_causal", "padding_causal_bottom_right") + and bias_type == "no_bias" + and dropout == 0.0 + ) + or ( + attn_mask_type + in ("causal_bottom_right", "padding_causal_bottom_right") + and max_seqlen_q <= max_seqlen_kv + ) + ) + ) + ) + + bias_mask_supported = not ( + cudnn_version >= 8906 + and attn_mask_type in ("padding", "padding_causal") + and bias_type == "post_scale_bias" + ) + format_supported = ( + qkv_format in ("sbhd", "bshd", "bhsd") + or ( + qkv_format == "thd" + and sm_arch >= 90 + and ( + (cudnn_version >= 90100 and num_attn_heads == num_gqa_groups) + or cudnn_version >= 90600 + ) + ) + or ( + q_format in ("sbhd", "bshd", "bhsd", "thd") + and kv_format in ("sbhd", "bshd", "bhsd", "thd") + and (q_format != "thd" or sm_arch >= 90) + and (kv_format != "thd" or sm_arch >= 90) + and cudnn_version >= 90700 + ) + ) + + sliding_window_supported = ( + ( + cudnn_version < 90200 + and window_size_left == -1 + and window_size_right in (-1, 0) + ) + or ( + cudnn_version >= 90200 + and ( + ( + window_size_left == -1 + and window_size_right == -1 + and attn_mask_type == "no_mask" + ) + or ( + window_size_left >= -1 + and window_size_right == 0 + and ( + attn_mask_type in ("no_mask", "causal") + or ( + attn_mask_type == "causal_bottom_right" + and max_seqlen_q == max_seqlen_kv + ) + ) + and max_seqlen_q <= max_seqlen_kv + and dropout == 0.0 + and bias_type == "no_bias" + and standard_format + ) + ) + ) + or ( + cudnn_version >= 90600 + and ( + (window_size_left == -1 and window_size_right in (-1, 0)) + or ( + window_size_left >= -1 + and window_size_right >= -1 + and ( + ( + attn_mask_type == "causal_bottom_right" + and ( + sm_arch < 100 + or ( + sm_arch >= 100 + and ( + ( + max_seqlen_q == max_seqlen_kv + and cudnn_version <= 90700 + ) + or cudnn_version > 90700 + ) + ) + ) + ) + or attn_mask_type in ("no_mask", "padding", "padding_causal") + or ( + attn_mask_type == "padding_causal_bottom_right" + and ( + sm_arch < 100 + or ( + sm_arch >= 100 + and ( + ( + max_seqlen_q == max_seqlen_kv + and cudnn_version <= 90700 + ) + or cudnn_version > 90700 + ) + ) + ) + ) + ) + and max_seqlen_q <= max_seqlen_kv + and bias_type == "no_bias" + and dropout == 0.0 + ) + ) + ) + ) + + softmax_supported = cudnn_version >= 91301 or softmax_type == "vanilla" + max_supported = not return_max_logit or cudnn_version >= 92100 + deterministic_supported = sm_arch < 100 or ( + not is_training + or ( + is_training + and not deterministic + and (dropout == 0.0 or bias_type == "no_bias") + ) + or ( + is_training + and deterministic + and cudnn_version >= 91801 + and dropout == 0.0 + and bias_type == "no_bias" + ) + ) + + supported = all( + ( + arch_supported, + seq_supported, + heads_supported, + dim_supported, + bias_supported, + mask_supported, + bias_mask_supported, + format_supported, + sliding_window_supported, + supported_ragged_offset_size, + cudnn_version not in (91000, 91001), + softmax_supported, + max_supported, + deterministic_supported, + ) + ) + backend = ( + FusedAttnBackend.F16_arbitrary_seqlen + if supported + else FusedAttnBackend.No_Backend + ) + + if cudnn_version < 8900 and backend == FusedAttnBackend.F16_arbitrary_seqlen: + backend = FusedAttnBackend.No_Backend + warnings.warn("FP16/BF16 fused attention requires cuDNN 8.9.0 or newer") + if ( + cudnn_version == 91400 + and max_seqlen_kv > 1024 + and window_size_left != -1 + and attn_mask_type not in ("causal", "causal_bottom_right") + ): + backend = FusedAttnBackend.No_Backend + warnings.warn( + "This non-causal sliding-window configuration requires cuDNN > 9.14.0" + ) + if ( + cudnn_version <= 91500 + and is_training + and standard_format + and max_seqlen_kv % 128 != 0 + and cuda_graph + and attn_mask_type + not in ("padding", "padding_causal", "padding_causal_bottom_right") + ): + backend = FusedAttnBackend.No_Backend + warnings.warn( + "This backward CUDA-graph configuration requires cuDNN 9.15.1 or newer" + ) + if backend == FusedAttnBackend.F16_arbitrary_seqlen and sm_arch == 120: + if cudnn_version < 91801: + backend = FusedAttnBackend.No_Backend + warnings.warn("SM120 fused attention requires cuDNN 9.18.1 or newer") + elif deterministic and is_training: + backend = FusedAttnBackend.No_Backend + warnings.warn( + "Deterministic fused-attention backward is not supported on SM120" + ) + elif qkv_layout in ("t3hd", "th3d"): + backend = FusedAttnBackend.No_Backend + warnings.warn("T3HD/TH3D fused attention is not supported on SM120") + return backend diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py new file mode 100644 index 00000000000..b276613f64e --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py @@ -0,0 +1,185 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Shared cuDNN Frontend Python graph runtime for PyTorch attention.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +import importlib +import threading +from typing import Any, Dict, Hashable, Optional, Tuple + +import torch + +_thread_state = threading.local() + + +def import_cudnn_frontend(): + """Import the cuDNN Frontend Python package lazily. + + PyTorch FusedAttention is optional at import time, so importing Transformer + Engine must not eagerly initialize cuDNN or fail on CPU-only processes. + """ + + try: + return importlib.import_module("cudnn") + except ImportError as exc: + raise ImportError( + "cuDNN Frontend Python package not found. Install " + "nvidia-cudnn-frontend>=1.27.0." + ) from exc + + +def _device_key(device: torch.device) -> Tuple[str, Optional[int]]: + device = torch.device(device) + if device.type == "cuda" and device.index is None: + return (device.type, torch.cuda.current_device()) + return (device.type, device.index) + + +def current_stream_handle(device: torch.device): + """Return this thread's cuDNN handle, bound to PyTorch's current stream.""" + + device = torch.device(device) + if device.type != "cuda": + raise ValueError(f"cuDNN attention only supports CUDA tensors, got {device}.") + device = torch.device("cuda", _device_key(device)[1]) + cudnn = import_cudnn_frontend() + + handles = getattr(_thread_state, "handles", None) + if handles is None: + handles = {} + _thread_state.handles = handles + handle = handles.get(device) + with torch.cuda.device(device): + if handle is None: + handle = cudnn.create_handle() + handles[device] = handle + cudnn.set_stream( + handle=handle, + stream=torch.cuda.current_stream(device).cuda_stream, + ) + return handle + + +def torch_to_cudnn_dtype(dtype: torch.dtype): + """Map a PyTorch scalar dtype to a cuDNN Frontend dtype.""" + + cudnn = import_cudnn_frontend() + mapping = { + torch.float16: cudnn.data_type.HALF, + torch.bfloat16: cudnn.data_type.BFLOAT16, + torch.float32: cudnn.data_type.FLOAT, + torch.int32: cudnn.data_type.INT32, + torch.int64: cudnn.data_type.INT64, + torch.uint8: cudnn.data_type.UINT8, + torch.float8_e4m3fn: cudnn.data_type.FP8_E4M3, + torch.float8_e5m2: cudnn.data_type.FP8_E5M2, + } + try: + return mapping[dtype] + except KeyError as exc: + raise ValueError(f"Unsupported cuDNN graph tensor dtype {dtype}.") from exc + + +def make_graph(io_dtype: Any, device: torch.device, *, name: str): + """Create an SDPA graph using FP32 intermediate and compute types.""" + + cudnn = import_cudnn_frontend() + return cudnn.pygraph( + name=name, + io_data_type=io_dtype, + intermediate_data_type=cudnn.data_type.FLOAT, + compute_data_type=cudnn.data_type.FLOAT, + handle=current_stream_handle(device), + ) + + +def finalize_graph(graph) -> int: + """Build a cuDNN graph and return its required workspace size.""" + + cudnn = import_cudnn_frontend() + graph.validate() + graph.build_operation_graph() + try: + graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError(f"cuDNN attention graph is not supported: {exc}") from exc + graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) + return max(int(graph.get_workspace_size()), 1) + + +@dataclass +class GraphEntry: + """Built graph plus named graph tensors and stream-local workspaces.""" + + graph: Any + tensors: Dict[str, Any] + workspace_size: int + _workspaces: Dict[int, torch.Tensor] = field(default_factory=dict, repr=False) + + def workspace(self, device: torch.device) -> torch.Tensor: + """Get a stable workspace for the current stream. + + A workspace may be reused by asynchronous launches on one stream, but + must not be shared by independent streams. CUDA graph capture also + keeps the allocation alive for CUDA graph replay. PyTorch's caching + allocator is capture-aware, so a stream-specific workspace may be + created on the first captured invocation after the graph itself has + been warmed and cached. + """ + + device = torch.device(device) + stream = torch.cuda.current_stream(device) + stream_key = int(stream.cuda_stream) + workspace = self._workspaces.get(stream_key) + if workspace is None: + workspace = torch.empty( + self.workspace_size, + dtype=torch.uint8, + device=device, + ) + self._workspaces[stream_key] = workspace + return workspace + + def execute(self, variant_pack: Dict[Any, Any], device: torch.device) -> None: + """Execute the graph on PyTorch's current stream.""" + + self.graph.execute( + variant_pack, + self.workspace(device), + handle=current_stream_handle(device), + ) + + +def graph_cache() -> Dict[Hashable, GraphEntry]: + """Return a thread-local graph cache.""" + + cache = getattr(_thread_state, "graph_cache", None) + if cache is None: + cache = {} + _thread_state.graph_cache = cache + return cache + + +def get_graph_entry(key: Hashable) -> Optional[GraphEntry]: + """Look up a graph in this thread's cache.""" + + return graph_cache().get(key) + + +def put_graph_entry(key: Hashable, entry: GraphEntry) -> GraphEntry: + """Insert and return a graph cache entry.""" + + graph_cache()[key] = entry + return entry + + +def clear_graph_cache() -> None: + """Clear thread-local handles and graphs. Intended for tests.""" + + _thread_state.handles = {} + _thread_state.graph_cache = {} diff --git a/transformer_engine/pytorch/attention/dot_product_attention/backends.py b/transformer_engine/pytorch/attention/dot_product_attention/backends.py index 8a219a6a4df..2034fd6f527 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/backends.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/backends.py @@ -40,7 +40,7 @@ QKVLayouts, dist_group_type, ) -from transformer_engine.pytorch.cpp_extensions.fused_attn import ( +from transformer_engine.pytorch.attention.dot_product_attention.cudnn_attention import ( fused_attn_fwd, fused_attn_bwd, FusedAttnBackend, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index ea89ca97eb7..6d566e01e8f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -14,7 +14,7 @@ nvtx_range_push, get_device_compute_capability, ) -from transformer_engine.pytorch.cpp_extensions.fused_attn import ( +from transformer_engine.pytorch.attention.dot_product_attention.cudnn_attention import ( fused_attn_fwd, fused_attn_bwd, FusedAttnBackend, diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py new file mode 100644 index 00000000000..ef2436228c2 --- /dev/null +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -0,0 +1,2833 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""PyTorch implementation of cuDNN-backed scaled dot-product attention. + +All cuDNN graph construction and execution in this module goes through the +``nvidia-cudnn-frontend`` Python API. TE common retains its C++ implementation +for other framework frontends, but PyTorch does not call it. +""" + +from __future__ import annotations + +from enum import IntEnum +import math +from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Union + +import torch +# The extension is used only to reserve PyTorch's graph-safe Philox state; +# graph construction, backend selection, and execution do not call TE common. +import transformer_engine_torch as tex + +from transformer_engine.pytorch.constants import ( + DType, + FP8BwdTensorIdx, + FP8FwdTensorIdx, + TE_DType_To_Torch, +) +from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage +from transformer_engine.pytorch.tensor.float8_tensor import ( + Float8CurrentScalingQuantizer, + Float8Quantizer, +) +from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer +from transformer_engine.pytorch.tensor.storage.float8_tensor_storage import ( + Float8TensorStorage, +) +from transformer_engine.pytorch.tensor.storage.mxfp8_tensor_storage import ( + MXFP8TensorStorage, +) + +from ._cudnn_graph import ( + GraphEntry, + finalize_graph, + get_graph_entry, + import_cudnn_frontend, + make_graph, + put_graph_entry, + torch_to_cudnn_dtype, +) + +__all__ = [ + "FusedAttnBackend", + "fused_attn_fwd", + "fused_attn_bwd", + "META_QKV", + "META_DQKV", + "META_O", + "META_DO", + "META_S", + "META_DP", +] + + +class FusedAttnBackend(IntEnum): + """PyTorch cuDNN attention implementation families. + + The numeric values intentionally preserve the historical TE-common ABI so + cached/backend-selection state and external Python callers remain source + compatible after removal of the pybind enum. + """ + + No_Backend = -1 + F16_arbitrary_seqlen = 1 + FP8 = 2 + + @classmethod + def cast(cls, backend: Union["FusedAttnBackend", int, Any]) -> "FusedAttnBackend": + if isinstance(backend, cls): + return backend + return cls(int(backend)) + + +META_QKV = FP8FwdTensorIdx.GEMM1_OUTPUT +META_DQKV = FP8BwdTensorIdx.GRAD_OUTPUT1 +META_O = FP8FwdTensorIdx.GEMM2_INPUT +META_DO = FP8BwdTensorIdx.GRAD_INPUT2 +META_S = FP8FwdTensorIdx.GEMM3_OUTPUT +META_DP = FP8BwdTensorIdx.GRAD_INPUT3 + +_F16_RNG_ELTS_PER_THREAD = 16 +_FP8_THREADS_PER_CTA = 128 + + +def _is_float8_tensor(tensor: Any) -> bool: + return isinstance(tensor, Float8TensorStorage) + + +def _is_mxfp8_tensor(tensor: Any) -> bool: + return isinstance(tensor, MXFP8TensorStorage) + + +def _quantized_data(tensor: Any, *, columnwise: bool = False) -> torch.Tensor: + if _is_float8_tensor(tensor): + data = tensor._transpose if columnwise else tensor._data + elif _is_mxfp8_tensor(tensor): + data = tensor._columnwise_data if columnwise else tensor._rowwise_data + else: + data = tensor + if data is None: + orientation = "columnwise" if columnwise else "rowwise" + raise ValueError(f"Attention input has no {orientation} data buffer.") + return data + + +def _quantized_scale_inv(tensor: Any, *, columnwise: bool = False) -> torch.Tensor: + if _is_float8_tensor(tensor): + return tensor._scale_inv + if _is_mxfp8_tensor(tensor): + scale = ( + tensor._columnwise_scale_inv if columnwise else tensor._rowwise_scale_inv + ) + if scale is None: + orientation = "columnwise" if columnwise else "rowwise" + raise ValueError( + f"MXFP8 attention input has no {orientation} scale-inverse buffer." + ) + return scale + raise TypeError(f"Expected an FP8 attention tensor, got {type(tensor).__name__}.") + + +def _fp8_cudnn_dtype(tensor: Any): + return torch_to_cudnn_dtype(TE_DType_To_Torch[DType.cast(tensor._fp8_dtype)]) + + +def _scalar_graph_tensor(graph, cudnn, name: str): + return graph.tensor( + name=name, + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.FLOAT, + ) + + +def _constant_graph_tensor(graph, cudnn, name: str): + """Create a scalar graph input that callers bind to a constant value.""" + return _scalar_graph_tensor(graph, cudnn, name) + + +def _format_stride(batch: int, heads: int, seqlen: int, dim: int, tensor_format: str): + if tensor_format in ("bshd", "thd"): + return (seqlen * heads * dim, dim, heads * dim, 1) + if tensor_format == "sbhd": + return (heads * dim, dim, batch * heads * dim, 1) + if tensor_format == "bhsd": + return (heads * seqlen * dim, seqlen * dim, dim, 1) + raise ValueError(f"Unsupported FP8 tensor format {tensor_format!r}.") + + +def _round_up(value: int, multiple: int) -> int: + return (value + multiple - 1) // multiple * multiple + + +def _max_ragged_tokens(num_tokens: int) -> int: + """Quantize THD token counts to the buckets used by TE's cuDNN path.""" + + if num_tokens <= 1024: + return 1024 + if num_tokens <= 32768: + return 1 << (num_tokens - 1).bit_length() + return _round_up(num_tokens, 32768) + + +def _max_ragged_batch(batch: int) -> int: + """Quantize THD batch sizes to the buckets used by TE's cuDNN path.""" + + if batch <= 32: + return 32 + if batch <= 512: + return 1 << (batch - 1).bit_length() + return _round_up(batch, 512) + + +def _padded_sequence_lengths(cu_seqlens: torch.Tensor, batch: int) -> torch.Tensor: + lengths = _sequence_lengths(cu_seqlens) + if lengths.numel() == batch: + return lengths + padding = torch.zeros( + batch - lengths.numel(), dtype=lengths.dtype, device=lengths.device + ) + return torch.cat((lengths, padding)) + + +def _element_ragged_offsets( + cu_seqlens_padded: torch.Tensor, + batch: int, + multiplier: int, +) -> torch.Tensor: + """Convert token offsets to padded int64 element offsets for legacy SDPA graphs.""" + + offsets = cu_seqlens_padded.to(dtype=torch.int64) + if offsets.numel() < batch + 1: + tail = offsets[-1:].expand(batch + 1 - offsets.numel()) + offsets = torch.cat((offsets, tail)) + return offsets * multiplier + + +def _mxfp8_padded_sizes(s_q: int, s_kv: int, d_qk: int, d_v: int) -> Dict[str, int]: + return { + "s_q_padded": _round_up(s_q, 128), + "s_kv_padded": _round_up(s_kv, 128), + "s_q_scale_padded": _round_up((s_q + 31) // 32, 4), + "s_kv_scale_padded": _round_up((s_kv + 31) // 32, 4), + "d_qk_padded": _round_up(d_qk, 128), + "d_v_padded": _round_up(d_v, 128), + "d_qk_scale_padded": _round_up((d_qk + 31) // 32, 4), + "d_v_scale_padded": _round_up((d_v + 31) // 32, 4), + } + + +def _make_mxfp8_scale_tensor( + graph, + cudnn, + *, + name: str, + batch: int, + heads: int, + seqlen: int, + dim: int, + tensor_format: str, +): + return graph.tensor( + name=name, + dim=(batch, heads, seqlen, dim), + stride=_format_stride(batch, heads, seqlen, dim, tensor_format), + data_type=cudnn.data_type.FP8_E8M0, + ).set_reordering_type(cudnn.tensor_reordering.F8_128x4) + + +def _make_float8_output(quantizer, shape, fake_dtype, device): + data = torch.empty(shape, dtype=torch.uint8, device=device) + return quantizer.create_tensor_from_data( + data, + fake_dtype=fake_dtype, + internal=bool(getattr(quantizer, "internal", False)), + ) + + +def _allocate_fp8_kernel_output(quantizer, shape, fake_dtype, device): + """Allocate the FP8 SDPA output and any hidden amax buffer. + + Delayed scaling asks cuDNN to produce FP8 directly. Current scaling and + MXFP8 preserve the historical attention contract and ask cuDNN for a + high-precision output, which the caller may quantize afterwards. + """ + if isinstance(quantizer, Float8Quantizer): + return _make_float8_output(quantizer, shape, fake_dtype, device), quantizer.amax + if isinstance(quantizer, Float8CurrentScalingQuantizer): + return torch.empty(shape, dtype=fake_dtype, device=device), torch.zeros( + 1, dtype=torch.float32, device=device + ) + if isinstance(quantizer, MXFP8Quantizer): + return torch.empty(shape, dtype=fake_dtype, device=device), None + raise TypeError( + f"Unsupported FP8 attention output quantizer {type(quantizer).__name__}." + ) + + +def _format_from_layout_component(component: str) -> str: + return "".join(char for char in component if char.isalpha()) + + +def _q_kv_formats(qkv_layout: str) -> Tuple[str, str]: + layout = qkv_layout.removeprefix("paged_kv_") + components = layout.split("_") + q_format = _format_from_layout_component(components[0]) + kv_format = ( + _format_from_layout_component(components[-1]) + if len(components) > 1 + else q_format + ) + return q_format, kv_format + + +def _is_paged_layout(qkv_layout: str) -> bool: + return qkv_layout.startswith("paged_kv_") + + +def _tensor_metadata(tensor: Optional[torch.Tensor]) -> Optional[Tuple[Any, ...]]: + if tensor is None: + return None + return ( + tuple(tensor.shape), + tuple(tensor.stride()), + tensor.dtype, + tensor.device.type, + tensor.device.index, + ) + + +def _logical_bhsd_desc( + tensor: torch.Tensor, + tensor_format: str, + *, + batch: int, + max_seqlen: int, +) -> Tuple[Tuple[int, ...], Tuple[int, ...]]: + """Describe a physical TE attention tensor as logical cuDNN BHSD.""" + + shape = tuple(tensor.shape) + stride = tuple(tensor.stride()) + if tensor_format == "sbhd": + return ( + (shape[1], shape[2], shape[0], shape[3]), + (stride[1], stride[2], stride[0], stride[3]), + ) + if tensor_format == "bshd": + return ( + (shape[0], shape[2], shape[1], shape[3]), + (stride[0], stride[2], stride[1], stride[3]), + ) + if tensor_format == "bhsd": + return (shape, stride) + if tensor_format == "thd": + # The ragged offset selects each batch's first token. The synthetic + # batch stride is only graph metadata; S/H/D use the real packed view. + return ( + (batch, shape[1], max_seqlen, shape[2]), + (max_seqlen * stride[0], stride[1], stride[0], stride[2]), + ) + raise ValueError(f"Unsupported attention tensor format {tensor_format!r}.") + + +def _make_bhsd_graph_tensor( + graph, + tensor: torch.Tensor, + tensor_format: str, + *, + batch: int, + max_seqlen: int, + data_type=None, + ragged_offset=None, + ragged_offset_multiplier: int = 1, + name: str, +): + dim, stride = _logical_bhsd_desc( + tensor, + tensor_format, + batch=batch, + max_seqlen=max_seqlen, + ) + graph_tensor = graph.tensor( + name=name, + dim=dim, + stride=stride, + data_type=data_type if data_type is not None else tensor.dtype, + ragged_offset=ragged_offset, + ragged_offset_multiplier=ragged_offset_multiplier, + ) + return graph_tensor + + +def _allocate_output( + q: torch.Tensor, + value_head_dim: int, + fake_dtype: torch.dtype, + fast_zero_fill: bool, +) -> torch.Tensor: + shape = (*q.shape[:-1], value_head_dim) + factory = torch.zeros if fast_zero_fill else torch.empty + return factory(shape, dtype=fake_dtype, device=q.device) + + +def _storage_span(tensor: torch.Tensor) -> int: + if tensor.numel() == 0: + return 0 + return 1 + sum( + (size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride()) + ) + + +def _allocate_grad_views( + inputs: Sequence[torch.Tensor], + *, + fast_zero_fill: bool, +) -> Tuple[torch.Tensor, ...]: + """Allocate gradients while preserving packed-QKV storage relationships.""" + + groups: Dict[Tuple[str, int, int], List[int]] = {} + for index, tensor in enumerate(inputs): + storage = tensor.untyped_storage() + key = (tensor.device.type, tensor.device.index or 0, storage.data_ptr()) + groups.setdefault(key, []).append(index) + + outputs: List[Optional[torch.Tensor]] = [None] * len(inputs) + for indices in groups.values(): + if len(indices) == 1: + inp = inputs[indices[0]] + out = torch.empty_strided( + inp.shape, inp.stride(), dtype=inp.dtype, device=inp.device + ) + if fast_zero_fill: + out.zero_() + outputs[indices[0]] = out + continue + + min_offset = min(inputs[index].storage_offset() for index in indices) + max_end = max( + inputs[index].storage_offset() + _storage_span(inputs[index]) + for index in indices + ) + exemplar = inputs[indices[0]] + base = torch.empty( + max_end - min_offset, dtype=exemplar.dtype, device=exemplar.device + ) + if fast_zero_fill: + base.zero_() + for index in indices: + inp = inputs[index] + outputs[index] = torch.as_strided( + base, + size=inp.shape, + stride=inp.stride(), + storage_offset=inp.storage_offset() - min_offset, + ) + + return tuple(output for output in outputs if output is not None) + + +def _reserve_philox_state( + device: torch.device, + rng_gen: Optional[torch.Generator], + increment: int, +) -> torch.Tensor: + """Reserve a Philox counter range and return CUDA ``[seed, offset]``.""" + + device = torch.device(device) + helper = getattr(tex, "get_cudnn_attention_rng_state", None) + if helper is not None: + with torch.cuda.device(device): + return helper(rng_gen, increment) + # Development-tree compatibility before the local extension has been + # rebuilt. Installed packages always expose the graph-safe helper above. + if rng_gen is None: + index = ( + device.index if device.index is not None else torch.cuda.current_device() + ) + rng_gen = torch.cuda.default_generators[index] + seed = rng_gen.initial_seed() + offset = rng_gen.get_offset() + rng_gen.set_offset(offset + increment) + return torch.tensor((seed, offset), dtype=torch.int64, device=device) + + +def _mask_options( + cudnn, + attn_mask_type: str, + window_size: Tuple[int, int], + bottom_right_diagonal: bool, + max_seqlen_q: int, + max_seqlen_kv: int, +) -> Dict[str, Any]: + is_causal = attn_mask_type in ("causal", "padding_causal") + is_bottom_right = attn_mask_type in ( + "causal_bottom_right", + "padding_causal_bottom_right", + ) + is_padding = attn_mask_type in ( + "padding", + "padding_causal", + "padding_causal_bottom_right", + ) + if is_bottom_right and max_seqlen_q == max_seqlen_kv and not is_padding: + is_causal = True + is_bottom_right = False + bottom_right_diagonal = False + + options: Dict[str, Any] = { + "use_causal_mask": is_causal, + "use_causal_mask_bottom_right": is_bottom_right, + "diagonal_alignment": ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if bottom_right_diagonal + else cudnn.diagonal_alignment.TOP_LEFT + ), + } + left, right = window_size + if left != -1: + options["diagonal_band_left_bound"] = left + 1 + # ``use_causal_mask`` already imposes a right bound of zero. The Python + # frontend rejects specifying that same bound through both attributes. + if right != -1 and not ((is_causal or is_bottom_right) and right == 0): + options["diagonal_band_right_bound"] = right + options["is_padding"] = is_padding + return options + + +def _sequence_lengths(cu_seqlens: torch.Tensor) -> torch.Tensor: + return cu_seqlens[1:] - cu_seqlens[:-1] + + +def _ragged_offset_tensor( + graph, + cu_seqlens_padded: torch.Tensor, + *, + multiplier: int, + name: str, + length: Optional[int] = None, + data_type=None, +): + return ( + graph.tensor( + name=name, + dim=(cu_seqlens_padded.numel() if length is None else length, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cu_seqlens_padded.dtype if data_type is None else data_type, + is_pass_by_value=False, + ), + multiplier, + ) + + +def _stats_layout( + *, + batch: int, + heads: int, + max_seqlen_q: int, + total_tokens_q: int, + ragged: bool, +) -> Tuple[Tuple[int, ...], Tuple[int, ...], Tuple[int, ...]]: + if ragged: + physical_shape = (total_tokens_q, heads, 1) + logical_dim = (batch, heads, max_seqlen_q, 1) + logical_stride = (heads * max_seqlen_q, 1, heads, 1) + else: + physical_shape = logical_dim = (batch, heads, max_seqlen_q, 1) + logical_stride = (heads * max_seqlen_q, max_seqlen_q, 1, 1) + return physical_shape, logical_dim, logical_stride + + +def _f16_fwd_key(**kwargs) -> Tuple[Any, ...]: + return ("f16_fwd",) + tuple(kwargs.items()) + + +def _build_f16_fwd_graph( + *, + is_training: bool, + max_seqlen_q: int, + max_seqlen_kv: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + output: torch.Tensor, + stats: Optional[torch.Tensor], + max_scores: Optional[torch.Tensor], + attn_bias: Optional[torch.Tensor], + cu_seqlens_q_padded: torch.Tensor, + cu_seqlens_kv_padded: torch.Tensor, + page_table_k: Optional[torch.Tensor], + page_table_v: Optional[torch.Tensor], + rng_state: torch.Tensor, + softmax_offset: Optional[torch.Tensor], + attn_scale: float, + dropout: float, + qkv_layout: str, + o_format: str, + attn_bias_type: str, + attn_mask_type: str, + softmax_type: str, + window_size: Tuple[int, int], + bottom_right_diagonal: bool, +) -> GraphEntry: + cudnn = import_cudnn_frontend() + graph = make_graph( + torch_to_cudnn_dtype(q.dtype), q.device, name="te_fused_attention_fwd" + ) + q_format, kv_format = _q_kv_formats(qkv_layout) + batch = cu_seqlens_q.numel() - 1 + is_ragged_q = q_format == "thd" + is_ragged_kv = kv_format == "thd" + use_ragged_stats = is_ragged_q and cudnn.backend_version() >= 90600 + use_token_buckets = ( + cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ) + use_direct_offsets = cudnn.backend_version() >= 92400 and dropout == 0.0 + use_legacy_offsets = (is_ragged_q or is_ragged_kv) and not use_direct_offsets + graph_batch = ( + _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch + ) + if not use_token_buckets: + use_ragged_stats = False + graph_seqlen_q = ( + _max_ragged_tokens(q.shape[0]) + if is_ragged_q and use_token_buckets + else max_seqlen_q + ) + graph_seqlen_kv = ( + _max_ragged_tokens(k.shape[0]) + if is_ragged_kv and use_token_buckets + else max_seqlen_kv + ) + + tensors: Dict[str, Any] = {} + tensors["_legacy_offsets"] = use_legacy_offsets + tensors["_graph_batch"] = graph_batch + offset_q = offset_o = offset_k = offset_v = offset_stats = None + if is_ragged_q: + offset_q, q_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1 if use_legacy_offsets else q.stride(0), + name="offset_q", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, + ) + offset_o, o_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1 if use_legacy_offsets else output.stride(0), + name="offset_o", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, + ) + tensors["offset_q"] = offset_q + tensors["offset_o"] = offset_o + else: + q_mult = o_mult = 1 + if is_ragged_kv: + offset_k, k_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1 if use_legacy_offsets else k.stride(0), + name="offset_k", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, + ) + offset_v, v_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1 if use_legacy_offsets else v.stride(0), + name="offset_v", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, + ) + tensors["offset_k"] = offset_k + tensors["offset_v"] = offset_v + else: + k_mult = v_mult = 1 + + q_t = _make_bhsd_graph_tensor( + graph, + q, + q_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, + ragged_offset=offset_q, + ragged_offset_multiplier=q_mult, + name="Q", + ) + + if _is_paged_layout(qkv_layout): + if kv_format == "bshd": + num_pages_k, page_size_k = k.shape[0], k.shape[1] + num_pages_v, page_size_v = v.shape[0], v.shape[1] + elif kv_format == "sbhd": + page_size_k, num_pages_k = k.shape[0], k.shape[1] + page_size_v, num_pages_v = v.shape[0], v.shape[1] + else: + raise ValueError(f"Paged attention does not support KV format {kv_format}.") + k_batch, k_seqlen = num_pages_k, page_size_k + v_batch, v_seqlen = num_pages_v, page_size_v + else: + k_batch = v_batch = graph_batch + k_seqlen = v_seqlen = max_seqlen_kv + + k_t = _make_bhsd_graph_tensor( + graph, + k, + kv_format, + batch=k_batch, + max_seqlen=graph_seqlen_kv if is_ragged_kv else k_seqlen, + ragged_offset=offset_k, + ragged_offset_multiplier=k_mult, + name="K", + ) + v_t = _make_bhsd_graph_tensor( + graph, + v, + kv_format, + batch=v_batch, + max_seqlen=graph_seqlen_kv if is_ragged_kv else v_seqlen, + ragged_offset=offset_v, + ragged_offset_multiplier=v_mult, + name="V", + ) + tensors.update(Q=q_t, K=k_t, V=v_t) + + options = _mask_options( + cudnn, + attn_mask_type, + window_size, + bottom_right_diagonal, + max_seqlen_q, + max_seqlen_kv, + ) + is_padding = options.pop("is_padding") + options.update( + generate_stats=is_training, + attn_scale=float(attn_scale), + use_padding_mask=is_padding, + use_alibi_mask=attn_bias_type == "alibi", + ) + + if attn_bias_type == "post_scale_bias": + bias_t = graph.tensor_like(attn_bias, name="Bias") + tensors["Bias"] = bias_t + options["bias"] = bias_t + + if is_padding: + seq_q = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + seq_kv = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) + seq_q_t = graph.tensor_like(seq_q, name="seq_len_q") + seq_kv_t = graph.tensor_like(seq_kv, name="seq_len_kv") + tensors["seq_len_q"] = seq_q_t + tensors["seq_len_kv"] = seq_kv_t + options["seq_len_q"] = seq_q_t + options["seq_len_kv"] = seq_kv_t + + if page_table_k is not None: + page_k_t = graph.tensor_like(page_table_k, name="page_table_k") + page_v_t = graph.tensor_like(page_table_v, name="page_table_v") + tensors["page_table_k"] = page_k_t + tensors["page_table_v"] = page_v_t + options["paged_attention_k_table"] = page_k_t + options["paged_attention_v_table"] = page_v_t + options["paged_attention_max_seq_len_kv"] = max_seqlen_kv + + if is_training and dropout != 0.0: + seed_t = graph.tensor( + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + offset_t = graph.tensor( + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + tensors["dropout_seed"] = seed_t + tensors["dropout_offset"] = offset_t + options["dropout"] = (float(dropout), seed_t, offset_t) + + if softmax_type != "vanilla": + if softmax_offset is None: + raise ValueError(f"softmax_type={softmax_type!r} requires softmax_offset.") + softmax_offset_t = graph.tensor_like(softmax_offset, name="softmax_offset") + tensors["softmax_offset"] = softmax_offset_t + options["sink_token"] = softmax_offset_t + + if max_scores is not None: + _, max_dim, max_stride = _stats_layout( + batch=graph_batch, + heads=q_t.get_dim()[1], + max_seqlen_q=graph_seqlen_q, + total_tokens_q=q.shape[0] if is_ragged_q else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + if use_ragged_stats: + offset_stats, stats_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1 if use_legacy_offsets else q_t.get_dim()[1], + name="offset_stats", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, + ) + tensors["offset_stats"] = offset_stats + else: + stats_mult = 1 + max_t = graph.tensor( + name="Max", + dim=max_dim, + stride=max_stride, + data_type=cudnn.data_type.FLOAT, + ragged_offset=offset_stats, + ragged_offset_multiplier=stats_mult, + ).set_output(True) + tensors["Max"] = max_t + options["score_max"] = max_t + + output_t, stats_t = graph.sdpa(name="te_sdpa", q=q_t, k=k_t, v=v_t, **options) + output_dim, output_stride = _logical_bhsd_desc( + output, + o_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, + ) + output_t.set_output(True).set_dim(output_dim).set_stride(output_stride) + if is_ragged_q: + output_t.set_ragged_offset(offset_o).set_ragged_offset_multiplier(o_mult) + tensors["O"] = output_t + + if is_training: + assert stats is not None + _, stats_dim, stats_stride = _stats_layout( + batch=graph_batch, + heads=output_dim[1], + max_seqlen_q=graph_seqlen_q, + total_tokens_q=q.shape[0] if is_ragged_q else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + if use_ragged_stats and offset_stats is None: + offset_stats, stats_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1 if use_legacy_offsets else output_dim[1], + name="offset_stats", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, + ) + tensors["offset_stats"] = offset_stats + stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( + stats_dim + ).set_stride(stats_stride) + if use_ragged_stats: + stats_t.set_ragged_offset(offset_stats).set_ragged_offset_multiplier( + stats_mult + ) + tensors["Stats"] = stats_t + + return GraphEntry( + graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) + ) + + +def _f16_forward( + is_training: bool, + max_seqlen_q: int, + max_seqlen_kv: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + fake_dtype: torch.dtype, + attn_bias: Optional[torch.Tensor], + cu_seqlens_q_padded: Optional[torch.Tensor], + cu_seqlens_kv_padded: Optional[torch.Tensor], + page_table_k: Optional[torch.Tensor], + page_table_v: Optional[torch.Tensor], + attn_scale: float, + dropout: float, + fast_zero_fill: bool, + qkv_layout: str, + o_format: str, + attn_bias_type: str, + attn_mask_type: str, + softmax_type: str, + window_size: Tuple[int, int], + bottom_right_diagonal: bool, + rng_gen: Optional[torch.Generator], + softmax_offset: Optional[torch.Tensor], + return_max_logit: bool, +) -> Tuple[torch.Tensor, List[torch.Tensor], Optional[torch.Tensor]]: + q_format, _ = _q_kv_formats(qkv_layout) + batch = cu_seqlens_q.numel() - 1 + heads = q.shape[1] if q_format == "bhsd" else q.shape[-2] + total_tokens_q = q.shape[0] if q_format == "thd" else batch * max_seqlen_q + cudnn = import_cudnn_frontend() + ragged_stats = ( + q_format == "thd" + and cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ) + output_shape = ( + (q.shape[0], heads, v.shape[-1]) + if o_format == "thd" + else _fp8_output_shape(batch, max_seqlen_q, heads, v.shape[-1], o_format) + ) + output_factory = torch.zeros if fast_zero_fill else torch.empty + output = output_factory(output_shape, dtype=fake_dtype, device=q.device) + stats_shape, _, _ = _stats_layout( + batch=batch, + heads=heads, + max_seqlen_q=max_seqlen_q, + total_tokens_q=total_tokens_q, + ragged=ragged_stats, + ) + stats = ( + torch.empty(stats_shape, dtype=torch.float32, device=q.device) + if is_training + else None + ) + max_scores = ( + torch.empty(stats_shape, dtype=torch.float32, device=q.device) + if return_max_logit + else None + ) + rng_state = _reserve_philox_state(q.device, rng_gen, _F16_RNG_ELTS_PER_THREAD) + cu_seqlens_q_padded = ( + cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded + ) + cu_seqlens_kv_padded = ( + cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded + ) + + key = _f16_fwd_key( + is_training=is_training, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + q=_tensor_metadata(q), + k=_tensor_metadata(k), + v=_tensor_metadata(v), + output=_tensor_metadata(output), + stats=_tensor_metadata(stats), + bias=_tensor_metadata(attn_bias), + qkv_layout=qkv_layout, + o_format=o_format, + attn_scale=float(attn_scale), + dropout=float(dropout), + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=tuple(window_size), + bottom_right_diagonal=bottom_right_diagonal, + return_max_logit=return_max_logit, + paged=page_table_k is not None, + ) + entry = get_graph_entry(key) + if entry is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "cuDNN attention graph must be built before CUDA graph capture." + ) + entry = _build_f16_fwd_graph( + is_training=is_training, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + q=q, + k=k, + v=v, + output=output, + stats=stats, + max_scores=max_scores, + attn_bias=attn_bias, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + page_table_k=page_table_k, + page_table_v=page_table_v, + rng_state=rng_state, + softmax_offset=softmax_offset, + attn_scale=attn_scale, + dropout=dropout, + qkv_layout=qkv_layout, + o_format=o_format, + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + ) + put_graph_entry(key, entry) + + tensors = entry.tensors + variant_pack: Dict[Any, Any] = { + tensors["Q"]: q, + tensors["K"]: k, + tensors["V"]: v, + tensors["O"]: output, + } + if is_training: + variant_pack[tensors["Stats"]] = stats + if attn_bias_type == "post_scale_bias": + variant_pack[tensors["Bias"]] = attn_bias + legacy_offsets = tensors["_legacy_offsets"] + graph_batch = tensors["_graph_batch"] + if "seq_len_q" in tensors: + variant_pack[tensors["seq_len_q"]] = _padded_sequence_lengths( + cu_seqlens_q, graph_batch + ) + variant_pack[tensors["seq_len_kv"]] = _padded_sequence_lengths( + cu_seqlens_kv, graph_batch + ) + if "offset_q" in tensors: + if legacy_offsets: + variant_pack[tensors["offset_q"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, q.stride(0) + ) + variant_pack[tensors["offset_o"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, output.stride(0) + ) + else: + variant_pack[tensors["offset_q"]] = cu_seqlens_q_padded + variant_pack[tensors["offset_o"]] = cu_seqlens_q_padded + if "offset_k" in tensors: + if legacy_offsets: + variant_pack[tensors["offset_k"]] = _element_ragged_offsets( + cu_seqlens_kv_padded, graph_batch, k.stride(0) + ) + variant_pack[tensors["offset_v"]] = _element_ragged_offsets( + cu_seqlens_kv_padded, graph_batch, v.stride(0) + ) + else: + variant_pack[tensors["offset_k"]] = cu_seqlens_kv_padded + variant_pack[tensors["offset_v"]] = cu_seqlens_kv_padded + if "offset_stats" in tensors: + variant_pack[tensors["offset_stats"]] = ( + _element_ragged_offsets(cu_seqlens_q_padded, graph_batch, heads) + if legacy_offsets + else cu_seqlens_q_padded + ) + if "page_table_k" in tensors: + variant_pack[tensors["page_table_k"]] = page_table_k + variant_pack[tensors["page_table_v"]] = page_table_v + if "dropout_seed" in tensors: + variant_pack[tensors["dropout_seed"]] = rng_state[:1] + variant_pack[tensors["dropout_offset"]] = rng_state[1:] + if "softmax_offset" in tensors: + variant_pack[tensors["softmax_offset"]] = softmax_offset + if "Max" in tensors: + variant_pack[tensors["Max"]] = max_scores + entry.execute(variant_pack, q.device) + + aux: List[torch.Tensor] = [] + if is_training: + aux.append(stats) + aux.append(rng_state) + if attn_bias_type not in ("no_bias", "alibi"): + aux.append(attn_bias) + if softmax_type != "vanilla": + aux.append(softmax_offset) + + max_logit = None + if return_max_logit: + if q_format == "thd" and max_scores.ndim == 4: + seqlens_q = _sequence_lengths(cu_seqlens_q).to(device=max_scores.device) + sq_idx = torch.arange(max_scores.shape[2], device=max_scores.device).view( + 1, 1, -1, 1 + ) + valid = sq_idx < seqlens_q.view(-1, 1, 1, 1) + max_scores_for_reduce = max_scores.masked_fill(~valid, float("-inf")) + else: + max_scores_for_reduce = max_scores + reduce_dims = (0, 2) if max_scores_for_reduce.ndim == 3 else (0, 2, 3) + max_logit = torch.amax(max_scores_for_reduce, dim=reduce_dims).to(output.dtype) + return output, aux, max_logit + + +def _fp8_output_shape( + batch: int, seqlen: int, heads: int, dim: int, tensor_format: str +): + if tensor_format == "bshd": + return (batch, seqlen, heads, dim) + if tensor_format == "sbhd": + return (seqlen, batch, heads, dim) + if tensor_format == "bhsd": + return (batch, heads, seqlen, dim) + raise ValueError(f"FP8 attention does not support output format {tensor_format!r}.") + + +def _allocate_attention_grad_data( + *, + batch: int, + heads: int, + kv_heads: int, + max_seqlen_q: int, + max_seqlen_kv: int, + head_dim_qk: int, + head_dim_v: int, + dqkv_layout: str, + dtype: torch.dtype, + device: torch.device, + zero: bool, +): + """Allocate dQ/dK/dV buffers with the requested packed storage layout.""" + + layout = dqkv_layout.removeprefix("paged_kv_") + components = layout.split("_") + q_format, kv_format = _q_kv_formats(layout) + q_shape = _fp8_output_shape(batch, max_seqlen_q, heads, head_dim_qk, q_format) + k_shape = _fp8_output_shape(batch, max_seqlen_kv, kv_heads, head_dim_qk, kv_format) + v_shape = _fp8_output_shape(batch, max_seqlen_kv, kv_heads, head_dim_v, kv_format) + factory = torch.zeros if zero else torch.empty + + if len(components) == 1: + if ( + head_dim_qk != head_dim_v + or heads != kv_heads + or max_seqlen_q != max_seqlen_kv + ): + raise ValueError( + f"Packed QKV gradient layout {layout!r} requires matching Q/K/V." + ) + packed_dim = components[0].index("3") + packed_shape = list(q_shape) + packed_shape.insert(packed_dim, 3) + packed = factory(packed_shape, dtype=dtype, device=device) + return tuple(packed.select(packed_dim, index) for index in range(3)) + if len(components) == 2: + if head_dim_qk != head_dim_v: + raise ValueError( + f"Packed KV gradient layout {layout!r} requires dQK == dV." + ) + q_out = factory(q_shape, dtype=dtype, device=device) + packed_dim = components[1].index("2") + packed_shape = list(k_shape) + packed_shape.insert(packed_dim, 2) + packed = factory(packed_shape, dtype=dtype, device=device) + return q_out, packed.select(packed_dim, 0), packed.select(packed_dim, 1) + return tuple( + factory(shape, dtype=dtype, device=device) + for shape in (q_shape, k_shape, v_shape) + ) + + +def _wrap_float8_grad_outputs(quantizer, tensors, fake_dtype): + return tuple( + quantizer.create_tensor_from_data( + tensor, + fake_dtype=fake_dtype, + internal=bool(getattr(quantizer, "internal", False)), + ) + for tensor in tensors + ) + + +def _build_fp8_fwd_graph( + *, + max_seqlen_q, + max_seqlen_kv, + q, + k, + v, + output, + stats, + amax_s, + amax_o, + s_quantizer, + o_quantizer, + qkv_layout, + o_format, + qkv_scale_inv_format, + attn_scale, + dropout, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + softmax_offset, + cu_seqlens_q, + cu_seqlens_kv, +): + cudnn = import_cudnn_frontend() + q_data = _quantized_data(q) + k_data = _quantized_data(k) + v_data = _quantized_data(v, columnwise=_is_mxfp8_tensor(v)) + q_format, kv_format = _q_kv_formats(qkv_layout) + batch = cu_seqlens_q.numel() - 1 + heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] + kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] + d_qk = q.shape[-1] + d_v = v.shape[-1] + graph = make_graph(_fp8_cudnn_dtype(q), q.device, name="te_fp8_sdpa_fwd") + tensors: Dict[str, Any] = {} + + q_t = _make_bhsd_graph_tensor( + graph, + q_data, + q_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=_fp8_cudnn_dtype(q), + name="Q", + ) + k_t = _make_bhsd_graph_tensor( + graph, + k_data, + kv_format, + batch=batch, + max_seqlen=max_seqlen_kv, + data_type=_fp8_cudnn_dtype(k), + name="K", + ) + v_t = _make_bhsd_graph_tensor( + graph, + v_data, + kv_format, + batch=batch, + max_seqlen=max_seqlen_kv, + data_type=_fp8_cudnn_dtype(v), + name="V", + ) + tensors.update(Q=q_t, K=k_t, V=v_t) + + options = _mask_options( + cudnn, + attn_mask_type, + window_size, + bottom_right_diagonal, + max_seqlen_q, + max_seqlen_kv, + ) + is_padding = options.pop("is_padding") + options.update( + generate_stats=True, + attn_scale=float(attn_scale), + use_padding_mask=is_padding, + ) + if is_padding: + seq_q = _sequence_lengths(cu_seqlens_q) + seq_kv = _sequence_lengths(cu_seqlens_kv) + seq_q_t = graph.tensor_like(seq_q, name="seq_len_q") + seq_kv_t = graph.tensor_like(seq_kv, name="seq_len_kv") + tensors.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) + options.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) + if dropout != 0.0: + seed_t = graph.tensor( + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + offset_t = graph.tensor( + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + tensors.update(dropout_seed=seed_t, dropout_offset=offset_t) + options["dropout"] = (float(dropout), seed_t, offset_t) + if softmax_type != "vanilla": + sink_t = graph.tensor_like(softmax_offset, name="softmax_offset") + tensors["softmax_offset"] = sink_t + options["sink_token"] = sink_t + + if _is_mxfp8_tensor(q): + # The MXFP8 binding has no padding-mask keyword. Avoid forwarding the + # false default through the generic graph capture layer. + options.pop("use_padding_mask", None) + if is_padding: + raise RuntimeError( + "The installed cuDNN Frontend Python MXFP8 graph API does not expose " + "padding sequence lengths." + ) + scale_format_q = qkv_scale_inv_format or q_format + scale_format_kv = qkv_scale_inv_format or kv_format + padded = _mxfp8_padded_sizes(max_seqlen_q, max_seqlen_kv, d_qk, d_v) + descale_q = _make_mxfp8_scale_tensor( + graph, + cudnn, + name="Descale_Q", + batch=batch, + heads=heads, + seqlen=padded["s_q_padded"], + dim=padded["d_qk_scale_padded"], + tensor_format=scale_format_q, + ) + descale_k = _make_mxfp8_scale_tensor( + graph, + cudnn, + name="Descale_K", + batch=batch, + heads=kv_heads, + seqlen=padded["s_kv_padded"], + dim=padded["d_qk_scale_padded"], + tensor_format=scale_format_kv, + ) + descale_v = _make_mxfp8_scale_tensor( + graph, + cudnn, + name="Descale_V", + batch=batch, + heads=kv_heads, + seqlen=padded["s_kv_scale_padded"], + dim=padded["d_v_padded"], + tensor_format=scale_format_kv, + ) + tensors.update(descale_q=descale_q, descale_k=descale_k, descale_v=descale_v) + output_t, stats_t, amax_o_t = graph.sdpa_mxfp8( + q_t, + k_t, + v_t, + descale_q, + descale_k, + descale_v, + name="te_sdpa_mxfp8", + **options, + ) + amax_o_t.set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) + else: + if "diagonal_band_left_bound" in options: + options["left_bound"] = options.pop("diagonal_band_left_bound") + if "diagonal_band_right_bound" in options: + options["right_bound"] = options.pop("diagonal_band_right_bound") + descale_q = _scalar_graph_tensor(graph, cudnn, "Descale_Q") + descale_k = _scalar_graph_tensor(graph, cudnn, "Descale_K") + descale_v = _scalar_graph_tensor(graph, cudnn, "Descale_V") + tensors.update(descale_q=descale_q, descale_k=descale_k, descale_v=descale_v) + if isinstance(s_quantizer, Float8Quantizer): + descale_s = _scalar_graph_tensor(graph, cudnn, "Descale_S") + scale_s = _scalar_graph_tensor(graph, cudnn, "Scale_S") + tensors.update(descale_s=descale_s, scale_s=scale_s) + else: + descale_s = _constant_graph_tensor(graph, cudnn, "Current_Descale_S") + scale_s = _constant_graph_tensor(graph, cudnn, "Current_Scale_S") + tensors.update(constant_descale_s=descale_s, constant_scale_s=scale_s) + if isinstance(o_quantizer, Float8Quantizer): + scale_o = _scalar_graph_tensor(graph, cudnn, "Scale_O") + tensors["scale_o"] = scale_o + else: + scale_o = _constant_graph_tensor(graph, cudnn, "Current_Scale_O") + tensors["constant_scale_o"] = scale_o + output_t, stats_t, amax_s_t, amax_o_t = graph.sdpa_fp8( + q_t, + k_t, + v_t, + descale_q, + descale_k, + descale_v, + descale_s, + scale_s, + scale_o, + name="te_sdpa_fp8", + **options, + ) + amax_s_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) + amax_o_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) + tensors.update(amax_s=amax_s_t, amax_o=amax_o_t) + + output_t.set_output(True).set_dim((batch, heads, max_seqlen_q, d_v)).set_stride( + _format_stride(batch, heads, max_seqlen_q, d_v, o_format) + ) + output_dtype = ( + _fp8_cudnn_dtype(output) if _is_float8_tensor(output) else output.dtype + ) + output_t.set_data_type(output_dtype) + stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( + (batch, heads, max_seqlen_q, 1) + ).set_stride((heads * max_seqlen_q, max_seqlen_q, 1, 1)) + tensors.update(O=output_t, Stats=stats_t) + return GraphEntry( + graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) + ) + + +def _fp8_forward( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + s_quantizer, + o_quantizer, + attn_scale, + dropout, + fast_zero_fill, + qkv_layout, + o_format, + qkv_scale_inv_format, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + rng_gen, + softmax_offset, +): + if not isinstance(q, QuantizedTensorStorage): + raise TypeError( + "The FP8 cuDNN attention backend requires quantized Q/K/V tensors." + ) + if _is_mxfp8_tensor(q) and "padding" in attn_mask_type: + # cuDNN Frontend 1.27's Python sdpa_mxfp8 forward binding omits the + # seq_len inputs that its C++ graph API exposes. Preserve functional + # padding semantics through the same Python graph API by dequantizing + # MXFP8 inputs and using the BF16/FP16 SDPA node for this configuration. + # Remove this fallback once the Python MXFP8 forward signature exposes + # padding sequence lengths. + q_hp, k_hp, v_hp = (tensor.dequantize(dtype=fake_dtype) for tensor in (q, k, v)) + output, aux, _ = _f16_forward( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q_hp, + k_hp, + v_hp, + fake_dtype, + None, + None, + None, + None, + None, + attn_scale, + dropout, + fast_zero_fill, + qkv_layout, + o_format, + "no_bias", + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + rng_gen, + softmax_offset, + False, + ) + return output, aux + del is_training, fast_zero_fill + q_format, _ = _q_kv_formats(qkv_layout) + batch = cu_seqlens_q.numel() - 1 + heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] + d_v = v.shape[-1] + output_shape = _fp8_output_shape(batch, max_seqlen_q, heads, d_v, o_format) + output, amax_o = _allocate_fp8_kernel_output( + o_quantizer, output_shape, fake_dtype, q.device + ) + stats = torch.empty( + (batch, heads, max_seqlen_q, 1), dtype=torch.float32, device=q.device + ) + amax_s = ( + s_quantizer.amax + if isinstance(s_quantizer, Float8Quantizer) + else ( + torch.zeros(1, dtype=torch.float32, device=q.device) + if not _is_mxfp8_tensor(q) + else None + ) + ) + rng_elts = ( + max_seqlen_q * max_seqlen_q + _FP8_THREADS_PER_CTA - 1 + ) // _FP8_THREADS_PER_CTA + rng_state = _reserve_philox_state(q.device, rng_gen, rng_elts) + + q_data = _quantized_data(q) + k_data = _quantized_data(k) + v_data = _quantized_data(v, columnwise=_is_mxfp8_tensor(v)) + output_data = _quantized_data(output) if _is_float8_tensor(output) else output + key = ( + "fp8_fwd", + max_seqlen_q, + max_seqlen_kv, + _tensor_metadata(q_data), + _tensor_metadata(k_data), + _tensor_metadata(v_data), + _tensor_metadata(output_data), + type(q).__name__, + type(o_quantizer).__name__, + qkv_layout, + o_format, + qkv_scale_inv_format, + float(attn_scale), + float(dropout), + attn_mask_type, + softmax_type, + tuple(window_size), + bottom_right_diagonal, + ) + entry = get_graph_entry(key) + if entry is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "cuDNN FP8 attention graph must be built before CUDA graph capture." + ) + entry = _build_fp8_fwd_graph( + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + q=q, + k=k, + v=v, + output=output, + stats=stats, + amax_s=amax_s, + amax_o=amax_o, + s_quantizer=s_quantizer, + o_quantizer=o_quantizer, + qkv_layout=qkv_layout, + o_format=o_format, + qkv_scale_inv_format=qkv_scale_inv_format, + attn_scale=attn_scale, + dropout=dropout, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + softmax_offset=softmax_offset, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + ) + put_graph_entry(key, entry) + + t = entry.tensors + variant_pack = { + t["Q"]: q_data, + t["K"]: k_data, + t["V"]: v_data, + t["O"]: output_data, + t["Stats"]: stats, + t["descale_q"]: _quantized_scale_inv(q), + t["descale_k"]: _quantized_scale_inv(k), + t["descale_v"]: _quantized_scale_inv(v, columnwise=_is_mxfp8_tensor(v)), + } + if "descale_s" in t: + variant_pack[t["descale_s"]] = torch.reciprocal(s_quantizer.scale) + variant_pack[t["scale_s"]] = s_quantizer.scale + if "scale_o" in t: + variant_pack[t["scale_o"]] = o_quantizer.scale + one = None + for name in ("constant_descale_s", "constant_scale_s", "constant_scale_o"): + if name in t: + if one is None: + one = torch.ones(1, dtype=torch.float32, device=q.device) + variant_pack[t[name]] = one + if "amax_s" in t: + variant_pack[t["amax_s"]] = amax_s + variant_pack[t["amax_o"]] = amax_o + if "seq_len_q" in t: + variant_pack[t["seq_len_q"]] = _sequence_lengths(cu_seqlens_q) + variant_pack[t["seq_len_kv"]] = _sequence_lengths(cu_seqlens_kv) + if "dropout_seed" in t: + variant_pack[t["dropout_seed"]] = rng_state[:1] + variant_pack[t["dropout_offset"]] = rng_state[1:] + if "softmax_offset" in t: + variant_pack[t["softmax_offset"]] = softmax_offset + entry.execute(variant_pack, q.device) + aux = [stats, rng_state] + if softmax_type != "vanilla": + aux.append(softmax_offset) + return output, aux + + +def fused_attn_fwd( + is_training: bool, + max_seqlen_q: int, + max_seqlen_kv: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + fake_dtype: torch.dtype, + fused_attention_backend: FusedAttnBackend, + attn_bias: torch.Tensor = None, + cu_seqlens_q_padded: torch.Tensor = None, + cu_seqlens_kv_padded: torch.Tensor = None, + page_table_k: torch.Tensor = None, + page_table_v: torch.Tensor = None, + s_quantizer=None, + o_quantizer=None, + attn_scale: float = None, + dropout: float = 0.0, + fast_zero_fill: bool = True, + qkv_layout: str = "sbh3d", + o_format: str = "sbhd", + qkv_scale_inv_format: str = None, + attn_bias_type: str = "no_bias", + attn_mask_type: str = "padding", + softmax_type: str = "vanilla", + window_size: Tuple[int, int] = (-1, -1), + bottom_right_diagonal: bool = None, + rng_gen: torch.Generator = None, + softmax_offset: torch.Tensor = None, + return_max_logit: bool = False, + cuda_graph: bool = False, +) -> Tuple[Union[torch.Tensor, None], ...]: + """Execute fused attention through cuDNN Frontend's Python graph API.""" + + del cuda_graph + backend = FusedAttnBackend.cast(fused_attention_backend) + if backend == FusedAttnBackend.No_Backend: + raise ValueError( + "No cuDNN fused-attention backend supports this configuration." + ) + if attn_scale is None: + attn_scale = 1.0 / math.sqrt(q.size(-1)) + if bottom_right_diagonal is None: + bottom_right_diagonal = attn_mask_type in ( + "causal_bottom_right", + "padding_causal_bottom_right", + ) + if backend == FusedAttnBackend.FP8: + if page_table_k is not None or page_table_v is not None: + raise ValueError("FP8 fused attention does not support paged KV cache.") + if attn_bias_type != "no_bias" or attn_bias is not None: + raise ValueError("FP8 fused attention does not support attention bias.") + if return_max_logit: + raise ValueError( + "FP8 fused attention does not support returning maximum logits." + ) + return _fp8_forward( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + s_quantizer, + o_quantizer, + attn_scale, + dropout, + fast_zero_fill, + qkv_layout, + o_format, + qkv_scale_inv_format, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + rng_gen, + softmax_offset, + ) + + output, aux, max_logit = _f16_forward( + is_training, + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + fake_dtype, + attn_bias, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, + page_table_k, + page_table_v, + attn_scale, + dropout, + fast_zero_fill, + qkv_layout, + o_format, + attn_bias_type, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + rng_gen, + softmax_offset, + return_max_logit, + ) + if return_max_logit: + return output, aux, max_logit + return output, aux + + +def _build_f16_bwd_graph( + *, + max_seqlen_q: int, + max_seqlen_kv: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + o: torch.Tensor, + d_o: torch.Tensor, + stats: torch.Tensor, + d_q: torch.Tensor, + d_k: torch.Tensor, + d_v: torch.Tensor, + attn_bias: Optional[torch.Tensor], + d_bias: Optional[torch.Tensor], + softmax_offset: Optional[torch.Tensor], + d_softmax_offset: Optional[torch.Tensor], + cu_seqlens_q_padded: torch.Tensor, + cu_seqlens_kv_padded: torch.Tensor, + attn_scale: float, + dropout: float, + qkv_layout: str, + o_format: str, + do_format: str, + dqkv_layout: str, + attn_bias_type: str, + attn_mask_type: str, + softmax_type: str, + window_size: Tuple[int, int], + bottom_right_diagonal: bool, + deterministic: bool, +) -> GraphEntry: + cudnn = import_cudnn_frontend() + graph = make_graph( + torch_to_cudnn_dtype(q.dtype), q.device, name="te_fused_attention_bwd" + ) + q_format, kv_format = _q_kv_formats(qkv_layout) + dq_format, dkv_format = _q_kv_formats(dqkv_layout) + batch = cu_seqlens_q.numel() - 1 + is_ragged_q = q_format == "thd" + is_ragged_kv = kv_format == "thd" + use_ragged_stats = is_ragged_q and cudnn.backend_version() >= 90600 + use_token_buckets = ( + cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ) + use_legacy_offsets = is_ragged_q or is_ragged_kv + graph_batch = ( + _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch + ) + if not use_token_buckets: + use_ragged_stats = False + graph_seqlen_q = ( + _max_ragged_tokens(q.shape[0]) + if is_ragged_q and use_token_buckets + else max_seqlen_q + ) + graph_seqlen_kv = ( + _max_ragged_tokens(k.shape[0]) + if is_ragged_kv and use_token_buckets + else max_seqlen_kv + ) + + tensors: Dict[str, Any] = { + "_legacy_offsets": use_legacy_offsets, + "_graph_batch": graph_batch, + } + offset_q = offset_o = offset_k = offset_v = offset_stats = None + if is_ragged_q: + offset_q, q_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_q", + length=graph_batch + 1, + data_type=torch.int64, + ) + offset_o, o_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_o", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors.update(offset_q=offset_q, offset_o=offset_o) + else: + q_mult = o_mult = 1 + if is_ragged_kv: + offset_k, k_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1, + name="offset_k", + length=graph_batch + 1, + data_type=torch.int64, + ) + offset_v, v_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1, + name="offset_v", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors.update(offset_k=offset_k, offset_v=offset_v) + else: + k_mult = v_mult = 1 + + q_t = _make_bhsd_graph_tensor( + graph, + q, + q_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, + ragged_offset=offset_q, + ragged_offset_multiplier=q_mult, + name="Q", + ) + k_t = _make_bhsd_graph_tensor( + graph, + k, + kv_format, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, + ragged_offset=offset_k, + ragged_offset_multiplier=k_mult, + name="K", + ) + v_t = _make_bhsd_graph_tensor( + graph, + v, + kv_format, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, + ragged_offset=offset_v, + ragged_offset_multiplier=v_mult, + name="V", + ) + o_t = _make_bhsd_graph_tensor( + graph, + o, + o_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, + ragged_offset=offset_o, + ragged_offset_multiplier=o_mult, + name="O", + ) + do_t = _make_bhsd_graph_tensor( + graph, + d_o, + do_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, + ragged_offset=offset_o, + ragged_offset_multiplier=1, + name="dO", + ) + tensors.update(Q=q_t, K=k_t, V=v_t, O=o_t, dO=do_t) + + stats_physical, stats_dim, stats_stride = _stats_layout( + batch=graph_batch, + heads=q_t.get_dim()[1], + max_seqlen_q=graph_seqlen_q, + total_tokens_q=q.shape[0] if is_ragged_q else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + del stats_physical + if use_ragged_stats: + offset_stats, stats_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_stats", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors["offset_stats"] = offset_stats + else: + stats_mult = 1 + stats_t = graph.tensor( + name="Stats", + dim=stats_dim, + stride=stats_stride, + data_type=cudnn.data_type.FLOAT, + ragged_offset=offset_stats, + ragged_offset_multiplier=stats_mult, + ) + tensors["Stats"] = stats_t + + options = _mask_options( + cudnn, + attn_mask_type, + window_size, + bottom_right_diagonal, + max_seqlen_q, + max_seqlen_kv, + ) + is_padding = options.pop("is_padding") + options.update( + attn_scale=float(attn_scale), + use_padding_mask=is_padding, + use_alibi_mask=attn_bias_type == "alibi", + use_deterministic_algorithm=deterministic, + ) + if use_ragged_stats: + options["max_total_seq_len_q"] = graph_seqlen_q + if is_ragged_kv and cudnn.backend_version() >= 90600: + options["max_total_seq_len_kv"] = graph_seqlen_kv + + if attn_bias_type == "post_scale_bias": + bias_t = graph.tensor_like(attn_bias, name="Bias") + tensors["Bias"] = bias_t + options["bias"] = bias_t + if d_bias is not None: + d_bias_t = graph.tensor_like(d_bias, name="dBias").set_output(True) + tensors["dBias"] = d_bias_t + options["dBias"] = d_bias_t + + if is_padding: + seq_q = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + seq_kv = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) + seq_q_t = graph.tensor_like(seq_q, name="seq_len_q") + seq_kv_t = graph.tensor_like(seq_kv, name="seq_len_kv") + tensors.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) + options.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) + + if dropout != 0.0: + seed_t = graph.tensor( + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + offset_t = graph.tensor( + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + tensors.update(dropout_seed=seed_t, dropout_offset=offset_t) + options["dropout"] = (float(dropout), seed_t, offset_t) + + if softmax_type != "vanilla": + if softmax_offset is None or d_softmax_offset is None: + raise ValueError(f"softmax_type={softmax_type!r} requires sink tensors.") + sink_t = graph.tensor_like(softmax_offset, name="softmax_offset") + dsink_t = graph.tensor_like( + d_softmax_offset, name="d_softmax_offset" + ).set_output(True) + tensors.update(softmax_offset=sink_t, d_softmax_offset=dsink_t) + options.update(sink_token=sink_t, dSink_token=dsink_t) + + dq_t, dk_t, dv_t = graph.sdpa_backward( + name="te_sdpa_backward", + q=q_t, + k=k_t, + v=v_t, + o=o_t, + dO=do_t, + stats=stats_t, + **options, + ) + dq_dim, dq_stride = _logical_bhsd_desc( + d_q, dq_format, batch=graph_batch, max_seqlen=graph_seqlen_q + ) + dk_dim, dk_stride = _logical_bhsd_desc( + d_k, dkv_format, batch=graph_batch, max_seqlen=graph_seqlen_kv + ) + dv_dim, dv_stride = _logical_bhsd_desc( + d_v, dkv_format, batch=graph_batch, max_seqlen=graph_seqlen_kv + ) + dq_t.set_output(True).set_dim(dq_dim).set_stride(dq_stride) + dk_t.set_output(True).set_dim(dk_dim).set_stride(dk_stride) + dv_t.set_output(True).set_dim(dv_dim).set_stride(dv_stride) + if is_ragged_q: + dq_t.set_ragged_offset(offset_q).set_ragged_offset_multiplier(1) + if is_ragged_kv: + dk_t.set_ragged_offset(offset_k).set_ragged_offset_multiplier(1) + dv_t.set_ragged_offset(offset_v).set_ragged_offset_multiplier(1) + tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) + + return GraphEntry( + graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) + ) + + +def _build_fp8_bwd_graph( + *, + max_seqlen_q, + max_seqlen_kv, + q, + k, + v, + o, + d_o, + d_o_f16, + stats, + d_q, + d_k, + d_v, + s_quantizer, + dp_quantizer, + dqkv_quantizer, + qkv_layout, + o_format, + do_format, + dqkv_layout, + qkv_scale_inv_format, + do_scale_inv_format, + attn_scale, + dropout, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + deterministic, + softmax_offset, + d_softmax_offset, + cu_seqlens_q, + cu_seqlens_kv, +): + cudnn = import_cudnn_frontend() + q_format, kv_format = _q_kv_formats(qkv_layout) + dq_format, dkv_format = _q_kv_formats(dqkv_layout) + batch = cu_seqlens_q.numel() - 1 + heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] + kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] + d_qk, d_value = q.shape[-1], v.shape[-1] + q_data = _quantized_data(q) + k_data = _quantized_data(k) + v_data = _quantized_data(v) + o_data = _quantized_data(o) if _is_float8_tensor(o) else o + do_data = _quantized_data(d_o) + dq_data = _quantized_data(d_q) if _is_float8_tensor(d_q) else d_q + dk_data = _quantized_data(d_k) if _is_float8_tensor(d_k) else d_k + dv_data = _quantized_data(d_v) if _is_float8_tensor(d_v) else d_v + graph = make_graph(_fp8_cudnn_dtype(q), q.device, name="te_fp8_sdpa_bwd") + tensors: Dict[str, Any] = {} + + q_t = _make_bhsd_graph_tensor( + graph, + q_data, + q_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=_fp8_cudnn_dtype(q), + name="Q", + ) + k_t = _make_bhsd_graph_tensor( + graph, + k_data, + kv_format, + batch=batch, + max_seqlen=max_seqlen_kv, + data_type=_fp8_cudnn_dtype(k), + name="K", + ) + v_t = _make_bhsd_graph_tensor( + graph, + v_data, + kv_format, + batch=batch, + max_seqlen=max_seqlen_kv, + data_type=_fp8_cudnn_dtype(v), + name="V", + ) + o_dtype = _fp8_cudnn_dtype(o) if _is_float8_tensor(o) else o.dtype + o_t = _make_bhsd_graph_tensor( + graph, + o_data, + o_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=o_dtype, + name="O", + ) + do_t = _make_bhsd_graph_tensor( + graph, + do_data, + do_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=_fp8_cudnn_dtype(d_o), + name="dO", + ) + stats_t = graph.tensor_like(stats, name="Stats") + tensors.update(Q=q_t, K=k_t, V=v_t, O=o_t, dO=do_t, Stats=stats_t) + + options = _mask_options( + cudnn, + attn_mask_type, + window_size, + bottom_right_diagonal, + max_seqlen_q, + max_seqlen_kv, + ) + is_padding = options.pop("is_padding") + if "diagonal_band_left_bound" in options: + options["left_bound"] = options.pop("diagonal_band_left_bound") + if "diagonal_band_right_bound" in options: + options["right_bound"] = options.pop("diagonal_band_right_bound") + options.update( + attn_scale=float(attn_scale), + use_padding_mask=is_padding, + use_deterministic_algorithm=deterministic, + ) + if is_padding: + seq_q = _sequence_lengths(cu_seqlens_q) + seq_kv = _sequence_lengths(cu_seqlens_kv) + seq_q_t = graph.tensor_like(seq_q, name="seq_len_q") + seq_kv_t = graph.tensor_like(seq_kv, name="seq_len_kv") + tensors.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) + options.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) + if dropout != 0.0: + seed_t = graph.tensor( + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + offset_t = graph.tensor( + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + data_type=cudnn.data_type.INT64, + ) + tensors.update(dropout_seed=seed_t, dropout_offset=offset_t) + options["dropout"] = (float(dropout), seed_t, offset_t) + if softmax_type != "vanilla": + sink_t = graph.tensor_like(softmax_offset, name="softmax_offset") + dsink_t = graph.tensor_like( + d_softmax_offset, name="d_softmax_offset" + ).set_output(True) + tensors.update(softmax_offset=sink_t, d_softmax_offset=dsink_t) + options.update(sink_token=sink_t, dSink_token=dsink_t) + + if _is_mxfp8_tensor(q): + if d_o_f16 is None: + raise ValueError( + "MXFP8 attention backward requires the high-precision dO tensor." + ) + q_col = _quantized_data(q, columnwise=True) + k_col = _quantized_data(k, columnwise=True) + do_col = _quantized_data(d_o, columnwise=True) + q_col_t = _make_bhsd_graph_tensor( + graph, + q_col, + q_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=_fp8_cudnn_dtype(q), + name="Q_T", + ) + k_col_t = _make_bhsd_graph_tensor( + graph, + k_col, + kv_format, + batch=batch, + max_seqlen=max_seqlen_kv, + data_type=_fp8_cudnn_dtype(k), + name="K_T", + ) + do_col_t = _make_bhsd_graph_tensor( + graph, + do_col, + do_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=_fp8_cudnn_dtype(d_o), + name="dO_T", + ) + do_f16_t = _make_bhsd_graph_tensor( + graph, + d_o_f16, + do_format, + batch=batch, + max_seqlen=max_seqlen_q, + data_type=d_o_f16.dtype, + name="dO_f16", + ) + scale_format_q = qkv_scale_inv_format or q_format + scale_format_kv = qkv_scale_inv_format or kv_format + scale_format_do = do_scale_inv_format or do_format + padded = _mxfp8_padded_sizes(max_seqlen_q, max_seqlen_kv, d_qk, d_value) + + def mx_scale(name, h, s, d, fmt): + tensor = _make_mxfp8_scale_tensor( + graph, + cudnn, + name=name, + batch=batch, + heads=h, + seqlen=padded[s], + dim=padded[d], + tensor_format=fmt, + ) + tensors[name] = tensor + return tensor + + descale_q = mx_scale( + "descale_q", heads, "s_q_padded", "d_qk_scale_padded", scale_format_q + ) + descale_q_t = mx_scale( + "descale_q_t", heads, "s_q_scale_padded", "d_qk_padded", scale_format_q + ) + descale_k = mx_scale( + "descale_k", kv_heads, "s_kv_padded", "d_qk_scale_padded", scale_format_kv + ) + descale_k_t = mx_scale( + "descale_k_t", kv_heads, "s_kv_scale_padded", "d_qk_padded", scale_format_kv + ) + descale_v = mx_scale( + "descale_v", kv_heads, "s_kv_padded", "d_v_scale_padded", scale_format_kv + ) + descale_do = mx_scale( + "descale_do", heads, "s_q_padded", "d_v_scale_padded", scale_format_do + ) + descale_do_t = mx_scale( + "descale_do_t", heads, "s_q_scale_padded", "d_v_padded", scale_format_do + ) + outputs = graph.sdpa_mxfp8_backward( + q_t, + q_col_t, + k_t, + k_col_t, + v_t, + o_t, + do_f16_t, + do_t, + do_col_t, + stats_t, + descale_q, + descale_q_t, + descale_k, + descale_k_t, + descale_v, + descale_do, + descale_do_t, + name="te_sdpa_mxfp8_backward", + **options, + ) + dq_t, dk_t, dv_t, *amax_outputs = outputs + for amax_t in amax_outputs: + amax_t.set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) + tensors.update(Q_T=q_col_t, K_T=k_col_t, dO_T=do_col_t, dO_f16=do_f16_t) + else: + scalar_names = ( + "descale_q", + "descale_k", + "descale_v", + "descale_o", + "descale_do", + ) + scalars = { + name: _scalar_graph_tensor(graph, cudnn, name) for name in scalar_names + } + tensors.update(scalars) + delayed = isinstance(dqkv_quantizer, Float8Quantizer) + if isinstance(s_quantizer, Float8Quantizer): + for name in ("descale_s", "scale_s"): + tensors[name] = _scalar_graph_tensor(graph, cudnn, name) + else: + for name in ("descale_s", "scale_s"): + tensors[f"constant_{name}"] = _constant_graph_tensor(graph, cudnn, name) + if isinstance(dp_quantizer, Float8Quantizer): + for name in ("descale_dp", "scale_dp"): + tensors[name] = _scalar_graph_tensor(graph, cudnn, name) + else: + for name in ("descale_dp", "scale_dp"): + tensors[f"constant_{name}"] = _constant_graph_tensor(graph, cudnn, name) + for name in ("scale_dq", "scale_dk", "scale_dv"): + key = name if delayed else f"constant_{name}" + tensors[key] = _scalar_graph_tensor(graph, cudnn, name) + descale_o_arg = tensors["descale_o"] + descale_s_arg = ( + tensors["descale_s"] + if "descale_s" in tensors + else tensors["constant_descale_s"] + ) + descale_dp_arg = ( + tensors["descale_dp"] + if "descale_dp" in tensors + else tensors["constant_descale_dp"] + ) + scale_s_arg = ( + tensors["scale_s"] if "scale_s" in tensors else tensors["constant_scale_s"] + ) + scale_dp_arg = ( + tensors["scale_dp"] + if "scale_dp" in tensors + else tensors["constant_scale_dp"] + ) + scale_dq_arg = ( + tensors["scale_dq"] + if "scale_dq" in tensors + else tensors["constant_scale_dq"] + ) + scale_dk_arg = ( + tensors["scale_dk"] + if "scale_dk" in tensors + else tensors["constant_scale_dk"] + ) + scale_dv_arg = ( + tensors["scale_dv"] + if "scale_dv" in tensors + else tensors["constant_scale_dv"] + ) + outputs = graph.sdpa_fp8_backward( + q_t, + k_t, + v_t, + o_t, + do_t, + stats_t, + tensors["descale_q"], + tensors["descale_k"], + tensors["descale_v"], + descale_o_arg, + tensors["descale_do"], + descale_s_arg, + descale_dp_arg, + scale_s_arg, + scale_dq_arg, + scale_dk_arg, + scale_dv_arg, + scale_dp_arg, + name="te_sdpa_fp8_backward", + **options, + ) + dq_t, dk_t, dv_t, amax_dq_t, amax_dk_t, amax_dv_t, amax_dp_t = outputs + for name, amax_t in ( + ("amax_dq", amax_dq_t), + ("amax_dk", amax_dk_t), + ("amax_dv", amax_dv_t), + ("amax_dp", amax_dp_t), + ): + amax_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) + tensors[name] = amax_t + + dq_t.set_output(True).set_data_type( + _fp8_cudnn_dtype(d_q) if _is_float8_tensor(d_q) else d_q.dtype + ).set_dim((batch, heads, max_seqlen_q, d_qk)).set_stride( + _format_stride(batch, heads, max_seqlen_q, d_qk, dq_format) + ) + dk_t.set_output(True).set_data_type( + _fp8_cudnn_dtype(d_k) if _is_float8_tensor(d_k) else d_k.dtype + ).set_dim((batch, kv_heads, max_seqlen_kv, d_qk)).set_stride( + _format_stride(batch, kv_heads, max_seqlen_kv, d_qk, dkv_format) + ) + dv_t.set_output(True).set_data_type( + _fp8_cudnn_dtype(d_v) if _is_float8_tensor(d_v) else d_v.dtype + ).set_dim((batch, kv_heads, max_seqlen_kv, d_value)).set_stride( + _format_stride(batch, kv_heads, max_seqlen_kv, d_value, dkv_format) + ) + tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) + return GraphEntry( + graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) + ) + + +def _fp8_backward( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + o, + d_o, + fake_dtype, + aux_ctx_tensors, + s_quantizer, + dp_quantizer, + dqkv_quantizer, + attn_scale, + dropout, + fast_zero_fill, + qkv_layout, + o_format, + do_format, + dqkv_layout, + qkv_scale_inv_format, + do_scale_inv_format, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + deterministic, +): + if _is_mxfp8_tensor(q) and "padding" in attn_mask_type: + q_hp, k_hp, v_hp = (tensor.dequantize(dtype=fake_dtype) for tensor in (q, k, v)) + d_o_f16 = aux_ctx_tensors[-1] + aux_count = 3 if softmax_type != "vanilla" else 2 + return fused_attn_bwd( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q_hp, + k_hp, + v_hp, + o, + d_o_f16, + fake_dtype, + aux_ctx_tensors[:aux_count], + FusedAttnBackend.F16_arbitrary_seqlen, + s_quantizer=None, + dp_quantizer=None, + dqkv_quantizer=None, + attn_scale=attn_scale, + dropout=dropout, + fast_zero_fill=fast_zero_fill, + qkv_layout=qkv_layout, + o_format=o_format, + do_format=do_format, + dqkv_layout=dqkv_layout, + attn_bias_type="no_bias", + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + deterministic=deterministic, + ) + q_format, kv_format = _q_kv_formats(qkv_layout) + batch = cu_seqlens_q.numel() - 1 + heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] + kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] + output_dtype = ( + torch.uint8 if isinstance(dqkv_quantizer, Float8Quantizer) else fake_dtype + ) + grad_data = _allocate_attention_grad_data( + batch=batch, + heads=heads, + kv_heads=kv_heads, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + head_dim_qk=q.shape[-1], + head_dim_v=v.shape[-1], + dqkv_layout=dqkv_layout, + dtype=output_dtype, + device=q.device, + zero=fast_zero_fill, + ) + if isinstance(dqkv_quantizer, Float8Quantizer): + d_q, d_k, d_v = _wrap_float8_grad_outputs(dqkv_quantizer, grad_data, fake_dtype) + else: + d_q, d_k, d_v = grad_data + stats, rng_state = aux_ctx_tensors[:2] + softmax_offset = aux_ctx_tensors[2] if softmax_type != "vanilla" else None + d_o_f16 = aux_ctx_tensors[-1] if _is_mxfp8_tensor(q) else None + d_softmax_offset = ( + torch.empty_like(softmax_offset) if softmax_offset is not None else None + ) + hidden_amax = [ + torch.zeros(1, dtype=torch.float32, device=q.device) for _ in range(4) + ] + + key = ( + "fp8_bwd", + max_seqlen_q, + max_seqlen_kv, + *( + _tensor_metadata(_quantized_data(x) if _is_float8_tensor(x) else x) + for x in (q, k, v, o, d_o) + ), + type(q).__name__, + type(dqkv_quantizer).__name__, + qkv_layout, + o_format, + do_format, + dqkv_layout, + qkv_scale_inv_format, + do_scale_inv_format, + float(attn_scale), + float(dropout), + attn_mask_type, + softmax_type, + tuple(window_size), + bottom_right_diagonal, + deterministic, + ) + entry = get_graph_entry(key) + if entry is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "cuDNN FP8 attention graph must be built before CUDA graph capture." + ) + entry = _build_fp8_bwd_graph( + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + q=q, + k=k, + v=v, + o=o, + d_o=d_o, + d_o_f16=d_o_f16, + stats=stats, + d_q=d_q, + d_k=d_k, + d_v=d_v, + s_quantizer=s_quantizer, + dp_quantizer=dp_quantizer, + dqkv_quantizer=dqkv_quantizer, + qkv_layout=qkv_layout, + o_format=o_format, + do_format=do_format, + dqkv_layout=dqkv_layout, + qkv_scale_inv_format=qkv_scale_inv_format, + do_scale_inv_format=do_scale_inv_format, + attn_scale=attn_scale, + dropout=dropout, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + deterministic=deterministic, + softmax_offset=softmax_offset, + d_softmax_offset=d_softmax_offset, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + ) + put_graph_entry(key, entry) + + t = entry.tensors + variant_pack = { + t["Q"]: _quantized_data(q), + t["K"]: _quantized_data(k), + t["V"]: _quantized_data(v), + t["O"]: _quantized_data(o) if _is_float8_tensor(o) else o, + t["dO"]: _quantized_data(d_o), + t["Stats"]: stats, + t["dQ"]: _quantized_data(d_q) if _is_float8_tensor(d_q) else d_q, + t["dK"]: _quantized_data(d_k) if _is_float8_tensor(d_k) else d_k, + t["dV"]: _quantized_data(d_v) if _is_float8_tensor(d_v) else d_v, + } + if _is_mxfp8_tensor(q): + variant_pack.update( + { + t["Q_T"]: _quantized_data(q, columnwise=True), + t["K_T"]: _quantized_data(k, columnwise=True), + t["dO_T"]: _quantized_data(d_o, columnwise=True), + t["dO_f16"]: d_o_f16, + t["descale_q"]: _quantized_scale_inv(q), + t["descale_q_t"]: _quantized_scale_inv(q, columnwise=True), + t["descale_k"]: _quantized_scale_inv(k), + t["descale_k_t"]: _quantized_scale_inv(k, columnwise=True), + t["descale_v"]: _quantized_scale_inv(v), + t["descale_do"]: _quantized_scale_inv(d_o), + t["descale_do_t"]: _quantized_scale_inv(d_o, columnwise=True), + } + ) + else: + variant_pack.update( + { + t["descale_q"]: _quantized_scale_inv(q), + t["descale_k"]: _quantized_scale_inv(k), + t["descale_v"]: _quantized_scale_inv(v), + t["descale_o"]: ( + _quantized_scale_inv(o) + if _is_float8_tensor(o) + else torch.ones(1, dtype=torch.float32, device=q.device) + ), + t["descale_do"]: _quantized_scale_inv(d_o), + } + ) + if "descale_s" in t: + variant_pack[t["descale_s"]] = torch.reciprocal(s_quantizer.scale) + variant_pack[t["scale_s"]] = s_quantizer.scale + if "descale_dp" in t: + variant_pack[t["descale_dp"]] = torch.reciprocal(dp_quantizer.scale) + variant_pack[t["scale_dp"]] = dp_quantizer.scale + for name in ("scale_dq", "scale_dk", "scale_dv"): + if name in t: + variant_pack[t[name]] = dqkv_quantizer.scale + one = torch.ones(1, dtype=torch.float32, device=q.device) + for name, graph_tensor in t.items(): + if name.startswith("constant_"): + variant_pack[graph_tensor] = one + amax_values = ( + ( + dqkv_quantizer.amax, + dqkv_quantizer.amax, + dqkv_quantizer.amax, + dp_quantizer.amax, + ) + if isinstance(dqkv_quantizer, Float8Quantizer) + else hidden_amax + ) + for name, value in zip( + ("amax_dq", "amax_dk", "amax_dv", "amax_dp"), amax_values + ): + variant_pack[t[name]] = value + if "seq_len_q" in t: + variant_pack[t["seq_len_q"]] = _sequence_lengths(cu_seqlens_q) + variant_pack[t["seq_len_kv"]] = _sequence_lengths(cu_seqlens_kv) + if "dropout_seed" in t: + variant_pack[t["dropout_seed"]] = rng_state[:1] + variant_pack[t["dropout_offset"]] = rng_state[1:] + if "softmax_offset" in t: + variant_pack[t["softmax_offset"]] = softmax_offset + variant_pack[t["d_softmax_offset"]] = d_softmax_offset + entry.execute(variant_pack, q.device) + return d_q, d_k, d_v, None, d_softmax_offset + + +def fused_attn_bwd( + max_seqlen_q: int, + max_seqlen_kv: int, + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + o: torch.Tensor, + d_o: torch.Tensor, + fake_dtype: torch.dtype, + aux_ctx_tensors: List[torch.Tensor], + fused_attention_backend: FusedAttnBackend, + cu_seqlens_q_padded: torch.Tensor = None, + cu_seqlens_kv_padded: torch.Tensor = None, + s_quantizer=None, + dp_quantizer=None, + dqkv_quantizer=None, + attn_scale: Optional[float] = None, + dropout: float = 0.0, + fast_zero_fill: bool = True, + qkv_layout: str = "sbh3d", + o_format: str = "sbhd", + do_format: str = "sbhd", + dqkv_layout: str = "sbh3d", + qkv_scale_inv_format: str = None, + do_scale_inv_format: str = None, + attn_bias_type: str = "no_bias", + attn_mask_type: str = "padding", + softmax_type: str = "vanilla", + window_size: Tuple[int, int] = (-1, -1), + bottom_right_diagonal: bool = None, + deterministic: bool = False, + cuda_graph: bool = False, +) -> Tuple[Union[torch.Tensor, None], ...]: + """Execute fused-attention backward through the Python graph API.""" + + del cuda_graph + backend = FusedAttnBackend.cast(fused_attention_backend) + if not aux_ctx_tensors: + raise ValueError("Fused-attention backward requires forward auxiliary tensors.") + if attn_scale is None: + attn_scale = 1.0 / math.sqrt(q.size(-1)) + if bottom_right_diagonal is None: + bottom_right_diagonal = attn_mask_type in ( + "causal_bottom_right", + "padding_causal_bottom_right", + ) + if backend == FusedAttnBackend.FP8: + if attn_bias_type != "no_bias": + raise ValueError( + "FP8 fused attention backward does not support attention bias." + ) + return _fp8_backward( + max_seqlen_q, + max_seqlen_kv, + cu_seqlens_q, + cu_seqlens_kv, + q, + k, + v, + o, + d_o, + fake_dtype, + aux_ctx_tensors, + s_quantizer, + dp_quantizer, + dqkv_quantizer, + attn_scale, + dropout, + fast_zero_fill, + qkv_layout, + o_format, + do_format, + dqkv_layout, + qkv_scale_inv_format, + do_scale_inv_format, + attn_mask_type, + softmax_type, + window_size, + bottom_right_diagonal, + deterministic, + ) + if backend != FusedAttnBackend.F16_arbitrary_seqlen: + raise ValueError( + "No cuDNN fused-attention backend supports this backward configuration." + ) + + stats = aux_ctx_tensors[0] + rng_state = aux_ctx_tensors[1] + aux_index = 2 + attn_bias = None + if attn_bias_type not in ("no_bias", "alibi"): + attn_bias = aux_ctx_tensors[aux_index] + aux_index += 1 + softmax_offset = None + if softmax_type != "vanilla": + softmax_offset = aux_ctx_tensors[aux_index] + + q_format, kv_format = _q_kv_formats(qkv_layout) + if q_format == "thd" or kv_format == "thd": + d_q, d_k, d_v = _allocate_grad_views((q, k, v), fast_zero_fill=fast_zero_fill) + else: + batch = cu_seqlens_q.numel() - 1 + heads = q.shape[1] if q_format == "bhsd" else q.shape[-2] + kv_heads = k.shape[1] if kv_format == "bhsd" else k.shape[-2] + d_q, d_k, d_v = _allocate_attention_grad_data( + batch=batch, + heads=heads, + kv_heads=kv_heads, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + head_dim_qk=q.shape[-1], + head_dim_v=v.shape[-1], + dqkv_layout=dqkv_layout, + dtype=q.dtype, + device=q.device, + zero=fast_zero_fill, + ) + d_bias = None + if attn_bias_type == "post_scale_bias": + # cuDNN does not support the [1,1,1,S] reduction form. + if not tuple(attn_bias.shape[:3]) == (1, 1, 1): + d_bias = torch.empty_like(attn_bias) + d_softmax_offset = ( + torch.empty_like(softmax_offset) if softmax_type != "vanilla" else None + ) + cu_seqlens_q_padded = ( + cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded + ) + cu_seqlens_kv_padded = ( + cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded + ) + + key = ( + "f16_bwd", + max_seqlen_q, + max_seqlen_kv, + _tensor_metadata(q), + _tensor_metadata(k), + _tensor_metadata(v), + _tensor_metadata(o), + _tensor_metadata(d_o), + _tensor_metadata(stats), + _tensor_metadata(d_q), + _tensor_metadata(d_k), + _tensor_metadata(d_v), + _tensor_metadata(attn_bias), + qkv_layout, + o_format, + do_format, + dqkv_layout, + float(attn_scale), + float(dropout), + attn_bias_type, + attn_mask_type, + softmax_type, + tuple(window_size), + bottom_right_diagonal, + deterministic, + ) + entry = get_graph_entry(key) + if entry is None: + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "cuDNN attention graph must be built before CUDA graph capture." + ) + entry = _build_f16_bwd_graph( + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + q=q, + k=k, + v=v, + o=o, + d_o=d_o, + stats=stats, + d_q=d_q, + d_k=d_k, + d_v=d_v, + attn_bias=attn_bias, + d_bias=d_bias, + softmax_offset=softmax_offset, + d_softmax_offset=d_softmax_offset, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, + attn_scale=attn_scale, + dropout=dropout, + qkv_layout=qkv_layout, + o_format=o_format, + do_format=do_format, + dqkv_layout=dqkv_layout, + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + window_size=window_size, + bottom_right_diagonal=bottom_right_diagonal, + deterministic=deterministic, + ) + put_graph_entry(key, entry) + + tensors = entry.tensors + variant_pack: Dict[Any, Any] = { + tensors["Q"]: q, + tensors["K"]: k, + tensors["V"]: v, + tensors["O"]: o, + tensors["dO"]: d_o, + tensors["Stats"]: stats, + tensors["dQ"]: d_q, + tensors["dK"]: d_k, + tensors["dV"]: d_v, + } + if "Bias" in tensors: + variant_pack[tensors["Bias"]] = attn_bias + if "dBias" in tensors: + variant_pack[tensors["dBias"]] = d_bias + graph_batch = tensors["_graph_batch"] + if "seq_len_q" in tensors: + variant_pack[tensors["seq_len_q"]] = _padded_sequence_lengths( + cu_seqlens_q, graph_batch + ) + variant_pack[tensors["seq_len_kv"]] = _padded_sequence_lengths( + cu_seqlens_kv, graph_batch + ) + if "offset_q" in tensors: + variant_pack[tensors["offset_q"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, q.stride(0) + ) + variant_pack[tensors["offset_o"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, o.stride(0) + ) + if "offset_k" in tensors: + variant_pack[tensors["offset_k"]] = _element_ragged_offsets( + cu_seqlens_kv_padded, graph_batch, k.stride(0) + ) + variant_pack[tensors["offset_v"]] = _element_ragged_offsets( + cu_seqlens_kv_padded, graph_batch, v.stride(0) + ) + if "offset_stats" in tensors: + stats_multiplier = stats.stride(0) + variant_pack[tensors["offset_stats"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, stats_multiplier + ) + if "dropout_seed" in tensors: + variant_pack[tensors["dropout_seed"]] = rng_state[:1] + variant_pack[tensors["dropout_offset"]] = rng_state[1:] + if "softmax_offset" in tensors: + variant_pack[tensors["softmax_offset"]] = softmax_offset + variant_pack[tensors["d_softmax_offset"]] = d_softmax_offset + entry.execute(variant_pack, q.device) + return d_q, d_k, d_v, d_bias, d_softmax_offset diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index b9593b42d9b..cc19ffd02a9 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -5,26 +5,26 @@ """cuDNN-backed Flex Attention helpers.""" from dataclasses import dataclass -import importlib import inspect from typing import Any, Callable, Dict, Optional, Tuple import torch -_cudnn_score_mod_handles: Dict[torch.device, Any] = {} +from ._cudnn_graph import ( + current_stream_handle, + finalize_graph, + import_cudnn_frontend, + make_graph, + torch_to_cudnn_dtype, +) + _cudnn_score_mod_graph_cache: Dict[Tuple[Any, ...], Any] = {} _SCORE_MOD_UNCACHEABLE = object() def _import_cudnn_frontend(): - """Import the cuDNN frontend Python package.""" - try: - return importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "cuDNN frontend Python package not found. " - "Install it with: pip install nvidia-cudnn-frontend" - ) from exc + """Compatibility wrapper around the shared attention graph runtime.""" + return import_cudnn_frontend() def _bhsd_dim_stride( @@ -41,7 +41,9 @@ def _bhsd_dim_stride( (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), ) - raise ValueError(f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}.") + raise ValueError( + f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}." + ) def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): @@ -164,10 +166,14 @@ def _score_mod_tensor_dict_metadata( """Describe score_mod tensor parameters without including their values.""" if tensors is None: return () - return tuple((name, _score_mod_tensor_metadata(tensor)) for name, tensor in tensors.items()) + return tuple( + (name, _score_mod_tensor_metadata(tensor)) for name, tensor in tensors.items() + ) -def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> Tuple[Any, ...]: +def _score_mod_bhsd_tensor_metadata( + tensor: torch.Tensor, tensor_format: str +) -> Tuple[Any, ...]: """Describe an SBHD/BSHD runtime tensor as a cuDNN BHSD graph tensor.""" dim, stride = _bhsd_dim_stride(tensor, tensor_format) return (dim, stride, tensor.dtype, _score_mod_device_key(tensor.device)) @@ -193,41 +199,18 @@ def _wrapped_score_mod(sdpa_graph, score_tensor): def _get_cudnn_current_stream_handle(cudnn, device: torch.device): - """Return a cuDNN handle for device, bound to PyTorch's current stream.""" - if device.type != "cuda": - raise ValueError(f"Flex Attention only supports CUDA tensors, got device {device}.") - if device.index is None: - device = torch.device("cuda", torch.cuda.current_device()) - - handle = _cudnn_score_mod_handles.get(device) - with torch.cuda.device(device): - if handle is None: - handle = cudnn.create_handle() - _cudnn_score_mod_handles[device] = handle - - stream = torch.cuda.current_stream(device).cuda_stream - cudnn.set_stream(handle=handle, stream=stream) - return handle + """Compatibility wrapper around the shared current-stream handle.""" + del cudnn + return current_stream_handle(device) def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" - cudnn = _import_cudnn_frontend() - - if dtype == torch.float16: - io_data_type = cudnn.data_type.HALF - elif dtype == torch.bfloat16: - io_data_type = cudnn.data_type.BFLOAT16 - else: - raise ValueError(f"Flex Attention only supports FP16/BF16 tensors, got {dtype}.") - - graph = cudnn.pygraph( - io_data_type=io_data_type, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - handle=_get_cudnn_current_stream_handle(cudnn, device), - ) - return graph + if dtype not in (torch.float16, torch.bfloat16): + raise ValueError( + f"Flex Attention only supports FP16/BF16 tensors, got {dtype}." + ) + return make_graph(torch_to_cudnn_dtype(dtype), device, name="te_flex_attention") @dataclass @@ -264,18 +247,8 @@ class _CudnnScoreModBwdGraphEntry: def _finalize_cudnn_graph(graph) -> int: - """Build a cuDNN frontend Python graph and return its workspace size.""" - cudnn = _import_cudnn_frontend() - - graph.validate() - graph.build_operation_graph() - try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN Flex Attention SDPA graph is not supported: {exc}") from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) - return max(graph.get_workspace_size(), 1) + """Compatibility wrapper around shared graph finalization.""" + return finalize_graph(graph) def _execute_cudnn_graph( @@ -357,7 +330,10 @@ def _cudnn_score_mod_bwd_cache_key( """Pre-build cache key for score_mod bprop execution plans.""" score_mod_key = _score_mod_callback_cache_key(score_mod) score_mod_bprop_key = _score_mod_callback_cache_key(score_mod_bprop) - if score_mod_key is _SCORE_MOD_UNCACHEABLE or score_mod_bprop_key is _SCORE_MOD_UNCACHEABLE: + if ( + score_mod_key is _SCORE_MOD_UNCACHEABLE + or score_mod_bprop_key is _SCORE_MOD_UNCACHEABLE + ): return None return ( "bwd", @@ -505,7 +481,9 @@ def _build_cudnn_score_mod_bwd_graph( else {} ) wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) - wrapped_score_mod_bprop = _wrap_score_mod(score_mod_bprop, score_mod_bprop_graph_tensors) + wrapped_score_mod_bprop = _wrap_score_mod( + score_mod_bprop, score_mod_bprop_graph_tensors + ) dq_layer = torch.empty_like(query_layer) dk_layer = torch.empty_like(key_layer) @@ -616,7 +594,9 @@ def forward( score_mod_tensors = dict(score_mod_tensors or {}) score_mod_bprop_tensors = dict(score_mod_bprop_tensors or {}) output_shape = (*query_layer.shape[:-1], value_layer.shape[-1]) - output_layer = torch.empty(output_shape, device=query_layer.device, dtype=query_layer.dtype) + output_layer = torch.empty( + output_shape, device=query_layer.device, dtype=query_layer.dtype + ) if is_training: stats = torch.empty( (*q_bhsd_dim[:-1], 1), diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index ba049c9aef9..385935626dc 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -22,11 +22,7 @@ import torch.nn.functional as F import transformer_engine_torch as tex import transformer_engine as te -from transformer_engine.pytorch.cpp_extensions.fused_attn import ( - QKVLayout, - AttnBiasType, - AttnMaskType, - SoftmaxType, +from transformer_engine.pytorch.attention.dot_product_attention.cudnn_attention import ( FusedAttnBackend, META_QKV, META_DQKV, @@ -35,6 +31,9 @@ META_S, META_DP, ) +from transformer_engine.pytorch.attention.dot_product_attention._cudnn_backend import ( + get_fused_attn_backend as get_cudnn_fused_attn_backend, +) from transformer_engine.pytorch.attention.inference import InferenceParams from transformer_engine.pytorch.cpu_offload import is_cpu_offload_enabled from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage @@ -365,7 +364,7 @@ def _get_fused_attn_backend( softmax_type, *args, ): - """Constant-foldable tex.get_fused_attn_backend: the result depends only on + """Constant-foldable Python backend selector: the result depends only on the attention config. Layout/bias/mask/softmax are taken as their string keys and resolved to the pybind enums here, so that every argument is a python literal or a python enum. @@ -377,14 +376,14 @@ def _get_fused_attn_backend( member comes out of the reconstruction corrupted (see the cast at the call site, which restores the enum).""" return int( - tex.get_fused_attn_backend( + get_cudnn_fused_attn_backend( is_training, q_type, kv_type, - QKVLayout[qkv_layout], - AttnBiasType[bias_type], - AttnMaskType[attn_mask_type], - SoftmaxType[softmax_type], + qkv_layout, + bias_type, + attn_mask_type, + softmax_type, *args, ) ) diff --git a/transformer_engine/pytorch/cpp_extensions/fused_attn.py b/transformer_engine/pytorch/cpp_extensions/fused_attn.py index 9a33df7634b..c509dadfd2e 100644 --- a/transformer_engine/pytorch/cpp_extensions/fused_attn.py +++ b/transformer_engine/pytorch/cpp_extensions/fused_attn.py @@ -2,30 +2,20 @@ # # See LICENSE for license information. -"""Python interface for fused attention extensions""" +"""Compatibility imports for the Python cuDNN attention implementation.""" -import math from enum import IntEnum -from typing import Tuple, List, Union, Optional + import torch -import transformer_engine_torch as tex -from transformer_engine_torch import ( - NVTE_QKV_Layout, - NVTE_QKV_Format, - NVTE_Bias_Type, - NVTE_Mask_Type, - NVTE_Softmax_Type, - NVTE_Fused_Attn_Backend, -) -from ..quantized_tensor import Quantizer -from ..constants import FP8BwdTensorIdx, FP8FwdTensorIdx, DType +from transformer_engine_torch import NVTE_QKV_Format + +from ..constants import DType, FP8BwdTensorIdx, FP8FwdTensorIdx -__all__ = [ - "fused_attn_fwd", - "fused_attn_bwd", -] +__all__ = ["fused_attn_fwd", "fused_attn_bwd"] +# Retained for rotary-position and custom-op compatibility. These operations +# still use TE's format enum even though attention execution itself does not. TORCH_DType = { DType.kFloat8E4M3: torch.uint8, DType.kFloat8E5M2: torch.uint8, @@ -47,122 +37,21 @@ "bhsd": NVTE_QKV_Format.NVTE_BHSD, } -QKVLayout = { - "sb3hd": NVTE_QKV_Layout.NVTE_SB3HD, - "sbh3d": NVTE_QKV_Layout.NVTE_SBH3D, - "sbhd_sb2hd": NVTE_QKV_Layout.NVTE_SBHD_SB2HD, - "sbhd_sbh2d": NVTE_QKV_Layout.NVTE_SBHD_SBH2D, - "sbhd_sbhd_sbhd": NVTE_QKV_Layout.NVTE_SBHD_SBHD_SBHD, - "bs3hd": NVTE_QKV_Layout.NVTE_BS3HD, - "bsh3d": NVTE_QKV_Layout.NVTE_BSH3D, - "bshd_bs2hd": NVTE_QKV_Layout.NVTE_BSHD_BS2HD, - "bshd_bsh2d": NVTE_QKV_Layout.NVTE_BSHD_BSH2D, - "bshd_bshd_bshd": NVTE_QKV_Layout.NVTE_BSHD_BSHD_BSHD, - "t3hd": NVTE_QKV_Layout.NVTE_T3HD, - "th3d": NVTE_QKV_Layout.NVTE_TH3D, - "thd_t2hd": NVTE_QKV_Layout.NVTE_THD_T2HD, - "thd_th2d": NVTE_QKV_Layout.NVTE_THD_TH2D, - "thd_thd_thd": NVTE_QKV_Layout.NVTE_THD_THD_THD, - "sbhd_bshd_bshd": NVTE_QKV_Layout.NVTE_SBHD_BSHD_BSHD, - "bshd_sbhd_sbhd": NVTE_QKV_Layout.NVTE_BSHD_SBHD_SBHD, - "thd_bshd_bshd": NVTE_QKV_Layout.NVTE_THD_BSHD_BSHD, - "thd_sbhd_sbhd": NVTE_QKV_Layout.NVTE_THD_SBHD_SBHD, - "paged_kv_bshd_bshd_bshd": NVTE_QKV_Layout.NVTE_Paged_KV_BSHD_BSHD_BSHD, - "paged_kv_bshd_sbhd_sbhd": NVTE_QKV_Layout.NVTE_Paged_KV_BSHD_SBHD_SBHD, - "paged_kv_sbhd_bshd_bshd": NVTE_QKV_Layout.NVTE_Paged_KV_SBHD_BSHD_BSHD, - "paged_kv_sbhd_sbhd_sbhd": NVTE_QKV_Layout.NVTE_Paged_KV_SBHD_SBHD_SBHD, - "paged_kv_thd_bshd_bshd": NVTE_QKV_Layout.NVTE_Paged_KV_THD_BSHD_BSHD, - "paged_kv_thd_sbhd_sbhd": NVTE_QKV_Layout.NVTE_Paged_KV_THD_SBHD_SBHD, - "bhsd_bhsd_bhsd": NVTE_QKV_Layout.NVTE_BHSD_BHSD_BHSD, -} - -AttnBiasType = { - "no_bias": NVTE_Bias_Type.NVTE_NO_BIAS, - "pre_scale_bias": NVTE_Bias_Type.NVTE_PRE_SCALE_BIAS, - "post_scale_bias": NVTE_Bias_Type.NVTE_POST_SCALE_BIAS, - "alibi": NVTE_Bias_Type.NVTE_ALIBI, -} - -AttnMaskType = { - "no_mask": NVTE_Mask_Type.NVTE_NO_MASK, - "padding": NVTE_Mask_Type.NVTE_PADDING_MASK, - "causal": NVTE_Mask_Type.NVTE_CAUSAL_MASK, - "padding_causal": NVTE_Mask_Type.NVTE_PADDING_CAUSAL_MASK, - "causal_bottom_right": NVTE_Mask_Type.NVTE_CAUSAL_BOTTOM_RIGHT_MASK, - "padding_causal_bottom_right": NVTE_Mask_Type.NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK, -} - -SoftmaxType = { - "vanilla": NVTE_Softmax_Type.NVTE_VANILLA_SOFTMAX, - "off-by-one": NVTE_Softmax_Type.NVTE_OFF_BY_ONE_SOFTMAX, - "learnable": NVTE_Softmax_Type.NVTE_LEARNABLE_SOFTMAX, -} - class FusedAttnBackend(IntEnum): - """Fused attention sub-backends. - - This is the canonical fused-attention backend enum for - ``transformer_engine.pytorch``. It mirrors the backend - ``transformer_engine_torch.NVTE_Fused_Attn_Backend`` (pybind11) enum - value-for-value, and instances of the two enums compare equal when they - share the same integer value. Unlike the pybind enum, a plain-python - ``IntEnum`` is traceable by ``torch.compile``: comparisons against a member - constant-fold cleanly. Lookup by name (``FusedAttnBackend["FP8"]``) works - the same way as with the dict this used to be. - - Members do not survive a graph break, though, so ``get_attention_backend`` - returns the sub-backend as a plain int and ``cast`` turns it back into a - member. - """ + """Legacy import-path mirror of the Python cuDNN attention backend enum.""" - No_Backend = int(NVTE_Fused_Attn_Backend.NVTE_No_Backend) - F16_arbitrary_seqlen = int(NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen) - FP8 = int(NVTE_Fused_Attn_Backend.NVTE_FP8) + No_Backend = -1 + F16_arbitrary_seqlen = 1 + FP8 = 2 @classmethod - def cast( - cls, backend: "Union[FusedAttnBackend, NVTE_Fused_Attn_Backend]" - ) -> "FusedAttnBackend": - """Normalize a backend value to the canonical ``FusedAttnBackend`` member. - - The pybind ``transformer_engine_torch.NVTE_Fused_Attn_Backend`` enum is - accepted as input for backward compatibility and mapped to the matching - ``FusedAttnBackend`` member. - """ + def cast(cls, backend): + """Convert an integer or another compatible enum to this enum.""" if isinstance(backend, cls): return backend return cls(int(backend)) - def __eq__(self, other: object) -> bool: - # ``FusedAttnBackend`` is an ``IntEnum`` while ``NVTE_Fused_Attn_Backend`` - # is a pybind11 enum. Compare by integer value so the two enums stay - # equivalent regardless of the pybind11 version (the pybind ``__eq__`` - # handles the reverse order). - if isinstance(other, NVTE_Fused_Attn_Backend): - return int(self) == int(other) - return int.__eq__(self, other) - - def __ne__(self, other: object) -> bool: - result = self.__eq__(other) - if result is NotImplemented: - return result - return not result - - def __hash__(self) -> int: - return int.__hash__(self) - - -# Fail fast at import time if a new enumerator is added on the C++ side -# without being mirrored above. -assert {f"NVTE_{m.name}" for m in FusedAttnBackend} == set(NVTE_Fused_Attn_Backend.__members__), ( - "FusedAttnBackend in python is out of sync with" - " transformer_engine_torch.NVTE_Fused_Attn_Backend defined on the C++ side." - " Please make sure TE C++ and python are in sync." -) - -BACKEND_FP8_THREADS_PER_CTA = 128 -BACKEND_F16arb_ELTS_PER_THREADS = 16 META_QKV = FP8FwdTensorIdx.GEMM1_OUTPUT META_DQKV = FP8BwdTensorIdx.GRAD_OUTPUT1 @@ -172,482 +61,21 @@ def __hash__(self) -> int: META_DP = FP8BwdTensorIdx.GRAD_INPUT3 -def fused_attn_fwd( - is_training: bool, - max_seqlen_q: int, - max_seqlen_kv: int, - cu_seqlens_q: torch.Tensor, - cu_seqlens_kv: torch.Tensor, - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - fake_dtype: torch.dtype, - fused_attention_backend: FusedAttnBackend, - attn_bias: torch.Tensor = None, - cu_seqlens_q_padded: torch.Tensor = None, - cu_seqlens_kv_padded: torch.Tensor = None, - page_table_k: torch.Tensor = None, - page_table_v: torch.Tensor = None, - s_quantizer: Quantizer = None, - o_quantizer: Quantizer = None, - attn_scale: float = None, - dropout: float = 0.0, - fast_zero_fill: bool = True, - qkv_layout: str = "sbh3d", - o_format: str = "sbhd", - qkv_scale_inv_format: str = None, - attn_bias_type: str = "no_bias", - attn_mask_type: str = "padding", - softmax_type: str = "vanilla", - window_size: Tuple[int, int] = (-1, -1), - bottom_right_diagonal: bool = None, - rng_gen: torch.Generator = None, - softmax_offset: torch.Tensor = None, - return_max_logit: bool = False, - cuda_graph: bool = False, -) -> Tuple[Union[torch.Tensor, None], ...]: - """Fused Attention FWD for separate QKV input. - - Parameters - ---------- - is_training : bool - if True, runs training and produces auxiliary tensors aux_ctx_tensors - for the backward; if False, runs inference and doesn't produce aux_ctx_tensors - max_seqlen_q : int - max sequence length for Q, used for padding; - may be larger than max(seqlens_q), - seqlens_q = cu_seqlens_q[1:] - cu_seqlens_q[:-1] - max_seqlen_kv : int - max sequence length for K and V, used for padding; - may be larger than max(seqlens_kv), - seqlens_kv = cu_seqlens_kv[1:] - cu_seqlens_kv[:-1] - cu_seqlens_q : torch.Tensor - cumulative sequence lengths for Q; shape [batch_size + 1] - cu_seqlens_kv : torch.Tensor - cumulative sequence lengths for K and V; shape [batch_size + 1] - q : torch.Tensor - input tensor Q; shape sbhd, bshd or thd (see `qkv_layout` for details) - k : torch.Tensor - input tensor K; shape sbhd, bshd or thd (see `qkv_layout` for details) - v : torch.Tensor - input tensor V; shape sbhd, bshd or thd (see `qkv_layout` for details) - fake_dtype : DType - data type of Q, K and V - in case of high precision, fake dtype in case of FP8; - in torch.dtype - fused_attention_backend : FusedAttnBackend - please see FusedAttention module for details on supported backends. - attn_bias : torch.Tensor, default = None - input tensor Bias when attn_bias_type is "pre_scale_bias" or "post_scale_bias"; - shape [1, num_heads, max_seqlen_q, max_seqlen_kv], same data type as q, k and v - cu_seqlens_q_padded : torch.Tensor, default = None - cumulative sequence offsets for Q; shape [batch_size + 1] - cu_seqlens_kv_padded : torch.Tensor, default = None - cumulative sequence offsets for KV; shape [batch_size + 1] - page_table_k : torch.Tensor, default = None - page table for K cache; shape [batch_size, max_pages_per_seq_k] - page_table_v : torch.Tensor, default = None - page table for V cache; shape [batch_size, max_pages_per_seq_v] - s_quantizer : Quantizer, default = None - Quantizer object for the intermediate value S. - o_quantizer : Quantizer, default = None - Quantizer object for the output of the attention. - attn_scale : float, default = None - if not None, use attn_scale as the attention scale for Q*K.T BMM; - if None, use 1.0/sqrt(head_dim_qk) as the default - dropout : float, default = 0.0 - dropout probability, 0.0 means no dropout, 1.0 means no output; - dropout must be 0.0 if is_training is False - fast_zero_fill : bool, default = True - if True, initializes the output tensor O to zero using the fast filling method; - if False, uses PyTorch's .fill_() method - qkv_layout : str, default = "sbh3d" - layout of Q, K and V; - {"sb3hd", "sbh3d", "sbhd_sb2hd", "sbhd_sbh2d", "sbhd_sbhd_sbhd", - "bs3hd", "bsh3d", "bshd_bs2hd", "bshd_bsh2d", "bshd_bshd_bshd", - "t3hd", "th3d", "thd_t2hd", "thd_th2d", "thd_thd_thd"} - o_format : str, default = "sbhd" - format of O; {"sbhd", "bshd", "thd"} - qkv_scale_inv_format : str, default = None - format of the scale-inverse tensors for QKV; {"sbhd", "bshd", "thd", "bhsd"}; - if None, defaults to the format inferred from qkv_layout. - attn_bias_type : str, default = "no_bias" - type of the bias; {"no_bias", "pre_scale_bias", "post_scale_bias", "alibi"} - attn_mask_type : str, default = "padding" - type of the attention mask; {"padding", "causal", "padding_causal", "no_mask"} - softmax_type : str, default = "vanilla" - type of the attention softmax; {"vanilla", "off-by-one", "learnable"} - window_size : Tuple[int, int], default = (-1, -1) - sliding window size for local attention, where query at position i attends to keys - in [i + seqlen_k - seqlen_q - window_size[0], i + seqlen_k - seqlen_q - + window_size[1]] inclusive. Special cases (-1, -1) and (-1, 0) mean no sliding - window and causal mask specifically. - bottom_right_diagonal: bool, default = None - whether to align sliding window and ALiBi diagonal to the top left (False) or - bottom right (True) corner of the softmax matrix. - rng_gen : torch.Generator, default = None - random number generator; - if None, uses the default CUDA generator from PyTorch; otherwise, uses rng_gen - softmax_offset : torch.Tensor, default = None - softmax offset tensor of shape [1, h_q, 1, 1]. - See softmax_type in DotProductAttention for details. - return_max_logit : bool, default = False - whether to return the maximum attention score - cuda_graph : bool, default = False - whether or not cuda graph capture is enabled. - - Returns - ---------- - o : torch.Tensor - output tensor O, of the attention calculation; same data type as Q, K and V; - same shape as Q - aux_ctx_tensors : List[torch.Tensor] - auxiliary output tensors used for the backward; - if is_training is True, aux_ctx_tensors = [softmax-related tensors, rng_state] - if is_training is False, aux_ctx_tensors = None - - softmax-related tensors: - 1. if fused_attention_backend == FusedAttnBackend["F16_arbitrary_seqlen"] - softmaxStats: torch.Tensor - log(sum(e^(x - max(x)))), where x=Q*K.T - shape [batch_size, num_heads, max_seqlen_q, 1], dtype float32 - Max: torch.Tensor, only when return_max_logit is True - shape [batch_size, num_heads, max_seqlen_q, 1], dtype float32 - 2. if fused_attention_backend == FusedAttnBackend["FP8"] - softmaxStats: torch.Tensor - log(sum(e^(x - max(x)))), where x=Q*K.T - shape [batch_size, num_heads, max_seqlen_q, 1], dtype float32 - rng_state: torch.Tensor - state of the random number generator; - [seed, offset], dtype uint64 - max_logit : if return_max_logit = True, shape [h] and same data type as O; otherwise None - """ - - if bottom_right_diagonal is None: - bottom_right_diagonal = attn_mask_type in { - "causal_bottom_right", - "padding_causal_bottom_right", - } - - if attn_scale is None: - d = q.size(-1) - attn_scale = 1.0 / math.sqrt(d) - - if attn_bias_type not in ["no_bias", "alibi"]: - if attn_bias is None: - raise ValueError( - f"attn_bias tensor cannot be None when attn_bias_type={attn_bias_type!r}." - ) - if attn_bias.dtype != q.dtype: - raise ValueError( - "attn_bias tensor must have the same dtype as q and kv: " - f"attn_bias.dtype={attn_bias.dtype} but q.dtype={q.dtype}." - ) - - # Accept the pybind enum for backward compatibility. - fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) - if fused_attention_backend == FusedAttnBackend["No_Backend"]: - raise ValueError( - "Fused attention does not support this input combination:" - f" qkv_layout={qkv_layout!r}, attn_bias_type={attn_bias_type!r}," - f" attn_mask_type={attn_mask_type!r}, q.shape={list(q.shape)}," - f" q.dtype={q.dtype}, backend={fused_attention_backend}." - ) - - if fused_attention_backend == FusedAttnBackend["F16_arbitrary_seqlen"]: - rng_elts_per_thread = BACKEND_F16arb_ELTS_PER_THREADS - # FP8 fused attention API from fmha_v2 - elif fused_attention_backend == FusedAttnBackend["FP8"]: - rng_elts_per_thread = ( - max_seqlen_q * max_seqlen_q + BACKEND_FP8_THREADS_PER_CTA - 1 - ) // BACKEND_FP8_THREADS_PER_CTA - else: - raise ValueError(f"Unsupported backend {fused_attention_backend}") - - # execute kernel - output_tensors = tex.fused_attn_fwd( - max_seqlen_q, - max_seqlen_kv, - is_training, - attn_scale, - dropout, - fast_zero_fill, - QKVLayout[qkv_layout], - QKVFormat[o_format], - QKVFormat[qkv_scale_inv_format], - AttnBiasType[attn_bias_type], - AttnMaskType[attn_mask_type], - SoftmaxType[softmax_type], - window_size, - bottom_right_diagonal, - cu_seqlens_q, - cu_seqlens_kv, - q, - k, - v, - fake_dtype, - cu_seqlens_q_padded, - cu_seqlens_kv_padded, - page_table_k, - page_table_v, - s_quantizer, - o_quantizer, - attn_bias, - softmax_offset, - rng_gen, - rng_elts_per_thread, - return_max_logit, - cuda_graph, +# Resolve lazily because cpp_extensions is imported while transformer_engine.pytorch +# itself is still initializing. +def fused_attn_fwd(*args, **kwargs): + """Run fused attention forward through the cuDNN frontend Python API.""" + from transformer_engine.pytorch.attention.dot_product_attention.cudnn_attention import ( + fused_attn_fwd as implementation, ) - if return_max_logit: - qkv_format = qkv_layout.replace("3", "").replace("2", "").split("_")[0] - # thd (newer cuDNN runtimes, non-sm120): output_tensors: out [tq, h, d], Stats [tq, h, 1], Max [tq, h, 1] - # thd (older cuDNN runtimes or sm120): output_tensors: out [tq, h, d], Stats [b, h, sq, 1], Max [b, h, sq, 1] - # bshd: output_tensors: out [b, sq, h, d], Stats [b, h, sq, 1], Max [b, h, sq, 1] - # sbhd: output_tensors: out [sq, b, h, d], Stats [b, h, sq, 1], Max [b, h, sq, 1] - aux_ctx_tensors = [output_tensors[1]] + list( - output_tensors[3:] - ) # Stats + rng_state + optional tensors - max_tensor = output_tensors[2] - amax_dims = (0, 2) if max_tensor.ndim == 3 else (0, 2, 3) - - if qkv_format == "thd": - if max_tensor.ndim == 4: - # For THD on cuDNN <= 9.6 or THD on sm120, Max tensor can be [b, h, sq, 1] - # with padded sequence positions. Exclude those padded positions when computing max_logit. - seqlens_q = (cu_seqlens_q[1:] - cu_seqlens_q[:-1]).to(device=max_tensor.device) - sq_idx = torch.arange(max_tensor.shape[2], device=max_tensor.device).view( - 1, 1, -1, 1 - ) - valid = sq_idx < seqlens_q.view(-1, 1, 1, 1) - max_tensor = max_tensor.masked_fill(~valid, float("-inf")) - elif max_tensor.ndim == 3: - if cu_seqlens_q_padded is not None: - # Exclude padding; CP may pass nonzero padded offsets. - actual_seqlens = (cu_seqlens_q[1:] - cu_seqlens_q[:-1]).to( - device=max_tensor.device - ) - tq = max_tensor.shape[0] - starts = cu_seqlens_q_padded[:-1].to(device=max_tensor.device) - ends = (starts + actual_seqlens).clamp(max=tq) - delta = torch.zeros(tq + 1, dtype=torch.int32, device=max_tensor.device) - updates = torch.ones_like(starts, dtype=torch.int32) - delta.scatter_add_(0, starts.clamp(max=tq), updates) - delta.scatter_add_(0, ends, -updates) - valid = delta[:-1].cumsum(0) > 0 - max_tensor = max_tensor.masked_fill(~valid.view(-1, 1, 1), float("-inf")) - - # Max -> max_logit [h] - max_logit = torch.amax(max_tensor, dim=amax_dims).to(dtype=output_tensors[0].dtype) - return output_tensors[0], aux_ctx_tensors, max_logit - - # out, aux_ctx_tensors - return output_tensors[0], output_tensors[1:] - - -def fused_attn_bwd( - max_seqlen_q: int, - max_seqlen_kv: int, - cu_seqlens_q: torch.Tensor, - cu_seqlens_kv: torch.Tensor, - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - o: torch.Tensor, - d_o: torch.Tensor, - fake_dtype: torch.dtype, - aux_ctx_tensors: List[torch.Tensor], - fused_attention_backend: FusedAttnBackend, - cu_seqlens_q_padded: torch.Tensor = None, - cu_seqlens_kv_padded: torch.Tensor = None, - s_quantizer: Quantizer = None, - dp_quantizer: Quantizer = None, - dqkv_quantizer: Quantizer = None, - attn_scale: Optional[float] = None, - dropout: float = 0.0, - fast_zero_fill: bool = True, - qkv_layout: str = "sbh3d", - o_format: str = "sbhd", - do_format: str = "sbhd", - dqkv_layout: str = "sbh3d", - qkv_scale_inv_format: str = None, - do_scale_inv_format: str = None, - attn_bias_type: str = "no_bias", - attn_mask_type: str = "padding", - softmax_type: str = "vanilla", - window_size: Tuple[int, int] = (-1, -1), - bottom_right_diagonal: bool = None, - deterministic: bool = False, - cuda_graph: bool = False, -) -> Tuple[Union[torch.Tensor, None], ...]: - """Fused Attention BWD for packed KV input. - - Parameters - ---------- - max_seqlen_q : int - max sequence length for Q, used for padding; may be larger than max(seqlens_q), - seqlens_q = cu_seqlens_q[1:] - cu_seqlens_q[:-1] - max_seqlen_kv : int - max sequence length for K and V, used for padding; - may be larger than max(seqlens_kv), - seqlens_kv = cu_seqlens_kv[1:] - cu_seqlens_kv[:-1] - cu_seqlens_q : torch.Tensor - cumulative sequence lengths for Q; shape [batch_size + 1] - cu_seqlens_kv : torch.Tensor - cumulative sequence lengths for K and V; shape [batch_size + 1] - q : torch.Tensor - input tensor Q; shape sbhd, bshd or thd (see `qkv_layout` for details) - k : torch.Tensor - input tensor K; shape sbhd, bshd or thd (see `qkv_layout` for details) - v : torch.Tensor - input tensor V; shape sbhd, bshd or thd (see `qkv_layout` for details) - o : torch.Tensor - input tensor O (output of forward); same data type as Q, K and V; - same shape as Q - d_o : torch.Tensor - input tensor dO (gradient of O); same data type as Q, K and V; - same shape as Q - fake_dtype : DType - data type of Q, K and V - in case of high precision, fake dtype in case of FP8; - in torch.dtype - aux_ctx_tensors : List[torch.Tensor] - auxiliary output tensors of the forward pass when its is_training is True, - e.g. aux_ctx_tensors = [S, Max, rng_state] - fused_attention_backend : FusedAttnBackend - please see FusedAttention module for details on supported backends. - cu_seqlens_q_padded : torch.Tensor, default = None - cumulative sequence offsets for Q; shape [batch_size + 1] - cu_seqlens_kv_padded : torch.Tensor, default = None - cumulative sequence offsets for KV; shape [batch_size + 1] - s_quantizer : Quantizer, default = None - Quantizer object for the intermediate value S. - dp_quantizer : Quantizer, default = None - Quantizer object for the intermediate value dP. - dqkv_quantizer : Quantizer, default = None - Quantizer object for the output values of the fused_attn_bwd. - attn_scale : float, default = None - if not None, use attn_scale as the attention scale for Q*K.T BMM; - if None, use 1.0/sqrt(head_dim_qk) as the default - dropout : float, default = 0.0 - dropout probability, 0.0 means no dropout, 1.0 means no output; - dropout must be 0.0 if is_training is False - fast_zero_fill : bool, default = True - if True, initializes the output tensor O to zero using the fast filling method; - if False, uses PyTorch's .fill_() method - qkv_layout : str, default = "sbh3d" - layout of Q, K and V; - {"sb3hd", "sbh3d", "sbhd_sb2hd", "sbhd_sbh2d", "sbhd_sbhd_sbhd", - "bs3hd", "bsh3d", "bshd_bs2hd", "bshd_bsh2d", "bshd_bshd_bshd", - "t3hd", "th3d", "thd_t2hd", "thd_th2d", "thd_thd_thd"} - o_format : str, default = "sbhd" - format of O; {"sbhd", "bshd", "thd"} - do_format : str, default = "sbhd" - format of dO; {"sbhd", "bshd", "thd"} - dqkv_layout : str, default = "sbh3d" - layout of dQ, dK and dV; - {"sb3hd", "sbh3d", "sbhd_sb2hd", "sbhd_sbh2d", "sbhd_sbhd_sbhd", - "bs3hd", "bsh3d", "bshd_bs2hd", "bshd_bsh2d", "bshd_bshd_bshd", - "t3hd", "th3d", "thd_t2hd", "thd_th2d", "thd_thd_thd"} - qkv_scale_inv_format : str, default = None - format of the scale-inverse tensors for QKV; {"sbhd", "bshd", "thd", "bhsd"}; - if None, defaults to the format inferred from qkv_layout. - do_scale_inv_format : str, default = None - format of the scale-inverse tensors for dO; {"sbhd", "bshd", "thd", "bhsd"}; - if None, defaults to the format inferred from the output layout. - attn_bias_type : str, default = "no_bias" - type of the bias; {"no_bias", "pre_scale_bias", "post_scale_bias", "alibi"} - attn_mask_type : str, default = "padding" - type of the attention mask; {"padding", "causal", "padding_causal", "no_mask"} - softmax_type : str, default = "vanilla" - type of the attention softmax; {"vanilla", "off-by-one", "learnable"} - window_size : Tuple[int, int], default = (-1, -1) - sliding window size for local attention, where query at position i attends to keys - in [i + seqlen_k - seqlen_q - window_size[0], i + seqlen_k - seqlen_q - + window_size[1]] inclusive. Special cases (-1, -1) and (-1, 0) mean no sliding - window and causal mask specifically. - bottom_right_diagonal: bool, default = None - whether to align sliding window and ALiBi diagonal to the top left (False) or - bottom right (True) corner of the softmax matrix. - deterministic : bool, default = False - whether to execute the backward pass with deterministic behaviours. - cuda_graph : bool, default = False - whether or not cuda graph capture is enabled. - - Returns - ---------- - d_q : torch.Tensor - gradient tensor of Q; same data type and shape as Q - d_k : torch.Tensor - gradient tensor of K; same data type and shape as K - d_v : torch.Tensor - gradient tensor of V; same data type and shape as V - d_bias : torch.Tensor, optional - gradient tensor of Bias when attn_bias_type is "pre_scale_bias" - or "post_scale_bias"; same data type and shape as Bias - d_softmax_offset : torch.Tensor, optional - gradient tensor of softmax offset of shape [1, h_q, 1, 1]. - See softmax_type in DotProductAttention for details. - """ - if bottom_right_diagonal is None: - bottom_right_diagonal = attn_mask_type in { - "causal_bottom_right", - "padding_causal_bottom_right", - } - - if attn_scale is None: - d = q.size(-1) - attn_scale = 1.0 / math.sqrt(d) - - # Accept the pybind enum for backward compatibility. - fused_attention_backend = FusedAttnBackend.cast(fused_attention_backend) - if fused_attention_backend == FusedAttnBackend["No_Backend"]: - raise ValueError( - "Fused attention backward does not support this input combination:" - f" qkv_layout={qkv_layout!r}, attn_bias_type={attn_bias_type!r}," - f" attn_mask_type={attn_mask_type!r}, q.shape={list(q.shape)}," - f" q.dtype={q.dtype}, backend={fused_attention_backend}." - ) + return implementation(*args, **kwargs) - if len(aux_ctx_tensors) < 1: - raise ValueError( - "aux_ctx_tensors must contain rng_state as its last element," - f" but got len(aux_ctx_tensors)={len(aux_ctx_tensors)}" - f" for backend={fused_attention_backend}." - ) - output_tensors = tex.fused_attn_bwd( - max_seqlen_q, - max_seqlen_kv, - attn_scale, - dropout, - fast_zero_fill, - QKVLayout[qkv_layout], - QKVFormat[o_format], - QKVFormat[do_format], - QKVLayout[dqkv_layout], - QKVFormat[qkv_scale_inv_format], - QKVFormat[do_scale_inv_format], - AttnBiasType[attn_bias_type], - AttnMaskType[attn_mask_type], - SoftmaxType[softmax_type], - window_size, - bottom_right_diagonal, - deterministic, - cu_seqlens_q, - cu_seqlens_kv, - q, - k, - v, - o, - d_o, - fake_dtype, - aux_ctx_tensors, - cu_seqlens_q_padded, - cu_seqlens_kv_padded, - s_quantizer, - dp_quantizer, - dqkv_quantizer, - cuda_graph, +def fused_attn_bwd(*args, **kwargs): + """Run fused attention backward through the cuDNN frontend Python API.""" + from transformer_engine.pytorch.attention.dot_product_attention.cudnn_attention import ( + fused_attn_bwd as implementation, ) - return output_tensors + return implementation(*args, **kwargs) diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index c23574759e2..cf2e4fa4181 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -94,40 +94,8 @@ std::tuple moe_unpermute_bwd(at::Tensor input_bwd, at::T * Attention **************************************************************************************************/ -NVTE_Fused_Attn_Backend get_fused_attn_backend( - bool is_training, const DType q_dtype, const DType kv_dtype, NVTE_QKV_Layout qkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, - float p_dropout, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool cuda_graph, bool deterministic); - -std::vector fused_attn_fwd( - size_t max_seqlen_q, size_t max_seqlen_kv, bool is_training, float attn_scale, float p_dropout, - bool set_zero, NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, const at::Tensor cu_seqlens_q, const at::Tensor cu_seqlens_kv, - const py::handle Q, const py::handle K, const py::handle V, const at::ScalarType fake_dtype, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, - const std::optional page_table_k, const std::optional page_table_v, - py::handle s_quantizer, py::handle o_quantizer, const std::optional Bias, - const std::optional SoftmaxOffset, const std::optional rng_gen, - size_t rng_elts_per_thread, bool return_max_logit, bool cuda_graph); - -std::vector fused_attn_bwd( - size_t max_seqlen_q, size_t max_seqlen_kv, float attn_scale, float p_dropout, bool set_zero, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, - NVTE_QKV_Layout dqkv_layout, NVTE_QKV_Format qkv_scale_inv_format, - NVTE_QKV_Format do_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, bool deterministic, const at::Tensor cu_seqlens_q, - const at::Tensor cu_seqlens_kv, const py::handle Q, const py::handle K, const py::handle V, - const py::handle O, const py::handle dO, const at::ScalarType fake_dtype, - const std::vector Aux_CTX_Tensors, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, - py::handle dp_quantizer, py::handle dqkv_quantizer, bool cuda_graph); +at::Tensor get_cudnn_attention_rng_state(const std::optional rng_gen, + size_t increment); at::Tensor fa_prepare_fwd(at::Tensor qkvi); at::Tensor fa_prepare_bwd(at::Tensor q, at::Tensor k, at::Tensor v); diff --git a/transformer_engine/pytorch/csrc/extensions/attention.cpp b/transformer_engine/pytorch/csrc/extensions/attention.cpp index eb8813d4a0a..0bd33a4e5a0 100644 --- a/transformer_engine/pytorch/csrc/extensions/attention.cpp +++ b/transformer_engine/pytorch/csrc/extensions/attention.cpp @@ -8,606 +8,8 @@ #include "common.h" #include "pybind.h" -namespace { - -constexpr int block_size = 512; - -// fast zero-fills of tensors -void mha_fill(const transformer_engine::TensorWrapper &self, const at::Tensor &start_index) { - std::vector shape = transformer_engine::pytorch::convertShape(self.shape()); - - auto max_tokens = shape[0]; - auto fcd_size = 1; - for (size_t i = 1; i <= shape.size(); i++) { - fcd_size *= shape[i]; - } - - NVTE_CHECK(fcd_size % block_size == 0, "input size not aligned to block size"); - - size_t element_size_bits = transformer_engine::pytorch::typeToNumBits(self.dtype()); - int32_t start_row = start_index.data_ptr()[0]; - void *base_ptr = static_cast(self.get_rowwise_data().data_ptr) + - static_cast(start_row) * fcd_size * element_size_bits / 8; - size_t num_rows_to_zero = max_tokens - start_row; - size_t total_bytes = num_rows_to_zero * fcd_size * element_size_bits / 8; - - NVTE_SCOPED_GIL_RELEASE( - { nvte_memset(base_ptr, 0, total_bytes, at::cuda::getCurrentCUDAStream()); }); -} - -} // namespace - namespace transformer_engine::pytorch { -// get the fused attention backend -NVTE_Fused_Attn_Backend get_fused_attn_backend( - bool is_training, const DType q_dtype, const DType kv_dtype, NVTE_QKV_Layout qkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, - float p_dropout, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool cuda_graph, bool deterministic) { - NVTE_Fused_Attn_Backend fused_attention_backend = nvte_get_fused_attn_backend( - is_training, static_cast(q_dtype), static_cast(kv_dtype), qkv_layout, - bias_type, attn_mask_type, softmax_type, p_dropout, num_attn_heads, num_gqa_groups, - max_seqlen_q, max_seqlen_kv, head_dim_qk, head_dim_v, window_size_left, window_size_right, - return_max_logit, cuda_graph, deterministic); - return fused_attention_backend; -} - -// helper function for S and dP quantizers -std::tuple> quantizer_helper( - py::handle quantizer, const std::vector &shape, DType dtype, bool create_hp_tensor, - std::optional data) { - std::unique_ptr T_quantizer = convert_quantizer(quantizer); - TensorWrapper te_T; - py::object py_T; - std::optional amax_buf; - if (quantizer.is_none()) { - // high precision - auto *none_quantizer = dynamic_cast(T_quantizer.get()); - if (data.has_value()) { - std::tie(te_T, py_T) = none_quantizer->create_tensor(shape, dtype, data.value()); - } else { - std::tie(te_T, py_T) = none_quantizer->create_tensor(shape, dtype); - } - } else if (detail::IsFloat8Quantizers(quantizer.ptr())) { - // delayed scaling; this helps initialize scale_inv - auto *T_quantizer_fp8 = dynamic_cast(T_quantizer.get()); - std::tie(te_T, py_T) = - T_quantizer_fp8->create_tensor(shape, dtype, data, std::nullopt, std::nullopt); - } else if (detail::IsFloat8CurrentScalingQuantizers(quantizer.ptr())) { - // current scaling - auto *T_quantizer_fp8 = dynamic_cast(T_quantizer.get()); - if (create_hp_tensor) { - if (data.has_value()) { - std::tie(te_T, py_T, amax_buf) = - T_quantizer_fp8->create_unquantized_tensor_with_amax(shape, dtype, data.value()); - } else { - std::tie(te_T, py_T, amax_buf) = - T_quantizer_fp8->create_unquantized_tensor_with_amax(shape, dtype); - } - } else { - std::tie(te_T, py_T) = T_quantizer_fp8->create_tensor(shape, dtype); - NVTE_CHECK( - !data.has_value(), - "Float8CurrentScalingQuantizer::create_tensor() does not take data tensor as input!"); - } - } else if (detail::IsMXFP8Quantizers(quantizer.ptr())) { - // MXFP8 - if (create_hp_tensor) { - if (data.has_value()) { - std::tie(te_T, py_T) = NoneQuantizer(py::none()).create_tensor(shape, dtype, data.value()); - } else { - std::tie(te_T, py_T) = NoneQuantizer(py::none()).create_tensor(shape, dtype); - } - } else { - auto *T_quantizer_fp8 = dynamic_cast(T_quantizer.get()); - std::tie(te_T, py_T) = T_quantizer_fp8->create_tensor(shape, dtype); - NVTE_CHECK(!data.has_value(), - "MXFP8Quantizer::create_tensor() does not take data tensor as input!"); - } - } - return {std::move(te_T), std::move(py_T), std::move(amax_buf)}; -} - -// fused attention FWD with separate Q, K and V tensors -std::vector fused_attn_fwd( - size_t max_seqlen_q, size_t max_seqlen_kv, bool is_training, float attn_scale, float p_dropout, - bool set_zero, NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, const at::Tensor cu_seqlens_q, const at::Tensor cu_seqlens_kv, - const py::handle Q, const py::handle K, const py::handle V, const at::ScalarType fake_dtype, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, - const std::optional page_table_k, const std::optional page_table_v, - py::handle s_quantizer, py::handle o_quantizer, const std::optional Bias, - const std::optional SoftmaxOffset, const std::optional rng_gen, - size_t rng_elts_per_thread, bool return_max_logit, bool cuda_graph) { - // Ensure that cuDNN handle is created on the correct device, - // overriding torch.cuda.set_device calls from user side. - // Assumes all tensors passed are on the same device. - at::cuda::CUDAGuard device_guard(cu_seqlens_q.device()); - - auto none = py::none(); - - // create QKV tensor wrappers - TensorWrapper te_Q, te_K, te_V; - te_Q = makeTransformerEngineTensor(Q, none); - te_K = makeTransformerEngineTensor(K, none); - te_V = makeTransformerEngineTensor(V, none); - const DType qkv_type = te_Q.dtype(); - - // create S tensor - auto [te_S, py_S, _] = quantizer_helper(s_quantizer, {0}, DType::kFloat32, false, std::nullopt); - - // create O tensor - std::unique_ptr O_quantizer = convert_quantizer(o_quantizer); - std::vector q_shape = convertShape(te_Q.shape()); - std::vector v_shape = convertShape(te_V.shape()); - auto o_shape_tmp = std::vector{q_shape.begin(), q_shape.end()}; - o_shape_tmp[o_shape_tmp.size() - 1] = v_shape[v_shape.size() - 1]; - auto o_shape = std::vector{o_shape_tmp.begin(), o_shape_tmp.end()}; - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - AttentionShape o_parsed(q_format, o_shape_tmp.data()); - size_t h = o_parsed.h(), d = o_parsed.d(); - o_parsed.to_format(o_format, o_shape.data()); - const DType fake_dtype_te = GetTransformerEngineDType(fake_dtype); - auto [te_O, py_O, o_amax_buf] = - quantizer_helper(o_quantizer, o_shape, fake_dtype_te, true, std::nullopt); - - // construct NVTE tensors - TensorWrapper te_Bias; - TensorWrapper te_cu_seqlens_q, te_cu_seqlens_kv; - TensorWrapper te_cu_seqlens_q_padded, te_cu_seqlens_kv_padded; - TensorWrapper te_page_table_k, te_page_table_v; - if (qkv_type == DType::kFloat8E4M3 || qkv_type == DType::kFloat8E5M2) { - // FP8 - if (set_zero && (o_format == NVTE_QKV_Format::NVTE_THD)) { - if ((h * d) % block_size == 0) { - mha_fill(te_O, cu_seqlens_q.index({torch::indexing::Slice(-1, torch::indexing::None)})); - } else { - te_O.zero_(at::cuda::getCurrentCUDAStream()); - } - } - } else if (qkv_type == DType::kBFloat16 || qkv_type == DType::kFloat16) { - if (o_format == NVTE_QKV_Format::NVTE_THD) { - te_O.zero_(at::cuda::getCurrentCUDAStream()); - } - } else { - NVTE_ERROR("Fused attention only supports FP8 and BF16/FP16 data types. \n"); - } - if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI) && (Bias.has_value())) { - auto bias_sizes = Bias.value().sizes().vec(); - std::vector bias_shape{bias_sizes.begin(), bias_sizes.end()}; - te_Bias = makeTransformerEngineTensor(Bias.value().data_ptr(), bias_shape, DType::kFloat32); - } - auto cu_seqlens_q_sizes = cu_seqlens_q.sizes().vec(); - std::vector cu_seqlens_q_shape{cu_seqlens_q_sizes.begin(), cu_seqlens_q_sizes.end()}; - auto cu_seqlens_kv_sizes = cu_seqlens_kv.sizes().vec(); - std::vector cu_seqlens_kv_shape{cu_seqlens_kv_sizes.begin(), cu_seqlens_kv_sizes.end()}; - te_cu_seqlens_q = - makeTransformerEngineTensor(cu_seqlens_q.data_ptr(), cu_seqlens_q_shape, DType::kInt32); - te_cu_seqlens_kv = - makeTransformerEngineTensor(cu_seqlens_kv.data_ptr(), cu_seqlens_kv_shape, DType::kInt32); - - if ((cu_seqlens_q_padded.has_value()) && (cu_seqlens_kv_padded.has_value())) { - auto cu_seqlens_q_padded_sizes = cu_seqlens_q_padded.value().sizes().vec(); - std::vector cu_seqlens_q_padded_shape{cu_seqlens_q_padded_sizes.begin(), - cu_seqlens_q_padded_sizes.end()}; - auto cu_seqlens_kv_padded_sizes = cu_seqlens_kv_padded.value().sizes().vec(); - std::vector cu_seqlens_kv_padded_shape{cu_seqlens_kv_padded_sizes.begin(), - cu_seqlens_kv_padded_sizes.end()}; - te_cu_seqlens_q_padded = makeTransformerEngineTensor(cu_seqlens_q_padded.value().data_ptr(), - cu_seqlens_q_padded_shape, DType::kInt32); - te_cu_seqlens_kv_padded = makeTransformerEngineTensor( - cu_seqlens_kv_padded.value().data_ptr(), cu_seqlens_kv_padded_shape, DType::kInt32); - } - - if ((page_table_k.has_value()) && (page_table_v.has_value())) { - auto page_table_k_sizes = page_table_k.value().sizes().vec(); - std::vector page_table_k_shape{page_table_k_sizes.begin(), page_table_k_sizes.end()}; - auto page_table_v_sizes = page_table_v.value().sizes().vec(); - std::vector page_table_v_shape{page_table_v_sizes.begin(), page_table_v_sizes.end()}; - te_page_table_k = - makeTransformerEngineTensor(page_table_k.value().data_ptr(), page_table_k_shape, - DType::kInt32, nullptr, nullptr, nullptr); - te_page_table_v = - makeTransformerEngineTensor(page_table_v.value().data_ptr(), page_table_v_shape, - DType::kInt32, nullptr, nullptr, nullptr); - } - - // softmax offset - TensorWrapper te_SoftmaxOffset; - if ((softmax_type != NVTE_VANILLA_SOFTMAX) && (SoftmaxOffset.has_value())) { - auto SoftmaxOffset_sizes = SoftmaxOffset.value().sizes().vec(); - std::vector SoftmaxOffset_shape{SoftmaxOffset_sizes.begin(), SoftmaxOffset_sizes.end()}; - te_SoftmaxOffset = - makeTransformerEngineTensor(SoftmaxOffset.value().data_ptr(), SoftmaxOffset_shape, - DType::kFloat32, nullptr, nullptr, nullptr); - } - - // extract rng seed and offset - auto gen = at::get_generator_or_default( - rng_gen, at::cuda::detail::getDefaultCUDAGenerator()); - at::PhiloxCudaState philox_args = init_philox_state(gen, rng_elts_per_thread); - auto options = torch::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); - auto rng_state = torch::empty({2}, options); - philox_unpack(philox_args, static_cast(rng_state.data_ptr())); - auto te_rng_state = makeTransformerEngineTensor(rng_state); - - // create auxiliary output tensors - NVTETensorPack nvte_aux_tensor_pack; - nvte_tensor_pack_create(&nvte_aux_tensor_pack); - - // create workspace - TensorWrapper workspace; - - // populate tensors with appropriate shapes and dtypes - NVTE_SCOPED_GIL_RELEASE({ - nvte_fused_attn_fwd( - te_Q.data(), te_K.data(), te_V.data(), te_Bias.data(), te_SoftmaxOffset.data(), te_S.data(), - te_O.data(), &nvte_aux_tensor_pack, te_cu_seqlens_q.data(), te_cu_seqlens_kv.data(), - te_cu_seqlens_q_padded.data(), te_cu_seqlens_kv_padded.data(), te_page_table_k.data(), - te_page_table_v.data(), te_rng_state.data(), max_seqlen_q, max_seqlen_kv, is_training, - return_max_logit, cuda_graph, attn_scale, p_dropout, qkv_layout, o_format, - qkv_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], - window_size[1], bottom_right_diagonal, workspace.data(), at::cuda::getCurrentCUDAStream()); - }); - - // allocate memory for workspace and auxiliary output tensors - auto workspace_data = allocateSpace(workspace.shape(), workspace.dtype()); - workspace = - makeTransformerEngineTensor(workspace_data.data_ptr(), workspace.shape(), workspace.dtype()); - - // output_tensors = [O, nvte_aux_tensor_pack.tensors] - std::vector output_tensors; - output_tensors.push_back(py_O); - auto set_tensor_param = [&](size_t i, const at::Tensor &output_tensor) { - output_tensors.push_back(py::cast(output_tensor)); - NVTEBasicTensor temp_data = {output_tensor.data_ptr(), - nvte_tensor_type(nvte_aux_tensor_pack.tensors[i]), - nvte_tensor_shape(nvte_aux_tensor_pack.tensors[i])}; - nvte_set_tensor_param(&nvte_aux_tensor_pack.tensors[i], kNVTERowwiseData, &temp_data); - }; - // allocate memory for nvte_aux_tensor_pack.tensors - // f16_arbitrary: S [b, h, sq, 1]/[tq, h, 1], (optional) Max [b, h, sq, 1]/[tq, h, 1], rng_state [2], (optional) Bias [1, h, sq, skv], (optional) SoftmaxOffset [1, h, 1, 1] - // fp8 : S [b, h, sq, 1], rng_state [2] - size_t i = 0; - at::Tensor output_tensor; - // intermediate softmax stats tensor S - output_tensor = - allocateSpace(nvte_shape_to_vector(nvte_tensor_shape(nvte_aux_tensor_pack.tensors[i])), - static_cast(nvte_tensor_type(nvte_aux_tensor_pack.tensors[i])), false); - set_tensor_param(i++, output_tensor); - // return_max_logit=true allocates Max after S - if (return_max_logit) { - output_tensor = - allocateSpace(nvte_shape_to_vector(nvte_tensor_shape(nvte_aux_tensor_pack.tensors[i])), - static_cast(nvte_tensor_type(nvte_aux_tensor_pack.tensors[i])), false); - set_tensor_param(i++, output_tensor); - } - // rng_state - if (i < nvte_aux_tensor_pack.size) { - set_tensor_param(i++, rng_state); - } - // bias (optional) - if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI) && (Bias.has_value())) { - set_tensor_param(i++, Bias.value()); - } - // softmax_offset (optional) - if ((softmax_type != NVTE_VANILLA_SOFTMAX) && (SoftmaxOffset.has_value())) { - set_tensor_param(i++, SoftmaxOffset.value()); - } - - // execute the kernel - NVTE_SCOPED_GIL_RELEASE({ - nvte_fused_attn_fwd( - te_Q.data(), te_K.data(), te_V.data(), te_Bias.data(), te_SoftmaxOffset.data(), te_S.data(), - te_O.data(), &nvte_aux_tensor_pack, te_cu_seqlens_q.data(), te_cu_seqlens_kv.data(), - te_cu_seqlens_q_padded.data(), te_cu_seqlens_kv_padded.data(), te_page_table_k.data(), - te_page_table_v.data(), te_rng_state.data(), max_seqlen_q, max_seqlen_kv, is_training, - return_max_logit, cuda_graph, attn_scale, p_dropout, qkv_layout, o_format, - qkv_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], - window_size[1], bottom_right_diagonal, workspace.data(), at::cuda::getCurrentCUDAStream()); - }); - - // destroy tensor wrappers, but not allocated memory - nvte_tensor_pack_destroy(&nvte_aux_tensor_pack); - - // if training, [O, softmax-related tensors, rng_state]; if inference, [O] - return output_tensors; -} - -// fused attention BWD with separate Q, K and V -std::vector fused_attn_bwd( - size_t max_seqlen_q, size_t max_seqlen_kv, float attn_scale, float p_dropout, bool set_zero, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, - NVTE_QKV_Layout dqkv_layout, NVTE_QKV_Format qkv_scale_inv_format, - NVTE_QKV_Format do_scale_inv_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, const std::vector window_size, - bool bottom_right_diagonal, bool deterministic, const at::Tensor cu_seqlens_q, - const at::Tensor cu_seqlens_kv, const py::handle Q, const py::handle K, const py::handle V, - const py::handle O, const py::handle dO, const at::ScalarType fake_dtype, - const std::vector Aux_CTX_Tensors, - const std::optional cu_seqlens_q_padded, - const std::optional cu_seqlens_kv_padded, py::handle s_quantizer, - py::handle dp_quantizer, py::handle dqkv_quantizer, bool cuda_graph) { - auto none = py::none(); - - // create QKV, O, dO tensor wrappers - TensorWrapper te_Q, te_K, te_V, te_O, te_dO; - te_Q = makeTransformerEngineTensor(Q, none); - te_K = makeTransformerEngineTensor(K, none); - te_V = makeTransformerEngineTensor(V, none); - te_O = makeTransformerEngineTensor(O, none); - te_dO = makeTransformerEngineTensor(dO, none); - - // create S and dP tensors - auto [te_S, py_S, _s] = quantizer_helper(s_quantizer, {0}, DType::kFloat32, false, std::nullopt); - auto [te_dP, py_dP, _dp] = - quantizer_helper(dp_quantizer, {0}, DType::kFloat32, false, std::nullopt); - - // create dQ, dK, dV tensors - TensorWrapper te_dQ, te_dK, te_dV; - py::object py_dQ, py_dK, py_dV; - std::optional dq_amax_buf, dk_amax_buf, dv_amax_buf; - std::unique_ptr dQKV_quantizer = convert_quantizer(dqkv_quantizer); - std::vector q_shape = convertShape(te_Q.shape()); - std::vector k_shape = convertShape(te_K.shape()); - std::vector v_shape = convertShape(te_V.shape()); - const DType dqkv_fake_dtype = GetTransformerEngineDType(fake_dtype); - size_t ndim_q = q_shape.size(); - size_t ndim_kv = k_shape.size(); - std::vector dQ_shape(ndim_q), dK_shape(ndim_kv), dV_shape(ndim_kv); - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - NVTE_QKV_Format dq_format = nvte_get_q_format(dqkv_layout); - NVTE_QKV_Format dkv_format = nvte_get_kv_format(dqkv_layout); - AttentionShape q_parsed(q_format, q_shape.data()); - size_t h_q = q_parsed.h(), d_qk = q_parsed.d(); - q_parsed.to_format(dq_format, dQ_shape.data()); - AttentionShape k_parsed(kv_format, k_shape.data()); - size_t h_kv = k_parsed.h(); - k_parsed.to_format(dkv_format, dK_shape.data()); - AttentionShape v_parsed(kv_format, v_shape.data()); - size_t d_v = v_parsed.d(); - v_parsed.to_format(dkv_format, dV_shape.data()); - at::Tensor dQ, dK, dV, dQKV, dKV; - // FP16/BF16: dqkv_fake_dtype = kFloat16/kBFloat16, dQ/dK/dV.dtype = torch.float16/torch.bfloat16 - // FP8DS: dqkv_fake_dtype = kFloat16/kBFloat16, dQ/dK/dV.dtype = torch.uint8 - // FP8CS/MXFP8: dqkv_fake_dtype = kFloat16/kBFloat16, dQ/dK/dV.dtype = torch.float16/torch.bfloat16 - auto options = torch::TensorOptions().dtype(fake_dtype).device(torch::kCUDA); - if (detail::IsFloat8Quantizers(dqkv_quantizer.ptr())) { - options = options.dtype(torch::kUInt8); - } - - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(dqkv_layout); - std::vector tmp_shape; - switch (layout_group) { - case NVTE_QKV_Layout_Group::NVTE_3HD: - tmp_shape = std::vector{dQ_shape.begin(), dQ_shape.end()}; - tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 2, int64_t(3)); - dQKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dQ = dQKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 3); - dK = dQKV.index({"...", torch::indexing::Slice(1, 2, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 3); - dV = dQKV.index({"...", torch::indexing::Slice(2, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 3); - break; - case NVTE_QKV_Layout_Group::NVTE_H3D: - tmp_shape = std::vector{dQ_shape.begin(), dQ_shape.end()}; - tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 1, int64_t(3)); - dQKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dQ = dQKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 2); - dK = dQKV.index({"...", torch::indexing::Slice(1, 2, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 2); - dV = dQKV.index({"...", torch::indexing::Slice(2, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 2); - break; - case NVTE_QKV_Layout_Group::NVTE_HD_2HD: - tmp_shape = std::vector(dQ_shape.begin(), dQ_shape.end()); - dQ = torch::empty(tmp_shape, options); - tmp_shape = std::vector{dK_shape.begin(), dK_shape.end()}; - tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 2, int64_t(2)); - dKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dK = dKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 3); - dV = dKV.index({"...", torch::indexing::Slice(1, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 3); - break; - case NVTE_QKV_Layout_Group::NVTE_HD_H2D: - tmp_shape = std::vector(dQ_shape.begin(), dQ_shape.end()); - dQ = torch::empty(tmp_shape, options); - tmp_shape = std::vector{dK_shape.begin(), dK_shape.end()}; - tmp_shape.insert(tmp_shape.begin() + tmp_shape.size() - 1, int64_t(2)); - dKV = torch::empty(c10::IntArrayRef(tmp_shape), options); - dK = dKV.index({"...", torch::indexing::Slice(0, 1, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 2); - dV = dKV.index({"...", torch::indexing::Slice(1, torch::indexing::None, 1), - torch::indexing::Slice(0, torch::indexing::None, 1)}) - .squeeze(tmp_shape.size() - 2); - break; - case NVTE_QKV_Layout_Group::NVTE_HD_HD_HD: - case NVTE_QKV_Layout_Group::NVTE_SD_SD_SD: - tmp_shape = std::vector(dQ_shape.begin(), dQ_shape.end()); - dQ = torch::empty(tmp_shape, options); - tmp_shape = std::vector(dK_shape.begin(), dK_shape.end()); - dK = torch::empty(tmp_shape, options); - tmp_shape = std::vector(dV_shape.begin(), dV_shape.end()); - dV = torch::empty(tmp_shape, options); - break; - default: - NVTE_ERROR("QKV layout not supported!"); - } - - std::tie(te_dQ, py_dQ, dq_amax_buf) = - quantizer_helper(dqkv_quantizer, dQ_shape, dqkv_fake_dtype, true, dQ); - std::tie(te_dK, py_dK, dk_amax_buf) = - quantizer_helper(dqkv_quantizer, dK_shape, dqkv_fake_dtype, true, dK); - std::tie(te_dV, py_dV, dv_amax_buf) = - quantizer_helper(dqkv_quantizer, dV_shape, dqkv_fake_dtype, true, dV); - - // construct NVTE tensors - if (detail::IsFloat8Quantizers(dqkv_quantizer.ptr())) { - // FP8 - if (set_zero) { - if (dq_format == NVTE_QKV_Format::NVTE_THD) { - if (((h_q * d_qk) % block_size == 0) && dQ.is_contiguous()) { - mha_fill(te_dQ, cu_seqlens_q.index({torch::indexing::Slice(-1, torch::indexing::None)})); - } else { - dQ.fill_(0); - } - } - if (dkv_format == NVTE_QKV_Format::NVTE_THD) { - if (((h_kv * d_qk) % block_size == 0) && ((h_kv * d_v) % block_size == 0) && - dK.is_contiguous() && dV.is_contiguous()) { - mha_fill(te_dK, cu_seqlens_kv.index({torch::indexing::Slice(-1, torch::indexing::None)})); - mha_fill(te_dV, cu_seqlens_kv.index({torch::indexing::Slice(-1, torch::indexing::None)})); - } else { - dK.fill_(0); - dV.fill_(0); - } - } - } - } else if (dqkv_quantizer.is_none() || - detail::IsFloat8CurrentScalingQuantizers(dqkv_quantizer.ptr()) || - detail::IsMXFP8Quantizers(dqkv_quantizer.ptr())) { - if (dq_format == NVTE_QKV_Format::NVTE_THD) { - dQ.fill_(0); - } - if (dkv_format == NVTE_QKV_Format::NVTE_THD) { - dK.fill_(0); - dV.fill_(0); - } - } else { - NVTE_ERROR("Fused attention only supports FP8 and BF16/FP16 data types. \n"); - } - - // create cu_seqlens tensorwrappers - auto cu_seqlens_q_sizes = cu_seqlens_q.sizes().vec(); - std::vector cu_seqlens_q_shape{cu_seqlens_q_sizes.begin(), cu_seqlens_q_sizes.end()}; - auto cu_seqlens_kv_sizes = cu_seqlens_kv.sizes().vec(); - std::vector cu_seqlens_kv_shape{cu_seqlens_kv_sizes.begin(), cu_seqlens_kv_sizes.end()}; - TensorWrapper te_cu_seqlens_q, te_cu_seqlens_kv; - te_cu_seqlens_q = makeTransformerEngineTensor(cu_seqlens_q.data_ptr(), cu_seqlens_q_shape, - DType::kInt32, nullptr, nullptr, nullptr); - te_cu_seqlens_kv = makeTransformerEngineTensor(cu_seqlens_kv.data_ptr(), cu_seqlens_kv_shape, - DType::kInt32, nullptr, nullptr, nullptr); - - TensorWrapper te_cu_seqlens_q_padded, te_cu_seqlens_kv_padded; - if ((cu_seqlens_q_padded.has_value()) && (cu_seqlens_kv_padded.has_value())) { - auto cu_seqlens_q_padded_sizes = cu_seqlens_q_padded.value().sizes().vec(); - std::vector cu_seqlens_q_padded_shape{cu_seqlens_q_padded_sizes.begin(), - cu_seqlens_q_padded_sizes.end()}; - auto cu_seqlens_kv_padded_sizes = cu_seqlens_kv_padded.value().sizes().vec(); - std::vector cu_seqlens_kv_padded_shape{cu_seqlens_kv_padded_sizes.begin(), - cu_seqlens_kv_padded_sizes.end()}; - te_cu_seqlens_q_padded = makeTransformerEngineTensor(cu_seqlens_q_padded.value().data_ptr(), - cu_seqlens_q_padded_shape, DType::kInt32); - te_cu_seqlens_kv_padded = makeTransformerEngineTensor( - cu_seqlens_kv_padded.value().data_ptr(), cu_seqlens_kv_padded_shape, DType::kInt32); - } - - // convert auxiliary tensors from forward to NVTETensors - NVTETensorPack nvte_aux_tensor_pack; - nvte_tensor_pack_create(&nvte_aux_tensor_pack); - nvte_aux_tensor_pack.size = Aux_CTX_Tensors.size(); - for (size_t i = 0; i < nvte_aux_tensor_pack.size; ++i) { - const std::vector &signed_shape = Aux_CTX_Tensors[i].sizes().vec(); - const std::vector tmp(signed_shape.begin(), signed_shape.end()); - - NVTEBasicTensor temp_data = { - Aux_CTX_Tensors[i].data_ptr(), - static_cast(GetTransformerEngineDType(Aux_CTX_Tensors[i].scalar_type())), - nvte_make_shape(tmp.data(), tmp.size())}; - nvte_set_tensor_param(&nvte_aux_tensor_pack.tensors[i], kNVTERowwiseData, &temp_data); - } - - // create dBias the same shape as Bias - at::Tensor dBias; - TensorWrapper te_dBias; - if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { - if (nvte_aux_tensor_pack.size >= 2) { - std::vector bias_shape(Aux_CTX_Tensors[nvte_aux_tensor_pack.size - 1].sizes().vec()); - dBias = torch::empty(bias_shape, options); - te_dBias = makeTransformerEngineTensor(dBias); - } else { - dBias = torch::empty({1, static_cast(h_q), static_cast(max_seqlen_q), - static_cast(max_seqlen_kv)}, - options); - te_dBias = makeTransformerEngineTensor(dBias); - } - if (nvte_get_qkv_format(qkv_layout) == NVTE_QKV_Format::NVTE_THD) { - dBias.fill_(0); - } - } - - // create dSoftmaxOffset in the same shape as SoftmaxOffset - at::Tensor dSoftmaxOffset; - TensorWrapper te_dSoftmaxOffset; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - options = torch::TensorOptions().dtype(at::kFloat).device(torch::kCUDA); - dSoftmaxOffset = torch::empty({1, static_cast(h_q), 1, 1}, options); - te_dSoftmaxOffset = makeTransformerEngineTensor(dSoftmaxOffset); - } - - // create workspace - TensorWrapper workspace; - - // populate tensors with appropriate shapes and dtypes - NVTE_SCOPED_GIL_RELEASE({ - nvte_fused_attn_bwd( - te_Q.data(), te_K.data(), te_V.data(), te_O.data(), te_dO.data(), te_S.data(), te_dP.data(), - &nvte_aux_tensor_pack, te_dQ.data(), te_dK.data(), te_dV.data(), te_dBias.data(), - te_dSoftmaxOffset.data(), te_cu_seqlens_q.data(), te_cu_seqlens_kv.data(), - te_cu_seqlens_q_padded.data(), te_cu_seqlens_kv_padded.data(), max_seqlen_q, max_seqlen_kv, - attn_scale, p_dropout, qkv_layout, o_format, do_format, dqkv_layout, qkv_scale_inv_format, - do_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], - window_size[1], bottom_right_diagonal, deterministic, cuda_graph, workspace.data(), - at::cuda::getCurrentCUDAStream()); - }); - - // allocate memory for workspace - auto workspace_data = allocateSpace(workspace.shape(), workspace.dtype()); - workspace = - makeTransformerEngineTensor(workspace_data.data_ptr(), workspace.shape(), workspace.dtype()); - - // execute kernel - NVTE_SCOPED_GIL_RELEASE({ - nvte_fused_attn_bwd( - te_Q.data(), te_K.data(), te_V.data(), te_O.data(), te_dO.data(), te_S.data(), te_dP.data(), - &nvte_aux_tensor_pack, te_dQ.data(), te_dK.data(), te_dV.data(), te_dBias.data(), - te_dSoftmaxOffset.data(), te_cu_seqlens_q.data(), te_cu_seqlens_kv.data(), - te_cu_seqlens_q_padded.data(), te_cu_seqlens_kv_padded.data(), max_seqlen_q, max_seqlen_kv, - attn_scale, p_dropout, qkv_layout, o_format, do_format, dqkv_layout, qkv_scale_inv_format, - do_scale_inv_format, bias_type, attn_mask_type, softmax_type, window_size[0], - window_size[1], bottom_right_diagonal, deterministic, cuda_graph, workspace.data(), - at::cuda::getCurrentCUDAStream()); - }); - - // destroy tensor wrappers - nvte_tensor_pack_destroy(&nvte_aux_tensor_pack); - - return {py_dQ, py_dK, py_dV, py::cast(dBias), py::cast(dSoftmaxOffset)}; -} - at::Tensor fa_prepare_fwd(at::Tensor qkvi) { NVTE_CHECK(qkvi.dim() == 4, "Expected 4-dim tensor."); NVTE_CHECK(qkvi.scalar_type() == at::ScalarType::Half || diff --git a/transformer_engine/pytorch/csrc/extensions/attention_rng.cpp b/transformer_engine/pytorch/csrc/extensions/attention_rng.cpp new file mode 100644 index 00000000000..ad3410a63a7 --- /dev/null +++ b/transformer_engine/pytorch/csrc/extensions/attention_rng.cpp @@ -0,0 +1,26 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +#include "../extensions.h" +#include "common.h" + +namespace transformer_engine::pytorch { + +at::Tensor get_cudnn_attention_rng_state(const std::optional rng_gen, + size_t increment) { + auto gen = at::get_generator_or_default( + rng_gen, at::cuda::detail::getDefaultCUDAGenerator()); + at::PhiloxCudaState philox_args = init_philox_state(gen, increment); + auto options = torch::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA); + auto rng_state = torch::empty({2}, options); + // PhiloxCudaState contains device pointers while CUDA graph capture is + // active. Unpacking on the current stream therefore preserves PyTorch's + // graph-safe intragraph offset semantics. + philox_unpack(philox_args, static_cast(rng_state.data_ptr())); + return rng_state; +} + +} // namespace transformer_engine::pytorch diff --git a/transformer_engine/pytorch/csrc/extensions/pybind.cpp b/transformer_engine/pytorch/csrc/extensions/pybind.cpp index 974b1c3b4c1..c8843df7492 100644 --- a/transformer_engine/pytorch/csrc/extensions/pybind.cpp +++ b/transformer_engine/pytorch/csrc/extensions/pybind.cpp @@ -458,8 +458,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("swap_first_dims", &transformer_engine::pytorch::swap_first_dims, "Swap first two tensor dimensions", py::arg("tensor"), py::kw_only(), py::arg("out"), py::call_guard()); - m.def("get_fused_attn_backend", &transformer_engine::pytorch::get_fused_attn_backend, - "Get Fused Attention backend", py::call_guard()); m.def("compute_amax", &transformer_engine::pytorch::compute_amax, "Compute absolute max value in tensor", py::arg("input"), py::arg("amax"), py::call_guard()); @@ -537,6 +535,11 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::call_guard()); // attention kernels + m.def("get_cudnn_attention_rng_state", + &transformer_engine::pytorch::get_cudnn_attention_rng_state, + "Reserve and unpack a graph-safe Philox state for Python cuDNN attention", + py::arg("rng_gen") = py::none(), py::arg("increment"), + py::call_guard()); m.def("fa_prepare_fwd", &transformer_engine::pytorch::fa_prepare_fwd, "Prepare QKV for Flash Attention", py::call_guard()); m.def("fa_prepare_bwd", &transformer_engine::pytorch::fa_prepare_bwd, @@ -550,10 +553,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("multi_tensor_pad_last_dim", &transformer_engine::pytorch::multi_tensor_pad_last_dim, "Pad multiple tensors' last dimension to a common alignment.", py::arg("inputs"), py::arg("alignment"), py::call_guard()); - m.def("fused_attn_fwd", &transformer_engine::pytorch::fused_attn_fwd, - "Fused Attention FP8/BF16/FP16 FWD with separate Q, K and V"); - m.def("fused_attn_bwd", &transformer_engine::pytorch::fused_attn_bwd, - "Fused Attention FP8/BF16/FP16 BWD with separate Q, K and V"); m.def("copy_to_kv_cache", &transformer_engine::pytorch::copy_to_kv_cache, "Copy new KV tokens to KV cache", py::call_guard()); m.def("convert_thd_to_bshd", &transformer_engine::pytorch::convert_thd_to_bshd, From 35d71587e47b0695c4fbc4bf21ffccf17e5ae310 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 26 Aug 2026 22:29:32 +0000 Subject: [PATCH 04/36] Fix JAX ragged cuDNN attention graphs Restore cuDNN-supported ragged graph bucketing and provide external element offsets from JAX for forward and backward execution. Use XLA FFI buffer byte sizes directly so uint32 RNG buffers do not require a TE dtype conversion. Signed-off-by: Vladimir Cherepanov --- .../jax/cpp_extensions/attention.py | 144 +++++++- .../jax/cpp_extensions/cudnn_attention.py | 327 ++++++------------ .../jax/csrc/extensions/attention.cpp | 3 +- 3 files changed, 238 insertions(+), 236 deletions(-) diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index bb9c82500ae..e2cb5eef463 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -40,7 +40,12 @@ with_sharding_constraint_by_logical_axes, ) from .base import BasePrimitive, register_primitive -from .cudnn_attention import build_bwd_graph, build_fwd_graph, is_fused_attn_supported +from .cudnn_attention import ( + build_bwd_graph, + build_fwd_graph, + is_fused_attn_supported, + ragged_graph_batch_size, +) from .misc import ( check_valid_batch_dims, get_all_device_compute_capability, @@ -254,6 +259,77 @@ def check_seed(self, seed, dropout_probability, is_training): return seed +def _multiply_offsets_as_uint64_words(offsets, multiplier): + """Return unsigned 64-bit products as interleaved low/high uint32 words.""" + values = offsets.astype(jnp.uint32) + scale = jnp.asarray(multiplier, dtype=jnp.uint32) + mask = jnp.asarray(0xFFFF, dtype=jnp.uint32) + + value_lo = values & mask + value_hi = values >> 16 + scale_lo = scale & mask + scale_hi = scale >> 16 + + product_lo = value_lo * scale_lo + product_mid_lo = value_lo * scale_hi + product_mid_hi = value_hi * scale_lo + product_hi = value_hi * scale_hi + carry = (product_lo >> 16) + (product_mid_lo & mask) + (product_mid_hi & mask) + + low_word = (product_lo & mask) | ((carry & mask) << 16) + high_word = ( + product_hi + (product_mid_lo >> 16) + (product_mid_hi >> 16) + (carry >> 16) + ) + return jnp.stack((low_word, high_word), axis=-1) + + +def _pack_ragged_offsets( + q_seq_offsets, + k_seq_offsets, + qkv_layout, + attn_heads, + num_gqa_groups, + q_head_dim, + v_head_dim, +): + """Pack Q/K/V/O/Stats element offsets into one JAX buffer.""" + q_multiplier = attn_heads * q_head_dim + if qkv_layout.is_qkvpacked(): + q_multiplier *= 3 + k_multiplier = v_multiplier = q_multiplier + elif qkv_layout.is_kvpacked(): + k_multiplier = v_multiplier = 2 * num_gqa_groups * q_head_dim + else: + k_multiplier = num_gqa_groups * q_head_dim + v_multiplier = num_gqa_groups * v_head_dim + offsets_and_multipliers = ( + (q_seq_offsets, q_multiplier), + (k_seq_offsets, k_multiplier), + (k_seq_offsets, v_multiplier), + (q_seq_offsets, attn_heads * v_head_dim), + (q_seq_offsets, attn_heads), + ) + if get_cudnn_version() < (9, 5, 0): + return jnp.stack( + tuple( + offsets.astype(jnp.int32) * jnp.asarray(multiplier, dtype=jnp.int32) + for offsets, multiplier in offsets_and_multipliers + ) + ) + return jnp.stack( + tuple( + _multiply_offsets_as_uint64_words(offsets, multiplier) + for offsets, multiplier in offsets_and_multipliers + ) + ) + + +def _pad_ragged_metadata(values, size, fill_value): + """Pad compact sequence metadata to the bucketed cuDNN graph extent.""" + values = values.flatten()[:size] + return jnp.pad(values, (0, size - values.size), constant_values=fill_value) + + class FusedAttnFwdPrimitive(BasePrimitive): """ Fused Attention Forward Primitive @@ -486,9 +562,15 @@ def convert_to_2d(offsets, batch, max_seqlen): ) return offsets_2d - batch, q_max_seqlen, kv_max_seqlen, *_ = FusedAttnHelper.parse_qkv_aval( - q, k, v, config.qkv_layout - ) + ( + batch, + q_max_seqlen, + kv_max_seqlen, + attn_heads, + num_gqa_groups, + q_head_dim, + v_head_dim, + ) = FusedAttnHelper.parse_qkv_aval(q, k, v, config.qkv_layout) assert len(batch) == 1, f"Expected len(batch) == 1, but got {len(batch)=}" kv_batch = q_batch = batch[0] @@ -521,6 +603,29 @@ def convert_to_2d(offsets, batch, max_seqlen): k_seq_offsets, k_seq_offsets >= 0, fill_value=kv_batch * kv_max_seqlen ) + graph_batch = ragged_graph_batch_size(q_batch, config.max_segments_per_seq) + q_seqlen = _pad_ragged_metadata(q_seqlen, graph_batch, fill_value) + kv_seqlen = _pad_ragged_metadata(kv_seqlen, graph_batch, fill_value) + q_seq_offsets = _pad_ragged_metadata( + q_seq_offsets, graph_batch + 1, q_batch * q_max_seqlen + ) + k_seq_offsets = _pad_ragged_metadata( + k_seq_offsets, graph_batch + 1, kv_batch * kv_max_seqlen + ) + + # Supported cuDNN ragged graphs require external element offsets. JAX disables + # x64 by default, so represent INT64 offsets as pairs of uint32 words and pack + # all five graph offsets into one otherwise-unused inner operand. + _q_segment_ids = _pack_ragged_offsets( + q_seq_offsets, + k_seq_offsets, + config.qkv_layout, + attn_heads, + num_gqa_groups, + q_head_dim, + v_head_dim, + ) + output, softmax_aux, max_tensor, rng_state, _ = FusedAttnFwdPrimitive.inner_primitive.bind( q, k, @@ -1020,9 +1125,15 @@ def convert_to_2d(offsets, batch, max_seqlen): ) return offsets_2d - batch, q_max_seqlen, kv_max_seqlen, *_ = FusedAttnHelper.parse_qkv_aval( - q, k, v, config.qkv_layout - ) + ( + batch, + q_max_seqlen, + kv_max_seqlen, + attn_heads, + num_gqa_groups, + q_head_dim, + v_head_dim, + ) = FusedAttnHelper.parse_qkv_aval(q, k, v, config.qkv_layout) assert ( len(batch) == 1 ), f"Expected len(batch) == 1, but got len(batch)={len(batch)}, batch={batch}" @@ -1056,6 +1167,25 @@ def convert_to_2d(offsets, batch, max_seqlen): k_seq_offsets, k_seq_offsets >= 0, fill_value=kv_batch * kv_max_seqlen ) + graph_batch = ragged_graph_batch_size(q_batch, config.max_segments_per_seq) + q_seqlen = _pad_ragged_metadata(q_seqlen, graph_batch, fill_value) + kv_seqlen = _pad_ragged_metadata(kv_seqlen, graph_batch, fill_value) + q_seq_offsets = _pad_ragged_metadata( + q_seq_offsets, graph_batch + 1, q_batch * q_max_seqlen + ) + k_seq_offsets = _pad_ragged_metadata( + k_seq_offsets, graph_batch + 1, kv_batch * kv_max_seqlen + ) + _q_segment_ids = _pack_ragged_offsets( + q_seq_offsets, + k_seq_offsets, + config.qkv_layout, + attn_heads, + num_gqa_groups, + q_head_dim, + v_head_dim, + ) + dq, dk, dv, dbias, dsoftmax_offset, _ = FusedAttnBwdPrimitive.inner_primitive.bind( q, k, diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index a9a32c89284..9e4f99407c9 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -51,12 +51,6 @@ _UID_DBIAS = 22 _UID_DSINK = 23 _UID_ATTN_SCALE = 24 -_UID_OFFSET_MULT_Q = 30 -_UID_OFFSET_MULT_K = 31 -_UID_OFFSET_MULT_V = 32 -_UID_OFFSET_MULT_O = 33 -_UID_OFFSET_MULT_STATS = 34 -_UID_OFFSET_MULT_MAX = 35 @dataclass(frozen=True) @@ -216,6 +210,28 @@ def _device_arch() -> int: return int(capabilities[0]) if capabilities else 0 +def ragged_graph_batch_size(input_batch: int, max_segments_per_seq: int) -> int: + """Match TE common's cuDNN graph batch-size bucket for ragged attention.""" + batch = int(input_batch) * int(max_segments_per_seq) + if _device_arch() == 120: + return batch + if batch <= 32: + return 32 + if batch <= 512: + return 1 << (batch - 1).bit_length() + return ((batch + 511) // 512) * 512 + + +def _ragged_graph_token_count(tokens: int) -> int: + """Match TE common's cuDNN graph token-count bucket for ragged attention.""" + tokens = int(tokens) + if tokens <= 1024: + return 1024 + if tokens <= 32768: + return 1 << (tokens - 1).bit_length() + return ((tokens + 32767) // 32768) * 32768 + + def _graph_dimensions(info: _LayoutInfo, config): """Return logical cuDNN B/H/S dimensions and physical auxiliary shapes.""" is_ragged = config.qkv_layout.is_thd() @@ -228,10 +244,13 @@ def _graph_dimensions(info: _LayoutInfo, config): graph_sq = info.q_max_seqlen graph_skv = info.kv_max_seqlen else: - # Static token extents replace common's quantized max-token buckets. JAX already - # compiles per local shape, so exact extents avoid unnecessary graph variants. - graph_sq = info.input_batch * info.q_max_seqlen - graph_skv = info.input_batch * info.kv_max_seqlen + graph_batch = ragged_graph_batch_size( + info.input_batch, config.max_segments_per_seq + ) + graph_sq = _ragged_graph_token_count(info.input_batch * info.q_max_seqlen) + graph_skv = _ragged_graph_token_count( + info.input_batch * info.kv_max_seqlen + ) else: graph_batch = info.input_batch graph_sq = info.q_max_seqlen @@ -269,66 +288,24 @@ def _tensor(graph, cudnn, *, name, dim, stride, dtype, uid): ) -def _ragged_offset(graph, cudnn, name: str, uid: int, graph_batch: int): +def _ragged_offset(graph, cudnn, name: str, uid: int, graph_batch: int, dtype): return _tensor( graph, cudnn, name=name, dim=(graph_batch + 1, 1, 1, 1), stride=(1, 1, 1, 1), - dtype=cudnn.data_type.INT32, + dtype=dtype, uid=uid, ) -def _set_ragged( - tensor, - offset, - multiplier: int, - *, - graph=None, - cudnn=None, - multiplier_uid=None, - scalar_uids=None, - scalar_values=None, - materialize_multiplier: bool = False, -): - """Attach token-unit offsets, optionally materializing element offsets in the graph. - - The unified SDPA engine understands ``ragged_offset_multiplier`` directly. cuDNN's - dropout forward path currently selects the composite engine, so for that case an - INT64 pointwise node performs the conversion that TE common used to launch as a - separate CUDA kernel. - """ - if materialize_multiplier: - use_int64 = get_cudnn_version() >= (9, 5, 0) - offset_data_type = cudnn.data_type.INT64 if use_int64 else cudnn.data_type.INT32 - offset_numpy_type = np.int64 if use_int64 else np.int32 - scale = _scalar_tensor( - graph, - cudnn, - f"ragged_multiplier_{multiplier_uid}", - multiplier_uid, - offset_data_type, - ) - effective_offset = graph.mul( - offset, - scale, - compute_data_type=offset_data_type, - name=f"scale_ragged_offset_{multiplier_uid}", - ) - effective_offset.set_data_type(offset_data_type) - effective_offset.set_dim((offset.get_dim()[0], 1, 1, 1)).set_stride( - (1, 1, 1, 1) - ) - tensor.set_ragged_offset(effective_offset) - scalar_uids.append(multiplier_uid) - scalar_values.append(np.asarray(multiplier, dtype=offset_numpy_type).tobytes()) - return effective_offset - else: - tensor.set_ragged_offset(offset) - tensor.set_ragged_offset_multiplier(int(multiplier)) - return offset +def _ragged_offset_spec(cudnn): + """Return the cuDNN datatype and byte size for external element offsets.""" + use_int64 = get_cudnn_version() >= (9, 5, 0) + dtype = cudnn.data_type.INT64 if use_int64 else cudnn.data_type.INT32 + itemsize = np.dtype(np.int64 if use_int64 else np.int32).itemsize + return dtype, itemsize def _mask_options(cudnn, info: _LayoutInfo, config): @@ -513,62 +490,29 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap offset_q = offset_k = offset_v = offset_o = offset_stats = None if config.qkv_layout.is_thd(): - offset_q = _ragged_offset(graph, cudnn, "offset_q", _UID_OFFSET_Q, graph_batch) - offset_k = _ragged_offset(graph, cudnn, "offset_k", _UID_OFFSET_K, graph_batch) - offset_v = _ragged_offset(graph, cudnn, "offset_v", _UID_OFFSET_V, graph_batch) - offset_o = _ragged_offset(graph, cudnn, "offset_o", _UID_OFFSET_O, graph_batch) - materialize_offsets = True - _set_ragged( - q, - offset_q, - info.q_heads * info.qk_dim * (3 if config.qkv_layout.is_qkvpacked() else 1), - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_Q, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=materialize_offsets, + offset_dtype, offset_itemsize = _ragged_offset_spec(cudnn) + offset_bytes = (graph_batch + 1) * offset_itemsize + offset_q = _ragged_offset( + graph, cudnn, "offset_q", _UID_OFFSET_Q, graph_batch, offset_dtype ) - if config.qkv_layout.is_qkvpacked(): - kv_offset_operand = 8 - kv_multiplier = 3 * info.q_heads * info.qk_dim - else: - kv_offset_operand = 9 - kv_multiplier = ( - 2 * info.kv_heads * info.qk_dim - if config.qkv_layout.is_kvpacked() - else info.kv_heads * info.qk_dim - ) - _set_ragged( - k, - offset_k, - kv_multiplier, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_K, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=materialize_offsets, + offset_k = _ragged_offset( + graph, cudnn, "offset_k", _UID_OFFSET_K, graph_batch, offset_dtype + ) + offset_v = _ragged_offset( + graph, cudnn, "offset_v", _UID_OFFSET_V, graph_batch, offset_dtype ) - _set_ragged( - v, - offset_v, - kv_multiplier - if config.qkv_layout.is_kvpacked() - else info.kv_heads * info.v_dim, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_V, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=materialize_offsets, + offset_o = _ragged_offset( + graph, cudnn, "offset_o", _UID_OFFSET_O, graph_batch, offset_dtype ) + q.set_ragged_offset(offset_q) + k.set_ragged_offset(offset_k) + v.set_ragged_offset(offset_v) input_bindings.extend( ( - GraphBinding(_UID_OFFSET_Q, 8), - GraphBinding(_UID_OFFSET_K, kv_offset_operand), - GraphBinding(_UID_OFFSET_V, kv_offset_operand), - GraphBinding(_UID_OFFSET_O, 8), + GraphBinding(_UID_OFFSET_Q, 10, 0), + GraphBinding(_UID_OFFSET_K, 10, offset_bytes), + GraphBinding(_UID_OFFSET_V, 10, 2 * offset_bytes), + GraphBinding(_UID_OFFSET_O, 10, 3 * offset_bytes), ) ) @@ -616,20 +560,15 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap ) if ragged_stats: offset_stats = _ragged_offset( - graph, cudnn, "offset_stats", _UID_OFFSET_STATS, graph_batch - ) - _set_ragged( - max_tensor, - offset_stats, - info.q_heads, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_MAX, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, + graph, + cudnn, + "offset_stats", + _UID_OFFSET_STATS, + graph_batch, + dtype=offset_dtype, ) - input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 8)) + max_tensor.set_ragged_offset(offset_stats) + input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 10, 4 * offset_bytes)) max_tensor.set_output(True) kwargs["score_max"] = max_tensor output_bindings.append(GraphBinding(_UID_MAX, 2)) @@ -640,38 +579,23 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap ).set_stride(_matrix_stride(info, config.qkv_layout, "o", graph_sq, graph_skv)) output.set_data_type(io_dtype) if config.qkv_layout.is_thd(): - _set_ragged( - output, - offset_o, - info.q_heads * info.v_dim, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_O, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, - ) + output.set_ragged_offset(offset_o) stats.set_output(True).set_uid(_UID_STATS).set_data_type(cudnn.data_type.FLOAT) stats.set_dim((graph_batch, info.q_heads, graph_sq, 1)) if ragged_stats: if offset_stats is None: offset_stats = _ragged_offset( - graph, cudnn, "offset_stats", _UID_OFFSET_STATS, graph_batch + graph, + cudnn, + "offset_stats", + _UID_OFFSET_STATS, + graph_batch, + dtype=offset_dtype, ) - input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 8)) + input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 10, 4 * offset_bytes)) stats.set_stride((info.q_heads * graph_sq, 1, info.q_heads, 1)) - _set_ragged( - stats, - offset_stats, - info.q_heads, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_STATS, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, - ) + stats.set_ragged_offset(offset_stats) else: stats.set_stride((info.q_heads * graph_sq, graph_sq, 1, 1)) @@ -900,95 +824,44 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): ) if config.qkv_layout.is_thd(): - offset_q = _ragged_offset(graph, cudnn, "offset_q", _UID_OFFSET_Q, graph_batch) - offset_k = _ragged_offset(graph, cudnn, "offset_k", _UID_OFFSET_K, graph_batch) - offset_v = _ragged_offset(graph, cudnn, "offset_v", _UID_OFFSET_V, graph_batch) - offset_o = _ragged_offset(graph, cudnn, "offset_o", _UID_OFFSET_O, graph_batch) - q_mult = ( - info.q_heads * info.qk_dim * (3 if config.qkv_layout.is_qkvpacked() else 1) - ) - kv_mult = ( - 3 * info.q_heads * info.qk_dim - if config.qkv_layout.is_qkvpacked() - else 2 * info.kv_heads * info.qk_dim - if config.qkv_layout.is_kvpacked() - else info.kv_heads * info.qk_dim + offset_dtype, offset_itemsize = _ragged_offset_spec(cudnn) + offset_bytes = (graph_batch + 1) * offset_itemsize + offset_q = _ragged_offset( + graph, cudnn, "offset_q", _UID_OFFSET_Q, graph_batch, offset_dtype ) - v_mult = ( - kv_mult - if not config.qkv_layout.is_separate() - else info.kv_heads * info.v_dim + offset_k = _ragged_offset( + graph, cudnn, "offset_k", _UID_OFFSET_K, graph_batch, offset_dtype ) - effective_q_offset = _set_ragged( - q, - offset_q, - q_mult, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_Q, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, + offset_v = _ragged_offset( + graph, cudnn, "offset_v", _UID_OFFSET_V, graph_batch, offset_dtype ) - effective_k_offset = _set_ragged( - k, - offset_k, - kv_mult, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_K, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, + offset_o = _ragged_offset( + graph, cudnn, "offset_o", _UID_OFFSET_O, graph_batch, offset_dtype ) - effective_v_offset = _set_ragged( - v, - offset_v, - v_mult, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_V, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, - ) - effective_o_offset = _set_ragged( - output, - offset_o, - info.q_heads * info.v_dim, - graph=graph, - cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_O, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, - ) - doutput.set_ragged_offset(effective_o_offset) - kv_operand = 11 if config.qkv_layout.is_qkvpacked() else 12 + q.set_ragged_offset(offset_q) + k.set_ragged_offset(offset_k) + v.set_ragged_offset(offset_v) + output.set_ragged_offset(offset_o) + doutput.set_ragged_offset(offset_o) input_bindings.extend( ( - GraphBinding(_UID_OFFSET_Q, 11), - GraphBinding(_UID_OFFSET_K, kv_operand), - GraphBinding(_UID_OFFSET_V, kv_operand), - GraphBinding(_UID_OFFSET_O, 11), + GraphBinding(_UID_OFFSET_Q, 13, 0), + GraphBinding(_UID_OFFSET_K, 13, offset_bytes), + GraphBinding(_UID_OFFSET_V, 13, 2 * offset_bytes), + GraphBinding(_UID_OFFSET_O, 13, 3 * offset_bytes), ) ) if ragged_stats: offset_stats = _ragged_offset( - graph, cudnn, "offset_stats", _UID_OFFSET_STATS, graph_batch - ) - _set_ragged( - stats, - offset_stats, - info.q_heads, graph=graph, cudnn=cudnn, - multiplier_uid=_UID_OFFSET_MULT_STATS, - scalar_uids=scalar_uids, - scalar_values=scalar_values, - materialize_multiplier=True, + name="offset_stats", + uid=_UID_OFFSET_STATS, + graph_batch=graph_batch, + dtype=offset_dtype, ) - input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 11)) + stats.set_ragged_offset(offset_stats) + input_bindings.append(GraphBinding(_UID_OFFSET_STATS, 13, 4 * offset_bytes)) if _is_dropout(config): seed = io_tensor( @@ -1027,9 +900,9 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): (graph_batch, info.kv_heads, graph_skv, info.v_dim) ).set_stride(v_stride) if config.qkv_layout.is_thd(): - dq.set_ragged_offset(effective_q_offset) - dk.set_ragged_offset(effective_k_offset) - dv.set_ragged_offset(effective_v_offset) + dq.set_ragged_offset(offset_q) + dk.set_ragged_offset(offset_k) + dv.set_ragged_offset(offset_v) workspace, data, version = finalize_graph( cudnn, graph, description="fused-attention backward" diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index bf2fc542f0d..130e8a9f663 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -216,8 +216,7 @@ void AppendRemainingBuffers(Variadic_Buffer_Type args, std::vector *ptrs } size_t BufferBytes(const Buffer_Type &buffer) { - return product(buffer.dimensions()) * - typeToSize(convert_ffi_datatype_to_te_dtype(buffer.element_type())); + return buffer.size_bytes(); } void MemsetResultAsync(cudaStream_t stream, Result_Type result, int value) { From 04188f2a746b275431a0774d832b2bd9a98ef247 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 19:59:28 +0000 Subject: [PATCH 05/36] Fix F16 attention inference auxiliaries Always generate and return softmax statistics and the Philox state for F16 forward graphs so context-parallel inference can combine partition outputs. Signed-off-by: Vladimir Cherepanov --- .../dot_product_attention/cudnn_attention.py | 64 ++++++++----------- 1 file changed, 27 insertions(+), 37 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index ef2436228c2..cd7b3ecae14 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -708,7 +708,7 @@ def _build_f16_fwd_graph( ) is_padding = options.pop("is_padding") options.update( - generate_stats=is_training, + generate_stats=True, attn_scale=float(attn_scale), use_padding_mask=is_padding, use_alibi_mask=attn_bias_type == "alibi", @@ -805,33 +805,30 @@ def _build_f16_fwd_graph( output_t.set_ragged_offset(offset_o).set_ragged_offset_multiplier(o_mult) tensors["O"] = output_t - if is_training: - assert stats is not None - _, stats_dim, stats_stride = _stats_layout( - batch=graph_batch, - heads=output_dim[1], - max_seqlen_q=graph_seqlen_q, - total_tokens_q=q.shape[0] if is_ragged_q else batch * max_seqlen_q, - ragged=use_ragged_stats, + assert stats is not None + _, stats_dim, stats_stride = _stats_layout( + batch=graph_batch, + heads=output_dim[1], + max_seqlen_q=graph_seqlen_q, + total_tokens_q=q.shape[0] if is_ragged_q else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + if use_ragged_stats and offset_stats is None: + offset_stats, stats_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1 if use_legacy_offsets else output_dim[1], + name="offset_stats", + length=graph_batch + 1, + data_type=torch.int64 if use_legacy_offsets else None, ) - if use_ragged_stats and offset_stats is None: - offset_stats, stats_mult = _ragged_offset_tensor( - graph, - cu_seqlens_q_padded, - multiplier=1 if use_legacy_offsets else output_dim[1], - name="offset_stats", - length=graph_batch + 1, - data_type=torch.int64 if use_legacy_offsets else None, - ) - tensors["offset_stats"] = offset_stats - stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( - stats_dim - ).set_stride(stats_stride) - if use_ragged_stats: - stats_t.set_ragged_offset(offset_stats).set_ragged_offset_multiplier( - stats_mult - ) - tensors["Stats"] = stats_t + tensors["offset_stats"] = offset_stats + stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( + stats_dim + ).set_stride(stats_stride) + if use_ragged_stats: + stats_t.set_ragged_offset(offset_stats).set_ragged_offset_multiplier(stats_mult) + tensors["Stats"] = stats_t return GraphEntry( graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) @@ -891,11 +888,7 @@ def _f16_forward( total_tokens_q=total_tokens_q, ragged=ragged_stats, ) - stats = ( - torch.empty(stats_shape, dtype=torch.float32, device=q.device) - if is_training - else None - ) + stats = torch.empty(stats_shape, dtype=torch.float32, device=q.device) max_scores = ( torch.empty(stats_shape, dtype=torch.float32, device=q.device) if return_max_logit @@ -975,8 +968,7 @@ def _f16_forward( tensors["V"]: v, tensors["O"]: output, } - if is_training: - variant_pack[tensors["Stats"]] = stats + variant_pack[tensors["Stats"]] = stats if attn_bias_type == "post_scale_bias": variant_pack[tensors["Bias"]] = attn_bias legacy_offsets = tensors["_legacy_offsets"] @@ -1028,10 +1020,8 @@ def _f16_forward( variant_pack[tensors["Max"]] = max_scores entry.execute(variant_pack, q.device) - aux: List[torch.Tensor] = [] + aux: List[torch.Tensor] = [stats, rng_state] if is_training: - aux.append(stats) - aux.append(rng_state) if attn_bias_type not in ("no_bias", "alibi"): aux.append(attn_bias) if softmax_type != "vanilla": From 31c043d82ce3010f8bc8977be0af8c5966ecda6a Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:00:47 +0000 Subject: [PATCH 06/36] Key F16 attention graphs by logical batch Include the effective bucketed THD graph batch in forward and backward cache keys to prevent reuse across incompatible ragged sequence counts. Signed-off-by: Vladimir Cherepanov --- .../dot_product_attention/cudnn_attention.py | 26 +++++++++++++++++-- 1 file changed, 24 insertions(+), 2 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index cd7b3ecae14..67eff1f1d9a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -864,11 +864,22 @@ def _f16_forward( softmax_offset: Optional[torch.Tensor], return_max_logit: bool, ) -> Tuple[torch.Tensor, List[torch.Tensor], Optional[torch.Tensor]]: - q_format, _ = _q_kv_formats(qkv_layout) + q_format, kv_format = _q_kv_formats(qkv_layout) batch = cu_seqlens_q.numel() - 1 heads = q.shape[1] if q_format == "bhsd" else q.shape[-2] total_tokens_q = q.shape[0] if q_format == "thd" else batch * max_seqlen_q cudnn = import_cudnn_frontend() + use_token_buckets = ( + cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ) + use_direct_offsets = cudnn.backend_version() >= 92400 and dropout == 0.0 + use_legacy_offsets = ( + q_format == "thd" or kv_format == "thd" + ) and not use_direct_offsets + graph_batch = ( + _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch + ) ragged_stats = ( q_format == "thd" and cudnn.backend_version() >= 90600 @@ -906,6 +917,7 @@ def _f16_forward( is_training=is_training, max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, + graph_batch=graph_batch, q=_tensor_metadata(q), k=_tensor_metadata(k), v=_tensor_metadata(v), @@ -2668,10 +2680,10 @@ def fused_attn_bwd( softmax_offset = aux_ctx_tensors[aux_index] q_format, kv_format = _q_kv_formats(qkv_layout) + batch = cu_seqlens_q.numel() - 1 if q_format == "thd" or kv_format == "thd": d_q, d_k, d_v = _allocate_grad_views((q, k, v), fast_zero_fill=fast_zero_fill) else: - batch = cu_seqlens_q.numel() - 1 heads = q.shape[1] if q_format == "bhsd" else q.shape[-2] kv_heads = k.shape[1] if kv_format == "bhsd" else k.shape[-2] d_q, d_k, d_v = _allocate_attention_grad_data( @@ -2701,11 +2713,21 @@ def fused_attn_bwd( cu_seqlens_kv_padded = ( cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded ) + cudnn = import_cudnn_frontend() + use_token_buckets = ( + cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ) + use_legacy_offsets = q_format == "thd" or kv_format == "thd" + graph_batch = ( + _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch + ) key = ( "f16_bwd", max_seqlen_q, max_seqlen_kv, + graph_batch, _tensor_metadata(q), _tensor_metadata(k), _tensor_metadata(v), From 87ef877c7abda6852834e3ddcf85c970b3cc8f3d Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:03:38 +0000 Subject: [PATCH 07/36] Use logical graph batch for dense attention tensors Describe dense peers of ragged QKV tensors with the bucketed graph batch while preserving their physical storage strides. Signed-off-by: Vladimir Cherepanov --- .../attention/dot_product_attention/cudnn_attention.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 67eff1f1d9a..f3955d80c66 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -311,16 +311,16 @@ def _logical_bhsd_desc( stride = tuple(tensor.stride()) if tensor_format == "sbhd": return ( - (shape[1], shape[2], shape[0], shape[3]), + (batch, shape[2], shape[0], shape[3]), (stride[1], stride[2], stride[0], stride[3]), ) if tensor_format == "bshd": return ( - (shape[0], shape[2], shape[1], shape[3]), + (batch, shape[2], shape[1], shape[3]), (stride[0], stride[2], stride[1], stride[3]), ) if tensor_format == "bhsd": - return (shape, stride) + return ((batch, shape[1], shape[2], shape[3]), stride) if tensor_format == "thd": # The ragged offset selects each batch's first token. The synthetic # batch stride is only graph metadata; S/H/D use the real packed view. From fd1aa126e0ecb1a8ec95d30aa1b3ba1ec601b6d7 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:04:06 +0000 Subject: [PATCH 08/36] Skip ragged KV token hint on SM120 Avoid passing the packed KV token bound to cuDNN THD backward on SM120, where token bucketing is disabled and the graph sequence extent is per-sequence. Signed-off-by: Vladimir Cherepanov --- .../attention/dot_product_attention/cudnn_attention.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index f3955d80c66..372df3fcb5e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -1876,7 +1876,11 @@ def _build_f16_bwd_graph( ) if use_ragged_stats: options["max_total_seq_len_q"] = graph_seqlen_q - if is_ragged_kv and cudnn.backend_version() >= 90600: + if ( + is_ragged_kv + and cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ): options["max_total_seq_len_kv"] = graph_seqlen_kv if attn_bias_type == "post_scale_bias": From 6b619967f08ade72df147d57f17bc703164ba4bf Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:06:34 +0000 Subject: [PATCH 09/36] Mask THD padding before max-logit reduction Build a device-side validity mask from padded query offsets so unwritten inter-sequence and tail gaps cannot affect ragged max-logit results. Signed-off-by: Vladimir Cherepanov --- .../dot_product_attention/cudnn_attention.py | 39 ++++++++++++++----- 1 file changed, 30 insertions(+), 9 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 372df3fcb5e..bdda677b182 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -1041,15 +1041,36 @@ def _f16_forward( max_logit = None if return_max_logit: - if q_format == "thd" and max_scores.ndim == 4: - seqlens_q = _sequence_lengths(cu_seqlens_q).to(device=max_scores.device) - sq_idx = torch.arange(max_scores.shape[2], device=max_scores.device).view( - 1, 1, -1, 1 - ) - valid = sq_idx < seqlens_q.view(-1, 1, 1, 1) - max_scores_for_reduce = max_scores.masked_fill(~valid, float("-inf")) - else: - max_scores_for_reduce = max_scores + max_scores_for_reduce = max_scores + if q_format == "thd": + if max_scores.ndim == 4: + seqlens_q = _sequence_lengths(cu_seqlens_q).to( + device=max_scores.device + ) + sq_idx = torch.arange( + max_scores.shape[2], device=max_scores.device + ).view(1, 1, -1, 1) + valid = sq_idx < seqlens_q.view(-1, 1, 1, 1) + max_scores_for_reduce = max_scores.masked_fill( + ~valid, float("-inf") + ) + elif max_scores.ndim == 3: + seqlens_q = _sequence_lengths(cu_seqlens_q).to( + device=max_scores.device + ) + total_tokens = max_scores.shape[0] + starts = cu_seqlens_q_padded[:-1].to(device=max_scores.device) + ends = (starts + seqlens_q).clamp(max=total_tokens) + delta = torch.zeros( + total_tokens + 1, dtype=torch.int32, device=max_scores.device + ) + updates = torch.ones_like(starts, dtype=torch.int32) + delta.scatter_add_(0, starts.clamp(max=total_tokens), updates) + delta.scatter_add_(0, ends, -updates) + valid = delta[:-1].cumsum(0) > 0 + max_scores_for_reduce = max_scores.masked_fill( + ~valid.view(-1, 1, 1), float("-inf") + ) reduce_dims = (0, 2) if max_scores_for_reduce.ndim == 3 else (0, 2, 3) max_logit = torch.amax(max_scores_for_reduce, dim=reduce_dims).to(output.dtype) return output, aux, max_logit From 8424cc46688e53ad808849c94db037bb0d6059c9 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:07:34 +0000 Subject: [PATCH 10/36] Zero THD outputs on the slow fill path Always initialize ragged F16 outputs and gradient storage because cuDNN leaves inter-sequence and tail padding unwritten, including when fast zero fill is disabled. Signed-off-by: Vladimir Cherepanov --- .../dot_product_attention/cudnn_attention.py | 15 +++++---------- 1 file changed, 5 insertions(+), 10 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index bdda677b182..984ec5ee3b9 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -381,10 +381,8 @@ def _storage_span(tensor: torch.Tensor) -> int: def _allocate_grad_views( inputs: Sequence[torch.Tensor], - *, - fast_zero_fill: bool, ) -> Tuple[torch.Tensor, ...]: - """Allocate gradients while preserving packed-QKV storage relationships.""" + """Allocate zeroed gradients while preserving packed-QKV storage relationships.""" groups: Dict[Tuple[str, int, int], List[int]] = {} for index, tensor in enumerate(inputs): @@ -399,8 +397,7 @@ def _allocate_grad_views( out = torch.empty_strided( inp.shape, inp.stride(), dtype=inp.dtype, device=inp.device ) - if fast_zero_fill: - out.zero_() + out.zero_() outputs[indices[0]] = out continue @@ -410,11 +407,9 @@ def _allocate_grad_views( for index in indices ) exemplar = inputs[indices[0]] - base = torch.empty( + base = torch.zeros( max_end - min_offset, dtype=exemplar.dtype, device=exemplar.device ) - if fast_zero_fill: - base.zero_() for index in indices: inp = inputs[index] outputs[index] = torch.as_strided( @@ -890,7 +885,7 @@ def _f16_forward( if o_format == "thd" else _fp8_output_shape(batch, max_seqlen_q, heads, v.shape[-1], o_format) ) - output_factory = torch.zeros if fast_zero_fill else torch.empty + output_factory = torch.zeros if fast_zero_fill or o_format == "thd" else torch.empty output = output_factory(output_shape, dtype=fake_dtype, device=q.device) stats_shape, _, _ = _stats_layout( batch=batch, @@ -2707,7 +2702,7 @@ def fused_attn_bwd( q_format, kv_format = _q_kv_formats(qkv_layout) batch = cu_seqlens_q.numel() - 1 if q_format == "thd" or kv_format == "thd": - d_q, d_k, d_v = _allocate_grad_views((q, k, v), fast_zero_fill=fast_zero_fill) + d_q, d_k, d_v = _allocate_grad_views((q, k, v)) else: heads = q.shape[1] if q_format == "bhsd" else q.shape[-2] kv_heads = k.shape[1] if kv_format == "bhsd" else k.shape[-2] From 95c2b5a18d398c6cfbe2ff7bd9ae0a2320bdeb84 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:08:14 +0000 Subject: [PATCH 11/36] Key FP8 backward graphs by encoding Include the FP8 dtypes of attention inputs, dO, and gradient outputs so delayed-scaling recipes with identical uint8 buffers cannot share incompatible graphs. Signed-off-by: Vladimir Cherepanov --- .../pytorch/attention/dot_product_attention/cudnn_attention.py | 1 + 1 file changed, 1 insertion(+) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 984ec5ee3b9..d637e107c53 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -2460,6 +2460,7 @@ def _fp8_backward( _tensor_metadata(_quantized_data(x) if _is_float8_tensor(x) else x) for x in (q, k, v, o, d_o) ), + *(getattr(x, "_fp8_dtype", None) for x in (q, k, v, o, d_o, d_q, d_k, d_v)), type(q).__name__, type(dqkv_quantizer).__name__, qkv_layout, From 31be28238a697ccc5693f3845a485e15e42d66da Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:08:36 +0000 Subject: [PATCH 12/36] Key paged attention graphs by table layout Record both page-table tensor descriptors in F16 forward cache keys so differing maximum-page dimensions and strides build distinct graphs. Signed-off-by: Vladimir Cherepanov --- .../pytorch/attention/dot_product_attention/cudnn_attention.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index d637e107c53..ce87f78fa5a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -929,7 +929,8 @@ def _f16_forward( window_size=tuple(window_size), bottom_right_diagonal=bottom_right_diagonal, return_max_logit=return_max_logit, - paged=page_table_k is not None, + page_table_k=_tensor_metadata(page_table_k), + page_table_v=_tensor_metadata(page_table_v), ) entry = get_graph_entry(key) if entry is None: From 92ed239cc175eb057df6ece59cd29c49f06e35e6 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Wed, 2 Sep 2026 20:17:37 +0000 Subject: [PATCH 13/36] Fix JAX cuDNN graph build compatibility Enable NVCC compilation for standalone JAX wheels, preserve legacy THD graph dimensions before cuDNN 9.6, and pin the runtime frontend package to the build-time version. Add regression coverage for the THD bucketing boundary. Signed-off-by: Vladimir Cherepanov --- build_tools/build_ext.py | 13 +++---- build_tools/jax.py | 10 ++++- tests/jax/test_fused_attn.py | 39 ++++++++++++++++++- .../jax/cpp_extensions/cudnn_attention.py | 13 ++++--- 4 files changed, 59 insertions(+), 16 deletions(-) diff --git a/build_tools/build_ext.py b/build_tools/build_ext.py index 30ea83a55b9..4d802b66e95 100644 --- a/build_tools/build_ext.py +++ b/build_tools/build_ext.py @@ -172,8 +172,9 @@ def run(self) -> None: os.remove(ext) def build_extensions(self): - # For core lib + JAX install, fix build_ext from pybind11.setup_helpers - # to handle CUDA files correctly. + # For JAX builds, fix build_ext from pybind11.setup_helpers to handle + # CUDA files correctly. This is also required by the standalone JAX + # wheel, where framework_extension_only is true. if "pytorch" not in get_frameworks(): # Ensure at least an empty list of flags for 'cxx' and 'nvcc' when # extra_compile_args is a dict. @@ -185,8 +186,7 @@ def build_extensions(self): # Define new _compile method that redirects to NVCC for .cu and .cuh files. original_compile_fn = self.compiler._compile - if not framework_extension_only: - self.compiler.src_extensions += [".cu", ".cuh"] + self.compiler.src_extensions += [".cu", ".cuh"] def _compile_fn(obj, src, ext, cc_args, extra_postargs, pp_opts) -> None: # Copy before we make any modifications. @@ -195,10 +195,7 @@ def _compile_fn(obj, src, ext, cc_args, extra_postargs, pp_opts) -> None: try: original_compiler = self.compiler.compiler_so - if ( - os.path.splitext(src)[1] in [".cu", ".cuh"] - and not framework_extension_only - ): + if os.path.splitext(src)[1] in [".cu", ".cuh"]: nvcc_bin = nvcc_path() self.compiler.set_executable("compiler_so", str(nvcc_bin)) if isinstance(cflags, dict): diff --git a/build_tools/jax.py b/build_tools/jax.py index 44297200ef3..829a4463b6e 100644 --- a/build_tools/jax.py +++ b/build_tools/jax.py @@ -5,6 +5,7 @@ """JAX related extensions.""" import os +from importlib.metadata import version as get_package_version from pathlib import Path from typing import List @@ -23,7 +24,14 @@ def install_requirements() -> List[str]: """Install dependencies for TE/JAX extensions.""" - return ["jax", "flax>=0.7.1", "nvidia-cudnn-frontend>=1.25.0"] + # Serialized cuDNN graphs use a version-specific wire format, so the Python + # frontend used at runtime must match the headers used to build the extension. + frontend_version = get_package_version("nvidia-cudnn-frontend") + return [ + "jax", + "flax>=0.7.1", + f"nvidia-cudnn-frontend=={frontend_version}", + ] def test_requirements() -> List[str]: diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index d68e409331c..3e0b14b682e 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -40,7 +40,7 @@ CPStrategy, ReorderStrategy, ) -from transformer_engine.jax.cpp_extensions import FusedAttnHelper +from transformer_engine.jax.cpp_extensions import FusedAttnHelper, cudnn_attention from transformer_engine_jax import ( NVTE_Fused_Attn_Backend, get_cudnn_version, @@ -383,6 +383,43 @@ def score_mod(_graph, score, _tensors): ) +@pytest.mark.parametrize( + "cudnn_version, expected_dimensions", + [ + ((9, 5, 1), (8, 768, 640, False)), + ((9, 6, 0), (32, 2048, 2048, True)), + ], +) +def test_thd_graph_bucketing_requires_cudnn_9_6( + monkeypatch, cudnn_version, expected_dimensions +): + """Pre-9.6 THD graphs retain dense dimensions and metadata extents.""" + monkeypatch.setattr(cudnn_attention, "get_cudnn_version", lambda: cudnn_version) + monkeypatch.setattr(cudnn_attention, "_device_arch", lambda: 90) + info = cudnn_attention._LayoutInfo( + batch_shape=(2,), + input_batch=2, + q_max_seqlen=768, + kv_max_seqlen=640, + q_heads=8, + kv_heads=8, + qk_dim=128, + v_dim=128, + ) + class Config: + qkv_layout = QKVLayout.THD_THD_THD + max_segments_per_seq = 4 + return_max_logit = False + + dimensions = cudnn_attention._graph_dimensions(info, Config()) + + assert dimensions[:4] == expected_dimensions + if cudnn_version < (9, 6, 0): + assert dimensions[4] == (2, 8, 768, 4) + else: + assert dimensions[4] == (2, 768, 8, 1) + + class BiasShape(Enum): """ Enum class to represent the different bias shapes used in the fused attention. diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index 9e4f99407c9..e65a4fd2a40 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -213,7 +213,9 @@ def _device_arch() -> int: def ragged_graph_batch_size(input_batch: int, max_segments_per_seq: int) -> int: """Match TE common's cuDNN graph batch-size bucket for ragged attention.""" batch = int(input_batch) * int(max_segments_per_seq) - if _device_arch() == 120: + # Bucketing is part of cuDNN's ragged-stats layout, introduced in 9.6. + # Older versions use dense stats and require the physical metadata extent. + if get_cudnn_version() < (9, 6, 0) or _device_arch() == 120: return batch if batch <= 32: return 32 @@ -239,14 +241,13 @@ def _graph_dimensions(info: _LayoutInfo, config): arch = _device_arch() use_ragged_stats = is_ragged and cudnn_version >= (9, 6, 0) and arch != 120 if is_ragged: - graph_batch = info.input_batch * int(config.max_segments_per_seq) - if arch == 120: + graph_batch = ragged_graph_batch_size( + info.input_batch, config.max_segments_per_seq + ) + if cudnn_version < (9, 6, 0) or arch == 120: graph_sq = info.q_max_seqlen graph_skv = info.kv_max_seqlen else: - graph_batch = ragged_graph_batch_size( - info.input_batch, config.max_segments_per_seq - ) graph_sq = _ragged_graph_token_count(info.input_batch * info.q_max_seqlen) graph_skv = _ragged_graph_token_count( info.input_batch * info.kv_max_seqlen From bfe6b7d7fd18e47aec0404348a77c298644feb4c Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Thu, 10 Sep 2026 21:43:43 +0000 Subject: [PATCH 14/36] Remove obsolete common fused attention code PyTorch and JAX now construct and execute cuDNN attention graphs in their framework-specific Python frontends, leaving the common graph implementations unused. Remove those implementations and their public entry points while retaining shared layout and support helpers. Move the graph-safe seed and offset extraction helper into the generic common utilities because PyTorch dropout and stochastic rounding still use it. Signed-off-by: Vladimir Cherepanov --- transformer_engine/common/CMakeLists.txt | 3 - transformer_engine/common/cudnn_utils.cpp | 23 - transformer_engine/common/cudnn_utils.h | 2 - .../common/fused_attn/fused_attn.cpp | 569 +----- .../fused_attn_f16_arbitrary_seqlen.cu | 1428 -------------- .../fused_attn_f16_arbitrary_seqlen.h | 52 - .../common/fused_attn/fused_attn_fp8.cu | 1711 ----------------- .../common/fused_attn/fused_attn_fp8.h | 46 - transformer_engine/common/fused_attn/utils.cu | 619 ------ transformer_engine/common/fused_attn/utils.h | 397 ---- .../include/transformer_engine/fused_attn.h | 277 +-- .../common/include/transformer_engine/utils.h | 18 + transformer_engine/common/util/utils.cu | 26 + .../jax/cpp_extensions/cudnn_attention.py | 4 +- .../dot_product_attention/_cudnn_backend.py | 4 +- .../dot_product_attention/cudnn_attention.py | 4 +- .../pytorch/attention/fused_mla_q_uproj.py | 2 +- 17 files changed, 58 insertions(+), 5127 deletions(-) delete mode 100644 transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu delete mode 100644 transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.h delete mode 100644 transformer_engine/common/fused_attn/fused_attn_fp8.cu delete mode 100644 transformer_engine/common/fused_attn/fused_attn_fp8.h delete mode 100644 transformer_engine/common/fused_attn/utils.cu delete mode 100644 transformer_engine/common/fused_attn/utils.h diff --git a/transformer_engine/common/CMakeLists.txt b/transformer_engine/common/CMakeLists.txt index ea7e1f6b498..1452d4b49f9 100644 --- a/transformer_engine/common/CMakeLists.txt +++ b/transformer_engine/common/CMakeLists.txt @@ -217,9 +217,6 @@ list(APPEND transformer_engine_cuda_sources dropout/dropout.cu fused_attn/context_parallel.cu fused_attn/kv_cache.cu - fused_attn/fused_attn_f16_arbitrary_seqlen.cu - fused_attn/fused_attn_fp8.cu - fused_attn/utils.cu gemm/cublaslt_gemm.cu gemm/cublaslt_grouped_gemm.cu normalization/layernorm/ln_bwd_semi_cuda_kernel.cu diff --git a/transformer_engine/common/cudnn_utils.cpp b/transformer_engine/common/cudnn_utils.cpp index 05ee35ccc7f..cf864a4b075 100644 --- a/transformer_engine/common/cudnn_utils.cpp +++ b/transformer_engine/common/cudnn_utils.cpp @@ -11,29 +11,6 @@ namespace transformer_engine { -// get cuDNN data type -cudnnDataType_t get_cudnn_dtype(const transformer_engine::DType t) { - using namespace transformer_engine; - switch (t) { - case DType::kInt32: - return CUDNN_DATA_INT32; - case DType::kInt64: - return CUDNN_DATA_INT64; - case DType::kFloat16: - return CUDNN_DATA_HALF; - case DType::kFloat32: - return CUDNN_DATA_FLOAT; - case DType::kBFloat16: - return CUDNN_DATA_BFLOAT16; - case DType::kFloat8E4M3: - return CUDNN_DATA_FP8_E4M3; - case DType::kFloat8E5M2: - return CUDNN_DATA_FP8_E5M2; - default: - NVTE_ERROR("Invalid cuDNN data type. \n"); - } -} - // get cuDNN data type cudnn_frontend::DataType_t get_cudnn_fe_dtype(const transformer_engine::DType t) { using namespace transformer_engine; diff --git a/transformer_engine/common/cudnn_utils.h b/transformer_engine/common/cudnn_utils.h index 0777d1e03d0..beba7de0916 100644 --- a/transformer_engine/common/cudnn_utils.h +++ b/transformer_engine/common/cudnn_utils.h @@ -23,8 +23,6 @@ void CreateCuDNNHandle(cudnnHandle_t* handle); } // namespace detail -cudnnDataType_t get_cudnn_dtype(const transformer_engine::DType t); - cudnn_frontend::DataType_t get_cudnn_fe_dtype(const transformer_engine::DType t); using cudnnExecutionPlanManager = detail::HandleManager; diff --git a/transformer_engine/common/fused_attn/fused_attn.cpp b/transformer_engine/common/fused_attn/fused_attn.cpp index 0e773fdc84a..c3601e8af6c 100644 --- a/transformer_engine/common/fused_attn/fused_attn.cpp +++ b/transformer_engine/common/fused_attn/fused_attn.cpp @@ -4,15 +4,11 @@ * See LICENSE for license information. ************************************************************************/ +#include + #include "transformer_engine/fused_attn.h" #include "../common.h" -#include "../cudnn_utils.h" -#include "../util/cuda_runtime.h" -#include "../util/system.h" -#include "fused_attn_f16_arbitrary_seqlen.h" -#include "fused_attn_fp8.h" -#include "utils.h" namespace transformer_engine { @@ -224,564 +220,3 @@ NVTE_QKV_Format nvte_get_kv_format(NVTE_QKV_Layout qkv_layout) { " in nvte_get_kv_format."); } } - -// select a backend for fused attention -NVTE_Fused_Attn_Backend nvte_get_fused_attn_backend( - bool is_training, NVTEDType q_dtype, NVTEDType kv_dtype, NVTE_QKV_Layout qkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, - float dropout, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool cuda_graph, bool deterministic) { - using namespace transformer_engine; - NVTE_Fused_Attn_Backend backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - const int device_id = cuda::current_device(); - const int sm_arch_ = cuda::sm_arch(device_id); - NVTE_CHECK(q_dtype == kv_dtype, "Q and KV must have the same data type."); - NVTE_QKV_Format qkv_format = nvte_get_qkv_format(qkv_layout); - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - const bool is_thd_layout = - q_format == NVTE_QKV_Format::NVTE_THD || kv_format == NVTE_QKV_Format::NVTE_THD; - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); - auto cudnn_runtime_version = cudnnGetVersion(); - - // For ragged offsets we only support 32-bit prior to cuDNN 9.5 - // Only used when THD format is requested. - const bool requires_64bit_ragged_offset = - (qkv_format == NVTE_THD && fused_attn::get_ragged_offset_dtype( - layout_group, num_attn_heads, num_gqa_groups, max_seqlen_q, - max_seqlen_kv, head_dim_qk, head_dim_v) == DType::kInt64); - const bool supported_ragged_offset_size = - (!requires_64bit_ragged_offset || cudnn_runtime_version >= 90500); - - if ((q_dtype == NVTEDType::kNVTEFloat8E4M3 || q_dtype == NVTEDType::kNVTEFloat8E5M2) && - sm_arch_ >= 90 && bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && - ( - // 9.2.1: {bshd, sbhd}, any seqlen, d=128, {no_mask, causal} - (cudnn_runtime_version >= 90201 && sm_arch_ < 100 && max_seqlen_q % 128 == 0 && - max_seqlen_kv % 128 == 0 && head_dim_qk == 128 && head_dim_v == 128 && - (attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK)) || - // 9.7: {bshd, sbhd}, any seqlen, d<=256 for sm90 and d<=128 for sm100, {padding, padding_causal} - (cudnn_runtime_version >= 90700 && - // TODO (cyang): add is_training to nvte_get_fused_attn_backend - // sm90: fwd d<=256, bwd d=128 only - // sm100: fwd d<=128, bwd d<=128 - ((sm_arch_ < 100 && (!is_training) && head_dim_qk <= 256 && head_dim_v <= 256) || - (sm_arch_ < 100 && is_training && head_dim_qk == 128 && head_dim_v == 128) || - (sm_arch_ >= 100 && head_dim_qk <= 128 && head_dim_v <= 128)) && - head_dim_qk % 16 == 0 && head_dim_v % 16 == 0 && - (attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - (sm_arch_ >= 100 && - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK))) || - // 9.21: d_qk=192, d_v=128 - (cudnn_runtime_version >= 92100 && sm_arch_ >= 100 && head_dim_qk <= 192 && - head_dim_v <= 128 && head_dim_qk % 16 == 0 && head_dim_v % 16 == 0 && - (attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK))) && - // pre-9.21: {bshd, sbhd}, {vanilla} - // 9.21+: {bshd, sbhd, bhsd}, {vanilla, off-by-one, learnable} - // 9.23+: {thd}; sm90 fwd only, sm100+ fwd/bwd - ((cudnn_runtime_version < 92100 && - (qkv_format == NVTE_QKV_Format::NVTE_BSHD || qkv_format == NVTE_QKV_Format::NVTE_SBHD) && - softmax_type == NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX) || - (cudnn_runtime_version >= 92100 && - (qkv_format == NVTE_QKV_Format::NVTE_BSHD || qkv_format == NVTE_QKV_Format::NVTE_SBHD || - qkv_format == NVTE_QKV_Format::NVTE_BHSD)) || - ((cudnn_runtime_version >= 92300 && (sm_arch_ >= 100 || (sm_arch_ >= 90 && !is_training))) && - qkv_format == NVTE_QKV_Format::NVTE_THD && supported_ragged_offset_size && - // cuDNN before 9.26 misindexes ragged FP8 Stats during sink-token backward. - (!is_training || softmax_type == NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX || - cudnn_runtime_version >= 92600) && - (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK))) && - // 9.10.0: known bugs with SDPA FP8 - (cudnn_runtime_version != 91000) && !return_max_logit) { - backend = NVTE_Fused_Attn_Backend::NVTE_FP8; - } else if ((q_dtype == NVTEDType::kNVTEFloat16) || (q_dtype == NVTEDType::kNVTEBFloat16)) { - bool flag_arb = false; - if ( - // TODO(cyang): replace with cudnn-frontend check_support for cleaner logic and better error messaging - // architecture - ((cudnn_runtime_version < 8903 && (sm_arch_ == 80 || sm_arch_ == 90)) || - (cudnn_runtime_version >= 8903 && sm_arch_ >= 80 && sm_arch_ < 100) || - (cudnn_runtime_version >= 90700 && sm_arch_ >= 100)) && - // sequence length - ((cudnn_runtime_version < 90000 && max_seqlen_q % 64 == 0 && max_seqlen_kv % 64 == 0) || - (cudnn_runtime_version >= 90000)) && - // number of heads - ((cudnn_runtime_version < 8907 && num_attn_heads == num_gqa_groups) || - (cudnn_runtime_version >= 8907)) && - // head dimension - // multiples of 8 - (head_dim_qk % 8 == 0 && head_dim_v % 8 == 0 && - // <= 128 - ((head_dim_qk <= 128 && head_dim_v <= 128) || - // 9.1: <= 256 + Hopper + fprop - // 9.5: <= 256 + Hopper + bprop - (head_dim_qk <= 256 && head_dim_v <= 256 && - ((!is_training && sm_arch_ == 90 && cudnn_runtime_version >= 90100) || - (is_training && sm_arch_ == 90 && cudnn_runtime_version >= 90500))) || - // 9.9: any head_dim + Blackwell + fprop + non_paged + sq > 1 - (!is_training && sm_arch_ >= 100 && cudnn_runtime_version >= 90900 && max_seqlen_q > 1 && - layout_group != NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD) || - // 9.10.2: any head_dim + any arch + fprop + paged - // 9.10.2: any head_dim + any arch + fprop + non_paged + sq > 1 - // 9.10.2: any head_dim + any arch + fprop + non_paged + sq = 1 + {no_mask, padding, BRCM, padding_BRCM} - (!is_training && cudnn_runtime_version >= 91002 && - (layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD || max_seqlen_q > 1 || - (max_seqlen_q == 1 && attn_mask_type != NVTE_Mask_Type::NVTE_CAUSAL_MASK && - attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK))) || - // 9.11: d_qk = 192, d_v = 128 + Blackwell + bprop + non-paged - (head_dim_qk == 192 && head_dim_v == 128 && is_training && sm_arch_ >= 100 && - cudnn_runtime_version >= 91100) || - // 9.23: d_qk = d_v = 256 + SM10x (cuDNN FE 1.24 / BE 9.23+) + bprop + non-paged. - // THD layouts require cuDNN FE 1.26 / BE 9.25+ for execution-plan support. - (head_dim_qk == 256 && head_dim_v == 256 && is_training && sm_arch_ >= 100 && - sm_arch_ < 110 && cudnn_runtime_version >= (is_thd_layout ? 92500 : 92300) && - layout_group != NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD && - // The FE forces this path onto the deterministic bprop algorithm, which on - // Blackwell rejects dBias, dropout, and ALiBi (and supports vanilla softmax only). - bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && dropout == 0.0 && - softmax_type == NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX && - // Non-causal D=256 supports only full-window attention; SWA is allowed only for causal masks. - ((window_size_left == -1 && window_size_right == -1) || - ((attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK) && - (window_size_right == -1 || window_size_right == 0))))) && - // 9.11+ bug: 128 < d_qk <= 256, 128 < d_v <= 256 + Hopper + bprop + MLA - // Conditional to temporarily use blanket cudnn_runtime_version >= 9.11 until fixed - (!((cudnn_runtime_version >= 91100) && is_training && sm_arch_ == 90 && - head_dim_qk >= 128 && head_dim_v >= 128 && !(head_dim_qk == 192 && head_dim_v == 128) && - head_dim_qk != head_dim_v))) && - // bias type - ((cudnn_runtime_version < 8906 && bias_type == NVTE_Bias_Type::NVTE_NO_BIAS) || - (cudnn_runtime_version >= 8906 && - (bias_type == NVTE_Bias_Type::NVTE_NO_BIAS || - (bias_type == NVTE_Bias_Type::NVTE_ALIBI && - attn_mask_type != NVTE_Mask_Type::NVTE_NO_MASK && - attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_MASK && - attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK && - attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK && - sm_arch_ >= 90) || - (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS && sm_arch_ >= 90))) || - (cudnn_runtime_version >= 90000 && - (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS && sm_arch_ >= 80))) && - // mask type - // pre-8.9.6: causal - ((cudnn_runtime_version < 8906 && attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK) || - // 8.9.6: {bshd, sbhd} + {no_mask, causal, padding, padding_causal} - (cudnn_runtime_version >= 8906 && - (qkv_format == NVTE_QKV_Format::NVTE_SBHD || qkv_format == NVTE_QKV_Format::NVTE_BSHD) && - (attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK)) || - // 9.1: adds thd + {padding, padding_causal} - (cudnn_runtime_version >= 90100 && qkv_format == NVTE_QKV_Format::NVTE_THD && - (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK)) || - // 9.3: adds {bshd, sbhd} + causal_bottom_right + self/cross-attn (sq <= skv) - (cudnn_runtime_version >= 90300 && - (qkv_format == NVTE_QKV_Format::NVTE_SBHD || qkv_format == NVTE_QKV_Format::NVTE_BSHD) && - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK && - max_seqlen_q % 64 == 0 && max_seqlen_kv % 64 == 0 && max_seqlen_q <= max_seqlen_kv && - bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && dropout == 0.0) || - // 9.5: adds {paged_kv_bshd, paged_kv_sbhd} + {padding, padding_causal, padding_causal_bottom_right} - (cudnn_runtime_version >= 90500 && - layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD && - (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK && - max_seqlen_q % 64 == 0 && max_seqlen_kv % 64 == 0 && max_seqlen_q <= max_seqlen_kv)) && - bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && dropout == 0.0) || - // 9.6: adds {bshd, sbhd, thd} + padding_causal_bottom_right + self/cross-attn (sq <= skv) - (cudnn_runtime_version >= 90600 && - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK && - max_seqlen_q % 64 == 0 && max_seqlen_kv % 64 == 0 && max_seqlen_q <= max_seqlen_kv && - bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && dropout == 0.0) || - // 9.7: removes s_q/s_kv % 64 = 0 for {causal_bottom_right, padding_causal_bottom_right} - // for any q_format/kv_format, and paged/non-paged - (cudnn_runtime_version >= 90700 && - (attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - ((attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK) && - bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && dropout == 0.0) || - ((attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK) && - max_seqlen_q <= max_seqlen_kv)))) && - // bias + mask combination - (!(cudnn_runtime_version >= 8906 && - (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) && - bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS)) && - // qkv format - (qkv_format == NVTE_QKV_Format::NVTE_SBHD || qkv_format == NVTE_QKV_Format::NVTE_BSHD || - qkv_format == NVTE_QKV_Format::NVTE_BHSD || - (qkv_format == NVTE_QKV_Format::NVTE_THD && sm_arch_ >= 90 && - ((cudnn_runtime_version >= 90100 && num_attn_heads == num_gqa_groups) || - cudnn_runtime_version >= 90600)) || - ((q_format == NVTE_QKV_Format::NVTE_SBHD || q_format == NVTE_QKV_Format::NVTE_BSHD || - q_format == NVTE_QKV_Format::NVTE_BHSD || - (q_format == NVTE_QKV_Format::NVTE_THD && sm_arch_ >= 90) || - kv_format == NVTE_QKV_Format::NVTE_SBHD || kv_format == NVTE_QKV_Format::NVTE_BSHD || - kv_format == NVTE_QKV_Format::NVTE_BHSD || - (kv_format == NVTE_QKV_Format::NVTE_THD && sm_arch_ >= 90)) && - cudnn_runtime_version >= 90700)) && - // sliding window - // pre-9.2: full attn, causal - ((cudnn_runtime_version < 90200 && window_size_left == -1 && - (window_size_right == -1 || window_size_right == 0)) || - // 9.2: SWA (left, 0) + top-left diagonal + {bshd, sbhd} - (cudnn_runtime_version >= 90200 && - ((window_size_left == -1 && window_size_right == -1 && - attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK) || - ((window_size_left == -1 || window_size_left >= 0) && window_size_right == 0 && - (attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK || - (attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK && - max_seqlen_q == max_seqlen_kv)) && - max_seqlen_q <= max_seqlen_kv && dropout == 0.0 && - bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && - (qkv_format == NVTE_QKV_Format::NVTE_BSHD || - qkv_format == NVTE_QKV_Format::NVTE_SBHD)))) || - // 9.6: SWA (left, 0) + top-left/bottom-right diagonal + {bshd, sbhd, thd} - (cudnn_runtime_version >= 90600 && - ((window_size_left == -1 && (window_size_right == -1 || window_size_right == 0)) || - ((window_size_left >= 0 || window_size_left == -1) && - (window_size_right >= 0 || window_size_right == -1) && - ((attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK && - // TODO(cyang): fix bug for BRCM + cross-attention on sm100 - (sm_arch_ < 100 || (sm_arch_ >= 100 && ((max_seqlen_q == max_seqlen_kv && - cudnn_runtime_version <= 90700) || - cudnn_runtime_version > 90700)))) || - attn_mask_type == NVTE_Mask_Type::NVTE_NO_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK || - attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK || - (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK && - (sm_arch_ < 100 || (sm_arch_ >= 100 && ((max_seqlen_q == max_seqlen_kv && - cudnn_runtime_version <= 90700) || - cudnn_runtime_version > 90700))))) && - max_seqlen_q <= max_seqlen_kv && bias_type == NVTE_Bias_Type::NVTE_NO_BIAS && - dropout == 0.0)))) && - // check 64-bit ragged offset support - (supported_ragged_offset_size) && - // 9.10.0/9.10.1: known bugs with SDPA F16 - (cudnn_runtime_version != 91000) && (cudnn_runtime_version != 91001) && - // softmax type - // pre-9.13.1: vanilla - // 9.13.1+: vanilla, off-by-one, learnable - (cudnn_runtime_version >= 91301 || - softmax_type == NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX) && - // max_logit - // pre-9.21: no (the composite softmax node rejects the Stats + Max output combination) - // 9.21+: yes (Stats + Max via the unified softmax node) - (!return_max_logit || cudnn_runtime_version >= 92100) && - // determinism on Blackwell - // pre-9.18.1: fwd: deterministic; bwd: non-deterministic - // 9.18.1+: fwd: deterministic; bwd: non-deterministic/deterministic - (sm_arch_ < 100 || - (sm_arch_ >= 100 && (!is_training || - (is_training && !deterministic && - (dropout == 0.0 || bias_type == NVTE_Bias_Type::NVTE_NO_BIAS)) || - (is_training && deterministic && cudnn_runtime_version >= 91801 && - dropout == 0.0 && bias_type == NVTE_Bias_Type::NVTE_NO_BIAS))))) { - flag_arb = true; - } - if (flag_arb) { - backend = NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen; - } - if (cudnn_runtime_version < 8900 && - backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - std::cout << "Warning: FP16/BF16 fused attention is supported by cuDNN 8.9.0+." - " Please upgrade your cuDNN version if possible." - << std::endl; - } - if ((cudnn_runtime_version == 91400) && (max_seqlen_kv > 1024) && (window_size_left != -1) && - (attn_mask_type != NVTE_Mask_Type::NVTE_CAUSAL_MASK) && - (attn_mask_type != NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK)) { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - std::cout << "Warning: Given combination of attention mask (non-causal) and " - "max_seqlen_kv (> 1024) does not support fused attention for cuDNN 9.14.0. " - " Please upgrade your cuDNN version if possible." - << std::endl; - } - if ((cudnn_runtime_version <= 91500) && is_training && - (qkv_format == NVTE_QKV_Format::NVTE_BSHD || qkv_format == NVTE_QKV_Format::NVTE_SBHD) && - (max_seqlen_kv % 128 != 0) && cuda_graph && - (attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_MASK) && - (attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) && - (attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)) { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - std::cout << "Warning: Given combination of attention mask (non-padding)," - " max_seqlen_kv (not divisible by 128), and qkv_format (BSHD/SBHD) for" - " backward fused attention with graph capture requires cuDNN 9.15.1+. " - "Please upgrade your cuDNN version if possible." - << std::endl; - } - if (backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen && sm_arch_ == 120) { - if (cudnn_runtime_version < 91801) { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - std::cout << "Warning: Given combination of sm_arch_ == 120 and cudnn_runtime_version < " - "91801 is not supported. " - << " Please upgrade your cuDNN version if possible." << std::endl; - } else if (deterministic && is_training) { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - std::cout << "Warning: Deterministic fused attention on SM120 is not supported." - << std::endl; - } else { - // Known missing support for T3HD/TH3D layouts on SM120 - const bool is_t3hd_or_th3d = - (qkv_layout == NVTE_QKV_Layout::NVTE_T3HD || qkv_layout == NVTE_QKV_Layout::NVTE_TH3D); - if (is_t3hd_or_th3d) { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - std::cout << "Warning: Given combination of T3HD/TH3D layouts on SM120 is not supported. " - << " Please consider using other THD layouts if possible." << std::endl; - } - } - } - } else { - backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend; - } - return backend; -} - -// NVTE fused attention FWD with separate Q, K and V -void nvte_fused_attn_fwd(const NVTETensor Q, const NVTETensor K, const NVTETensor V, - const NVTETensor Bias, const NVTETensor SoftmaxOffset, NVTETensor S, - NVTETensor O, NVTETensorPack *Aux_CTX_Tensors, - const NVTETensor cu_seqlens_q, const NVTETensor cu_seqlens_kv, - const NVTETensor cu_seqlens_q_padded, - const NVTETensor cu_seqlens_kv_padded, const NVTETensor page_table_k, - const NVTETensor page_table_v, const NVTETensor rng_state, - size_t max_seqlen_q, size_t max_seqlen_kv, bool is_training, - bool return_max_logit, bool cuda_graph, float attn_scale, float dropout, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_Bias_Type bias_type, - NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, - int64_t window_size_left, int64_t window_size_right, - bool bottom_right_diagonal, NVTETensor workspace, cudaStream_t stream) { - NVTE_API_CALL(nvte_flash_attn_fwd); - using namespace transformer_engine; - const Tensor *input_cu_seqlens_q = convertNVTETensorCheck(cu_seqlens_q); - const Tensor *input_cu_seqlens_kv = convertNVTETensorCheck(cu_seqlens_kv); - const Tensor *input_cu_seqlens_q_padded = convertNVTETensorCheck(cu_seqlens_q_padded); - const Tensor *input_cu_seqlens_kv_padded = convertNVTETensorCheck(cu_seqlens_kv_padded); - const Tensor *input_page_table_k = convertNVTETensorCheck(page_table_k); - const Tensor *input_page_table_v = convertNVTETensorCheck(page_table_v); - const Tensor *input_rng_state = convertNVTETensorCheck(rng_state); - const Tensor *input_Q = convertNVTETensorCheck(Q); - const Tensor *input_K = convertNVTETensorCheck(K); - const Tensor *input_V = convertNVTETensorCheck(V); - const Tensor *input_Bias = convertNVTETensorCheck(Bias); - const Tensor *input_SoftmaxOffset = convertNVTETensorCheck(SoftmaxOffset); - Tensor *input_output_S = convertNVTETensorCheck(S); - Tensor *output_O = convertNVTETensorCheck(O); - Tensor *wkspace = convertNVTETensor(workspace); - - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - auto *q_dims = input_Q->data.shape.data(); - auto *k_dims = input_K->data.shape.data(); - auto *v_dims = input_V->scaling_mode != NVTE_MXFP8_1D_SCALING - ? input_V->data.shape.data() - : input_V->columnwise_data.shape.data(); - AttentionShape q_shape(q_format, q_dims); - AttentionShape k_shape(kv_format, k_dims); - AttentionShape v_shape(kv_format, v_dims); - size_t b = q_shape.b(), h_q = q_shape.h(), d_qk = q_shape.d(), t_q = q_shape.t(); - size_t h_kv = k_shape.h(), t_kv = k_shape.t(), d_v = v_shape.d(); - if (q_format == NVTE_QKV_Format::NVTE_THD) { - b = input_cu_seqlens_q->data.shape[0] - 1; - } else if (kv_format == NVTE_QKV_Format::NVTE_THD) { - b = input_cu_seqlens_kv->data.shape[0] - 1; - } - - int64_t num_pages_k = 0; - int64_t num_pages_v = 0; - int64_t page_size_k = 0; - int64_t page_size_v = 0; - int64_t max_pages_per_seq_k = 0; - int64_t max_pages_per_seq_v = 0; - if (input_page_table_k->data.dptr != nullptr) { - max_pages_per_seq_k = input_page_table_k->data.shape[1]; - } - if (input_page_table_v->data.dptr != nullptr) { - max_pages_per_seq_v = input_page_table_v->data.shape[1]; - } - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); - if (layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD) { - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - if (kv_format == NVTE_QKV_Format::NVTE_BSHD) { - num_pages_k = input_K->data.shape[0]; - page_size_k = input_K->data.shape[1]; - num_pages_v = input_V->data.shape[0]; - page_size_v = input_V->data.shape[1]; - } else if (kv_format == NVTE_QKV_Format::NVTE_SBHD) { - num_pages_k = input_K->data.shape[1]; - page_size_k = input_K->data.shape[0]; - num_pages_v = input_V->data.shape[1]; - page_size_v = input_V->data.shape[0]; - } - } - - auto handle = cudnnExecutionPlanManager::Instance().GetHandle(); - const NVTEDType Q_type = static_cast(input_Q->data.dtype); - const NVTEDType KV_type = static_cast(input_K->data.dtype); - - NVTE_Fused_Attn_Backend fused_attention_backend = nvte_get_fused_attn_backend( - is_training, Q_type, KV_type, qkv_layout, bias_type, attn_mask_type, softmax_type, dropout, - h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, window_size_left, window_size_right, - return_max_logit, cuda_graph, false); - - if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) { - fused_attn_arbitrary_seqlen_fwd( - b, h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, t_q, t_kv, num_pages_k, num_pages_v, - page_size_k, page_size_v, max_pages_per_seq_k, max_pages_per_seq_v, is_training, - return_max_logit, attn_scale, dropout, qkv_layout, o_format, bias_type, attn_mask_type, - softmax_type, window_size_left, window_size_right, bottom_right_diagonal, input_Q, input_K, - input_V, input_Bias, input_SoftmaxOffset, output_O, Aux_CTX_Tensors, input_cu_seqlens_q, - input_cu_seqlens_kv, input_cu_seqlens_q_padded, input_cu_seqlens_kv_padded, - input_page_table_k, input_page_table_v, input_rng_state, wkspace, stream, handle); - } else if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_FP8) { - fused_attn_fp8_fwd(b, h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, t_q, t_kv, is_training, - attn_scale, dropout, qkv_layout, o_format, qkv_scale_inv_format, bias_type, - attn_mask_type, softmax_type, window_size_left, window_size_right, - bottom_right_diagonal, input_Q, input_K, input_V, input_SoftmaxOffset, - input_output_S, output_O, Aux_CTX_Tensors, input_cu_seqlens_q, - input_cu_seqlens_kv, input_cu_seqlens_q_padded, input_cu_seqlens_kv_padded, - input_rng_state, wkspace, stream, handle); - } else { - NVTE_ERROR("Invalid combination of data type and sequence length for fused attention. \n"); - } -} -// NVTE fused attention BWD with separate Q, K and V -void nvte_fused_attn_bwd(const NVTETensor Q, const NVTETensor K, const NVTETensor V, - const NVTETensor O, const NVTETensor dO, const NVTETensor S, NVTETensor dP, - const NVTETensorPack *Aux_CTX_Tensors, NVTETensor dQ, NVTETensor dK, - NVTETensor dV, NVTETensor dBias, NVTETensor dSoftmaxOffset, - const NVTETensor cu_seqlens_q, const NVTETensor cu_seqlens_kv, - const NVTETensor cu_seqlens_q_padded, - const NVTETensor cu_seqlens_kv_padded, size_t max_seqlen_q, - size_t max_seqlen_kv, float attn_scale, float dropout, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal, bool deterministic, - bool cuda_graph, NVTETensor workspace, cudaStream_t stream) { - NVTE_API_CALL(nvte_flash_attn_bwd); - using namespace transformer_engine; - const Tensor *input_cu_seqlens_q = convertNVTETensorCheck(cu_seqlens_q); - const Tensor *input_cu_seqlens_kv = convertNVTETensorCheck(cu_seqlens_kv); - const Tensor *input_cu_seqlens_q_padded = convertNVTETensorCheck(cu_seqlens_q_padded); - const Tensor *input_cu_seqlens_kv_padded = convertNVTETensorCheck(cu_seqlens_kv_padded); - const Tensor *input_Q = convertNVTETensorCheck(Q); - const Tensor *input_K = convertNVTETensorCheck(K); - const Tensor *input_V = convertNVTETensorCheck(V); - const Tensor *input_O = convertNVTETensorCheck(O); - const Tensor *input_dO = convertNVTETensorCheck(dO); - const Tensor *input_S = convertNVTETensorCheck(S); - Tensor *input_output_dP = convertNVTETensorCheck(dP); - Tensor *output_dQ = convertNVTETensorCheck(dQ); - Tensor *output_dK = convertNVTETensorCheck(dK); - Tensor *output_dV = convertNVTETensorCheck(dV); - Tensor *output_dBias = convertNVTETensorCheck(dBias); - Tensor *output_dSoftmaxOffset = convertNVTETensorCheck(dSoftmaxOffset); - Tensor *wkspace = convertNVTETensor(workspace); - - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - auto *q_dims = input_Q->data.shape.data(); - auto *k_dims = input_K->data.shape.data(); - auto *v_dims = input_V->data.shape.data(); - AttentionShape q_shape(q_format, q_dims); - AttentionShape k_shape(kv_format, k_dims); - AttentionShape v_shape(kv_format, v_dims); - size_t b = q_shape.b(), h_q = q_shape.h(), d_qk = q_shape.d(), t_q = q_shape.t(); - size_t h_kv = k_shape.h(), t_kv = k_shape.t(), d_v = v_shape.d(); - if (q_format == NVTE_QKV_Format::NVTE_THD) { - b = input_cu_seqlens_q->data.shape[0] - 1; - } else if (kv_format == NVTE_QKV_Format::NVTE_THD) { - b = input_cu_seqlens_kv->data.shape[0] - 1; - } - - auto handle = cudnnExecutionPlanManager::Instance().GetHandle(); - const NVTEDType Q_type = static_cast(input_Q->data.dtype); - const NVTEDType KV_type = static_cast(input_K->data.dtype); - - NVTE_Fused_Attn_Backend fused_attention_backend = nvte_get_fused_attn_backend( - true, Q_type, KV_type, qkv_layout, bias_type, attn_mask_type, softmax_type, dropout, h_q, - h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, window_size_left, window_size_right, false, - cuda_graph, deterministic); - - if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) { - size_t i = 0; - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - Tensor *input_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - Tensor *input_Bias, *input_SoftmaxOffset; - if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { - input_Bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - } - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - input_SoftmaxOffset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - } - fused_attn_arbitrary_seqlen_bwd( - b, h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, t_q, t_kv, attn_scale, dropout, - qkv_layout, o_format, do_format, dqkv_layout, bias_type, attn_mask_type, softmax_type, - window_size_left, window_size_right, bottom_right_diagonal, deterministic, input_Q, input_K, - input_V, input_O, input_dO, input_Bias, input_SoftmaxOffset, output_S, output_dQ, output_dK, - output_dV, output_dBias, output_dSoftmaxOffset, input_cu_seqlens_q, input_cu_seqlens_kv, - input_cu_seqlens_q_padded, input_cu_seqlens_kv_padded, input_rng_state, wkspace, stream, - handle); - } else if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_FP8) { - size_t i = 0; - const Tensor *input_M = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - const Tensor *input_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - const Tensor *input_SoftmaxOffset = nullptr; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - input_SoftmaxOffset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - } - const Tensor *input_dO_f16 = nullptr; - if (input_dO->scaling_mode == NVTE_MXFP8_1D_SCALING) { - input_dO_f16 = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - } - fused_attn_fp8_bwd(b, h_q, h_kv, max_seqlen_q, max_seqlen_kv, d_qk, d_v, t_q, t_kv, attn_scale, - dropout, qkv_layout, o_format, do_format, dqkv_layout, qkv_scale_inv_format, - do_scale_inv_format, bias_type, attn_mask_type, softmax_type, - window_size_left, window_size_right, bottom_right_diagonal, deterministic, - input_Q, input_K, input_V, input_O, input_dO, input_dO_f16, input_M, input_S, - input_SoftmaxOffset, input_output_dP, output_dQ, output_dK, output_dV, - output_dSoftmaxOffset, input_cu_seqlens_q, input_cu_seqlens_kv, - input_cu_seqlens_q_padded, input_cu_seqlens_kv_padded, input_rng_state, - wkspace, stream, handle); - } else { - NVTE_ERROR("Invalid combination of data type and sequence length for fused attention. \n"); - } -} - -uint32_t nvte_get_runtime_num_segments(NVTETensor cu_seqlen, NVTETensor workspace, size_t len, - cudaStream_t stream) { - NVTE_API_CALL(nvte_get_runtime_num_segments); - using namespace transformer_engine::fused_attn; - return GetRuntimeNumSegments(cu_seqlen, workspace, len, stream); -} - -void nvte_populate_rng_state_async(NVTETensor rng_state_dst, const NVTETensor seed, - size_t q_max_seqlen, size_t kv_max_seqlen, - NVTE_Fused_Attn_Backend backend, cudaStream_t stream) { - NVTE_API_CALL(nvte_populate_rng_state_async); - using namespace transformer_engine::fused_attn; - PopulateRngStateAsync(rng_state_dst, seed, q_max_seqlen, kv_max_seqlen, backend, stream); -} diff --git a/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu b/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu deleted file mode 100644 index bf34758a352..00000000000 --- a/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu +++ /dev/null @@ -1,1428 +0,0 @@ -/************************************************************************* - * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * - * See LICENSE for license information. - ************************************************************************/ - -#include -#include -#include -#include - -#include -#include - -#include "../common.h" -#include "../cudnn_utils.h" -#include "../util/cuda_runtime.h" -#include "../util/system.h" -#include "fused_attn_f16_arbitrary_seqlen.h" -#include "utils.h" - -#define Q_ID 1 -#define K_ID 2 -#define V_ID 3 -#define O_ID 4 -#define S_ID 5 -#define B_ID 6 -#define D_CONST_ID 7 -#define S_CONST_ID 8 -#define Q_SEQLEN_ID 9 -#define K_SEQLEN_ID 10 -#define dQ_ID 11 -#define dK_ID 12 -#define dV_ID 13 -#define dO_ID 14 -#define MASK_VAL_ID 15 -#define dS_ID 16 -#define D_SEED_ID 17 -#define D_OFFSET_ID 18 -#define S_STATS_ID 19 -#define S_SUM_ID 20 -#define SCALE_PROB 21 -#define K_TRANSPOSE_ID 22 -#define dQ_ACCUM_ID 23 - -#define VIRTUAL_ID 30 - -namespace transformer_engine { -namespace fused_attn { -void fused_attn_arbitrary_seqlen_fwd_impl( - int64_t b, int64_t h, int64_t hg, int64_t s_q, int64_t s_kv, int64_t d_qk, int64_t d_v, - int64_t max_b, int64_t max_t_q, int64_t max_t_kv, int64_t num_pages_k, int64_t num_pages_v, - int64_t page_size_k, int64_t page_size_v, int64_t max_pages_per_seq_k, - int64_t max_pages_per_seq_v, int64_t bias_b, int64_t bias_h, int64_t bias_sq, int64_t bias_skv, - bool is_training, bool return_max_logit, float scaling_factor, float dropout_probability, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal, void *devPtrQ, void *devPtrK, - void *devPtrV, void *devPtrBias, void *devPtrSoftmaxOffset, void *devPtrS1, void *devPtrS2, - void *devPtrO, void *devPtrDropoutSeed, void *devPtrDropoutOffset, void *devPtrCuSeqlensQ, - void *devPtrCuSeqlensKV, void *devPtrPageTableK, void *devPtrPageTableV, - void *devPtrSeqOffsetsQ, void *devPtrSeqOffsetsKV, cudnn_frontend::DataType_t tensorType, - void *workspace, size_t *workspace_size, cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - - bool is_bias = (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS); - bool is_alibi = (bias_type == NVTE_Bias_Type::NVTE_ALIBI); - bool is_causal = ((mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK)); - bool is_bottom_right = ((mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - bool is_padding = ((mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - if (is_bottom_right && s_q == s_kv && !is_padding) { - is_causal = true; - is_bottom_right = false; - bottom_right_diagonal = false; - } - bool is_softmax_offset = (softmax_type != NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX); - bool is_dropout = (is_training && dropout_probability != 0.0f); - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - bool is_ragged_q = (q_format == NVTE_QKV_Format::NVTE_THD); - bool is_ragged_kv = (kv_format == NVTE_QKV_Format::NVTE_THD); - const auto cudnn_runtime_version = cudnnGetVersion(); - const int device_id = cuda::current_device(); - const int sm_arch_ = cuda::sm_arch(device_id); - bool use_ragged_stats = is_ragged_q && cudnn_runtime_version >= 90600 && sm_arch_ != 120; - - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); - bool is_paged_kv = (layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD); - if (is_paged_kv) { - NVTE_CHECK(is_padding, "Paged attention requires padding mask!"); - } - - // Newer versions of cuDNN SDPA can accept sequence lengths directly as a cumulative - // tensor, and can accept ragged offsets in arbitrary units (such as tokens) instead - // of elements. Take advantage of this if possible to avoid 2 extra kernel calls. - const bool use_cu_seqlens_directly = - CUDNN_FRONTEND_VERSION >= 12500 && - // The frontend gates cu_seq_len support on min(compile-time, runtime) cuDNN - // version, so we'll do the same. - (CUDNN_VERSION >= 92400 && cudnn_runtime_version >= 92400) && - // This extra restriction is needed because cuDNN frontend doesn't yet allow - // the combination of dropout and stats generation for the fprop unified engine, - // so any such request would always get routed to the old composite SDPA engine - // (which doesn't support cu_seqlens). Remove this restriction when possible. - !is_dropout; - - // keep original batch size because cu_seqlens are created with [b+1] shape - int64_t actual_b = b; - if ((is_ragged_q || is_ragged_kv) && cudnn_runtime_version >= 90600) { - NVTE_CHECK(is_padding, "Ragged QKV input requires padding or padding_causal mask!"); - // On SM 120, cuDNN support check treats layouts with stride[0] > dim[1]*dim[2]*dim[3] - // as interleaved and rejects them. Use BHSD-like dimensions/strides with max_seqlen at plan build - // so the check passes; ragged offset still provides variable-length boundaries. - if (sm_arch_ != 120) { - // replace batch size and maximum sequence lengths with maximum token counts - // for query and key/value so the graph is static within each quantization bucket. - // When passing cu_seqlens* directly to cuDNN SDPA, keep the true batch size: - // cuDNN reads the user's [actual_b+1] cu_seqlens buffers, so a quantized batch - // would read out of bounds. - if (!use_cu_seqlens_directly) { - b = max_b; - } - s_q = is_ragged_q ? max_t_q : s_q; - s_kv = is_ragged_kv ? max_t_kv : s_kv; - } - } - - const DType ragged_offset_type = - use_cu_seqlens_directly - ? DType::kInt32 // cu_seqlens* are given to us as int32; keep it that way. - : (cudnn_runtime_version >= 90500 ? DType::kInt64 : DType::kInt32); - - // Ragged offset multipliers (elements per token); shared with the legacy conversion - // kernel (cu_seqlens_padded_to_offsets) so the two paths cannot drift apart. - const RaggedOffsetMultipliers offset_mults(layout_group, h, hg, d_qk, d_v); - - bool generate_stats = true; // Always return stats - try { - FADescriptor_v1 descriptor{ - b, - h, - hg, - s_q, - s_kv, - d_qk, - d_v, - num_pages_k, - num_pages_v, - page_size_k, - page_size_v, - max_pages_per_seq_k, - max_pages_per_seq_v, - bias_b, - bias_h, - bias_sq, - bias_skv, - scaling_factor, - is_training, - dropout_probability, - qkv_layout, - o_format, - NVTE_QKV_Format_NOT_SET, - NVTE_QKV_Layout_NOT_SET, - NVTE_QKV_Format_NOT_SET, - NVTE_QKV_Format_NOT_SET, - bias_type, - mask_type, - softmax_type, - window_size_left, - window_size_right, - bottom_right_diagonal, - true, - tensorType, - cudnn_frontend::DataType_t::NOT_SET, - cudnn_frontend::DataType_t::NOT_SET, - cudnn_frontend::DataType_t::NOT_SET, - return_max_logit, - }; - - namespace fe = cudnn_frontend; - using graph_and_tensors = - std::tuple, - std::shared_ptr, // Q - std::shared_ptr, // K - std::shared_ptr, // V - std::shared_ptr, // attn_scale - std::shared_ptr, // O - std::shared_ptr, // S1 - std::shared_ptr, // S2 - std::shared_ptr, // bias - std::shared_ptr, // softmax_offset - std::shared_ptr, // seq_q / cu_seq_len_q - std::shared_ptr, // seq_kv / cu_seq_len_kv - std::shared_ptr, // page_table_k - std::shared_ptr, // page_table_v - std::shared_ptr, // offset_q - std::shared_ptr, // offset_k - std::shared_ptr, // offset_v - std::shared_ptr, // offset_o - std::shared_ptr, // offset_stats - std::shared_ptr, // dropout_seed - std::shared_ptr>; // dropout_offset - - using CacheType = std::map; - static thread_local CacheType sdpa_f16_fprop_cache; - - // Get plan from cache if cache is available, otherwise create one - auto get_graph = [&](CacheType &cache, const FADescriptor_v1 &descriptor) -> graph_and_tensors { - // if hit, return - auto it = cache.find(descriptor); - if (it != cache.end()) { - auto graph = it->second; - return graph; - } - - // otherwise, build the op_graph and the plan. Then update cache - auto mha_graph = std::make_shared(); - mha_graph->set_io_data_type(tensorType) - .set_intermediate_data_type(fe::DataType_t::FLOAT) - .set_compute_data_type(fe::DataType_t::FLOAT); - - std::shared_ptr Q, K, V, attn_scale, softmax_offset; - std::shared_ptr bias, seq_q, seq_kv; - std::shared_ptr page_table_k, page_table_v; - std::shared_ptr offset_q, offset_k, offset_v, offset_o, - offset_stats; - std::shared_ptr dropout_seed, dropout_offset; - - std::vector q_stride(4); - std::vector k_stride(4); - std::vector v_stride(4); - generateMatrixStrides(b, h, s_q, s_kv, d_qk, q_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_Q_Matrix); - if (is_paged_kv) { - generateMatrixStrides(num_pages_k, hg, page_size_k, page_size_v, d_qk, k_stride.data(), - qkv_layout, NVTE_QKV_Matrix::NVTE_K_Matrix); - generateMatrixStrides(num_pages_v, hg, page_size_k, page_size_v, d_v, v_stride.data(), - qkv_layout, NVTE_QKV_Matrix::NVTE_V_Matrix); - } else { - generateMatrixStrides(b, hg, s_q, s_kv, d_qk, k_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_K_Matrix); - generateMatrixStrides(b, hg, s_q, s_kv, d_v, v_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_V_Matrix); - } - - Q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Q") - .set_dim({b, h, s_q, d_qk}) - .set_stride(q_stride)); - if (is_ragged_q) { - offset_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_q") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - Q->set_ragged_offset(offset_q); - if (use_cu_seqlens_directly) { - Q->set_ragged_offset_multiplier(offset_mults.q); - } - } - K = mha_graph->tensor(fe::graph::Tensor_attributes().set_name("K").set_stride(k_stride)); - V = mha_graph->tensor(fe::graph::Tensor_attributes().set_name("V").set_stride(v_stride)); - if (is_paged_kv) { - K->set_dim({num_pages_k, hg, page_size_k, d_qk}); - V->set_dim({num_pages_v, hg, page_size_v, d_v}); - } else if (is_ragged_kv) { - offset_k = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_k") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - offset_v = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_v") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - K->set_dim({b, hg, s_kv, d_qk}).set_ragged_offset(offset_k); - V->set_dim({b, hg, s_kv, d_v}).set_ragged_offset(offset_v); - if (use_cu_seqlens_directly) { - K->set_ragged_offset_multiplier(offset_mults.k); - V->set_ragged_offset_multiplier(offset_mults.v); - } - } else { - K->set_dim({b, hg, s_kv, d_qk}); - V->set_dim({b, hg, s_kv, d_v}); - } - - attn_scale = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("attn_scale") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_is_pass_by_value(true) - .set_data_type(fe::DataType_t::FLOAT)); - - fe::graph::SDPA_attributes sdpa_options; - sdpa_options = fe::graph::SDPA_attributes() - .set_name("flash_attention") - .set_generate_stats(generate_stats) - .set_causal_mask(is_causal) - .set_causal_mask_bottom_right(is_bottom_right) - .set_attn_scale(attn_scale); - - fe::DiagonalAlignment_t const &diagonal_alignment = - bottom_right_diagonal ? fe::DiagonalAlignment_t::BOTTOM_RIGHT - : fe::DiagonalAlignment_t::TOP_LEFT; - sdpa_options.set_diagonal_alignment(diagonal_alignment); - if (cudnn_runtime_version >= 90200 && window_size_left != -1) { - sdpa_options.set_diagonal_band_left_bound(window_size_left + 1); - } - if (cudnn_runtime_version >= 90600 && window_size_right != -1) { - sdpa_options.set_diagonal_band_right_bound(window_size_right); - } - - sdpa_options.set_alibi_mask(is_alibi); - - if (is_bias) { - bias = mha_graph->tensor( - fe::graph::Tensor_attributes() - .set_name("bias") - .set_dim({bias_b, bias_h, bias_sq, bias_skv}) - .set_stride({bias_h * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1})); - sdpa_options.set_bias(bias); - } - - if (is_padding) { - if (use_cu_seqlens_directly) { - // seq_q/seq_kv keep their tuple slots but hold (b+1)-shaped cu_seqlen tensors. - seq_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("cu_seq_len_q") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - seq_kv = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("cu_seq_len_kv") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_options.set_padding_mask(is_padding) - .set_cu_seq_len_q(seq_q) - .set_cu_seq_len_kv(seq_kv); - // cu_seq_len (and the ragged offset multiplier) are unified-engine-only. - // Pin the implementation so an unsupported config fails with the unified - // engine's specific error instead of auto-selection's generic failure. - sdpa_options.set_implementation(fe::AttentionImplementation_t::UNIFIED); - } else { - seq_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_q") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - seq_kv = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_kv") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_options.set_padding_mask(is_padding).set_seq_len_q(seq_q).set_seq_len_kv(seq_kv); - } - } - - if (is_paged_kv) { - page_table_k = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("page_table_k") - .set_dim({b, 1, max_pages_per_seq_k, 1}) - .set_stride({{max_pages_per_seq_k, max_pages_per_seq_v, 1, 1}}) - .set_data_type(fe::DataType_t::INT32)); - page_table_v = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("page_table_v") - .set_dim({b, 1, max_pages_per_seq_v, 1}) - .set_stride({{max_pages_per_seq_v, max_pages_per_seq_v, 1, 1}}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_options.set_paged_attention_k_table(page_table_k); - sdpa_options.set_paged_attention_v_table(page_table_v); - sdpa_options.set_paged_attention_max_seq_len_kv(static_cast(s_kv)); - } - - if (is_dropout) { - dropout_seed = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Seed") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - dropout_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Offset") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - sdpa_options.set_dropout(dropout_probability, dropout_seed, dropout_offset); - } - - if (is_softmax_offset) { - softmax_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("softmax_offset") - .set_dim({1, h, 1, 1}) - .set_stride({h, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - sdpa_options.set_sink_token(softmax_offset); - } - - std::shared_ptr Max; - if (use_ragged_stats) { - offset_stats = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_stats") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - } - if (return_max_logit) { - Max = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Max") - .set_dim({b, h, s_q, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - if (use_ragged_stats) { - Max->set_stride({h * s_q, 1, h, 1}).set_ragged_offset(offset_stats); - if (use_cu_seqlens_directly) { - Max->set_ragged_offset_multiplier(offset_mults.stats); - } - } else { - Max->set_stride({h * s_q, s_q, 1, 1}); - } - sdpa_options.set_logit_max(Max); - } - - auto [O, Stats] = mha_graph->sdpa(Q, K, V, std::move(sdpa_options)); - - std::vector o_stride(4); - generateMatrixStrides(b, h, s_q, s_kv, d_v, o_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_O_Matrix); - O->set_output(true).set_dim({b, h, s_q, d_v}).set_stride(o_stride); - if (is_ragged_q) { - offset_o = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_o") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - O->set_ragged_offset(offset_o); - if (use_cu_seqlens_directly) { - O->set_ragged_offset_multiplier(offset_mults.o); - } - } - - Stats->set_output(true).set_data_type(fe::DataType_t::FLOAT).set_dim({b, h, s_q, 1}); - if (use_ragged_stats) { - Stats->set_stride({h * s_q, 1, h, 1}).set_ragged_offset(offset_stats); - if (use_cu_seqlens_directly) { - Stats->set_ragged_offset_multiplier(offset_mults.stats); - } - } else { - Stats->set_stride({h * s_q, s_q, 1, 1}); - } - - std::tuple, // Q - std::shared_ptr, // K - std::shared_ptr, // V - std::shared_ptr, // attn_scale - std::shared_ptr> // O - key_tensors_tuple = std::make_tuple(Q, K, V, attn_scale, O); - auto Stats_tuple = - return_max_logit ? std::make_tuple(Stats, Max) : std::make_tuple(Stats, nullptr); - auto bias_tuple = is_bias ? std::make_tuple(bias) : std::make_tuple(nullptr); - auto softmax_offset_tuple = - is_softmax_offset ? std::make_tuple(softmax_offset) : std::make_tuple(nullptr); - auto padding_tuple = - is_padding ? std::make_tuple(seq_q, seq_kv) : std::make_tuple(nullptr, nullptr); - auto page_table_tuple = is_paged_kv ? std::make_tuple(page_table_k, page_table_v) - : std::make_tuple(nullptr, nullptr); - auto offset_qo_tuple = - is_ragged_q ? std::make_tuple(offset_q, offset_o) : std::make_tuple(nullptr, nullptr); - auto offset_kv_tuple = - is_ragged_kv ? std::make_tuple(offset_k, offset_v) : std::make_tuple(nullptr, nullptr); - auto offset_s_tuple = - use_ragged_stats ? std::make_tuple(offset_stats) : std::make_tuple(nullptr); - auto dropout_tuple = is_dropout ? std::make_tuple(dropout_seed, dropout_offset) - : std::make_tuple(nullptr, nullptr); - - NVTE_CHECK_CUDNN_FE(mha_graph->validate()); - NVTE_CHECK_CUDNN_FE(mha_graph->build_operation_graph(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->create_execution_plans({fe::HeurMode_t::A})); - NVTE_CHECK_CUDNN_FE(mha_graph->check_support(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->build_plans(handle)); - - auto return_tuple = - std::tuple_cat(std::make_tuple(mha_graph), key_tensors_tuple, Stats_tuple, bias_tuple, - softmax_offset_tuple, padding_tuple, page_table_tuple, offset_qo_tuple, - offset_kv_tuple, offset_s_tuple, dropout_tuple); - cache.insert({descriptor, return_tuple}); - - return return_tuple; - }; - - auto [mha_graph, Q, K, V, attn_scale, O, S1, S2, bias, softmax_offset, seq_q, seq_kv, - page_table_k, page_table_v, offset_q, offset_o, offset_k, offset_v, offset_stats, - dropout_seed, dropout_offset] = get_graph(sdpa_f16_fprop_cache, descriptor); - - // Exit to request upper level API to allocate memory if needed - // n.b. Care should be taken to align each of the added worksapce tensors to their type. - // We do this by adding padding at the end of each separate allocation. - // When passing cu_seqlens* directly to cuDNN SDPA, no conversion workspace is - // needed: cuDNN consumes the user's cu_seqlens buffers as-is. - auto plan_workspace_size = alignTo<16>(mha_graph->get_workspace_size()); - const size_t num_bytes_per_seqlen = alignTo<16>(b * sizeof(int32_t)); - const size_t num_bytes_per_ragged_offset = - alignTo<16>(((b + 1) * typeToNumBits(ragged_offset_type)) / 8); - size_t actual_seqlen_workspace_size = 0; - size_t seqlen_offsets_workspace_size = 0; - if (!use_cu_seqlens_directly) { - if (is_padding) { - actual_seqlen_workspace_size = 2 * num_bytes_per_seqlen; - } - if (is_ragged_q || is_ragged_kv) { - const size_t count = - 2 * (static_cast(is_ragged_q) + static_cast(is_ragged_kv)); - seqlen_offsets_workspace_size = - (use_ragged_stats ? count + 1 : count) * num_bytes_per_ragged_offset; - } - } - if (workspace == nullptr) { - *workspace_size = - plan_workspace_size + actual_seqlen_workspace_size + seqlen_offsets_workspace_size; - return; - } - - // cuDNN stream check needs to be moved here to support dummy kernel calls with - // null streams for sizing the cuDNN workspace. - NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); - - // Build variant pack - std::unordered_map, void *> variant_pack = { - {Q, devPtrQ}, {K, devPtrK}, {V, devPtrV}, {attn_scale, &scaling_factor}, - {O, devPtrO}, {S1, devPtrS1}}; - - if (return_max_logit) { - variant_pack[S2] = devPtrS2; - } - - if (is_bias) { - variant_pack[bias] = devPtrBias; - } - - if (is_padding) { - if (use_cu_seqlens_directly) { - variant_pack[seq_q] = devPtrCuSeqlensQ; - variant_pack[seq_kv] = devPtrCuSeqlensKV; - } else { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block - 1) / nthreads_per_block; - void *devActualSeqlenQ = static_cast(workspace) + plan_workspace_size; - void *devActualSeqlenKV = static_cast(devActualSeqlenQ) + num_bytes_per_seqlen; - cu_seqlens_to_actual_seqlens<<>>( - actual_b, b, static_cast(devPtrCuSeqlensQ), - static_cast(devPtrCuSeqlensKV), - static_cast(devActualSeqlenQ), static_cast(devActualSeqlenKV)); - NVTE_CHECK_CUDA(cudaGetLastError()); - variant_pack[seq_q] = devActualSeqlenQ; - variant_pack[seq_kv] = devActualSeqlenKV; - } - } - - if (is_paged_kv) { - variant_pack[page_table_k] = devPtrPageTableK; - variant_pack[page_table_v] = devPtrPageTableV; - } - - if (use_cu_seqlens_directly) { - // The token-unit cu_seqlens_padded buffers serve as the ragged offsets; the engine - // applies the per-tensor multipliers set at graph build time. - if (is_ragged_q) { - variant_pack[offset_q] = devPtrSeqOffsetsQ; - variant_pack[offset_o] = devPtrSeqOffsetsQ; - } - if (is_ragged_kv) { - void *devOffsetsKV = offset_mults.kv_from_q ? devPtrSeqOffsetsQ : devPtrSeqOffsetsKV; - variant_pack[offset_k] = devOffsetsKV; - variant_pack[offset_v] = devOffsetsKV; - } - if (use_ragged_stats) { - variant_pack[offset_stats] = devPtrSeqOffsetsQ; - } - } else if (is_ragged_q || is_ragged_kv) { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block) / nthreads_per_block; - void *devOffsets = - static_cast(workspace) + plan_workspace_size + actual_seqlen_workspace_size; - void *devOffsetsQ = nullptr; - void *devOffsetsO = nullptr; - if (is_ragged_q) { - devOffsetsQ = devOffsets; - devOffsetsO = static_cast(devOffsetsQ) + num_bytes_per_ragged_offset; - } - void *devOffsetsK = nullptr; - void *devOffsetsV = nullptr; - if (is_ragged_kv) { - devOffsetsK = static_cast(devOffsets) + - static_cast(is_ragged_q) * 2 * num_bytes_per_ragged_offset; - devOffsetsV = static_cast(devOffsetsK) + num_bytes_per_ragged_offset; - } - void *devOffsetsS = nullptr; - if (use_ragged_stats) { - devOffsetsS = static_cast(devOffsets) + - (static_cast(is_ragged_q) + static_cast(is_ragged_kv)) * 2 * - num_bytes_per_ragged_offset; - } - cu_seqlens_padded_to_offsets<<>>( - offset_mults, actual_b, b, static_cast(devPtrSeqOffsetsQ), - static_cast(devPtrSeqOffsetsKV), ragged_offset_type, devOffsetsQ, devOffsetsK, - devOffsetsV, devOffsetsO, devOffsetsS); - NVTE_CHECK_CUDA(cudaGetLastError()); - if (is_ragged_q) { - variant_pack[offset_q] = devOffsetsQ; - variant_pack[offset_o] = devOffsetsO; - } - if (is_ragged_kv) { - variant_pack[offset_k] = devOffsetsK; - variant_pack[offset_v] = devOffsetsV; - } - if (use_ragged_stats) { - variant_pack[offset_stats] = devOffsetsS; - } - } - - if (is_dropout) { - variant_pack[dropout_seed] = devPtrDropoutSeed; - variant_pack[dropout_offset] = devPtrDropoutOffset; - } - - if (is_softmax_offset) { - variant_pack[softmax_offset] = devPtrSoftmaxOffset; - } - - NVTE_CHECK_CUDNN_FE(mha_graph->execute(handle, variant_pack, workspace)); - } catch (cudnn_frontend::cudnnException &e) { - NVTE_ERROR(e.what()); - } -} // NOLINT(readability/fn_size) - -void fused_attn_arbitrary_seqlen_bwd_impl( - int64_t b, int64_t h, int64_t hg, int64_t s_q, int64_t s_kv, int64_t d_qk, int64_t d_v, - int64_t max_b, int64_t max_t_q, int64_t max_t_kv, int64_t bias_b, int64_t bias_h, - int64_t bias_sq, int64_t bias_skv, float scaling_factor, float dropout_probability, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, - NVTE_QKV_Layout dqkv_layout, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, - NVTE_Softmax_Type softmax_type, int64_t window_size_left, int64_t window_size_right, - bool bottom_right_diagonal, bool deterministic, void *devPtrQ, void *devPtrKTranspose, - void *devPtrVTranspose, void *devPtrO, void *devPtrSoftmaxStats, void *devPtrBias, - void *devPtrSoftmaxOffset, void *devPtrdQ, void *devPtrdK, void *devPtrdV, void *devPtrdO, - void *devPtrdBias, void *devPtrdSoftmaxOffset, void *devPtrDropoutSeed, - void *devPtrDropoutOffset, void *devPtrCuSeqlensQ, void *devPtrCuSeqlensKV, - void *devPtrSeqOffsetsQ, void *devPtrSeqOffsetsKV, cudnn_frontend::DataType_t tensorType, - void *workspace, size_t *workspace_size, cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - - bool is_bias = (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS); - bool is_alibi = (bias_type == NVTE_Bias_Type::NVTE_ALIBI); - bool is_causal = ((mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK)); - bool is_bottom_right = ((mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - bool is_padding = ((mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - if (is_bottom_right && s_q == s_kv && !is_padding) { - is_causal = true; - is_bottom_right = false; - bottom_right_diagonal = false; - } - bool is_softmax_offset = (softmax_type != NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX); - bool is_dropout = (dropout_probability != 0.0f); - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - bool is_ragged_q = (q_format == NVTE_QKV_Format::NVTE_THD); - bool is_ragged_kv = (kv_format == NVTE_QKV_Format::NVTE_THD); - const auto cudnn_runtime_version = cudnnGetVersion(); - const int device_id = cuda::current_device(); - const int sm_arch_ = cuda::sm_arch(device_id); - bool use_ragged_stats = is_ragged_q && cudnn_runtime_version >= 90600 && sm_arch_ != 120; - - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); - bool is_paged_kv = (layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD); - if (is_paged_kv) { - NVTE_CHECK(is_padding, "Paged attention requires padding mask!"); - } - - // keep original batch size because cu_seqlens are created with [b+1] shape - int64_t actual_b = b; - if ((is_ragged_q || is_ragged_kv) && cudnn_runtime_version >= 90600) { - NVTE_CHECK(is_padding, "Ragged QKV input requires padding or padding_causal mask!"); - // On SM 120, cuDNN support check requires BHSD-like strides with max_seqlen (see fwd). - if (sm_arch_ != 120) { - // replace batch size and maximum sequence lengths with maximum token counts - // for query and key/value so the graph is static within each quantization bucket - b = max_b; - s_q = is_ragged_q ? max_t_q : s_q; - s_kv = is_ragged_kv ? max_t_kv : s_kv; - } - } - // We choose between 32-bit and 64-bit offsets depending on need. - // This allows us to support older cuDNN runtimes gracefully. - const DType ragged_offset_type = cudnn_runtime_version >= 90500 ? DType::kInt64 : DType::kInt32; - - try { - FADescriptor_v1 descriptor{ - b, - h, - hg, - s_q, - s_kv, - d_qk, - d_v, - 0, - 0, - 0, - 0, - 0, - 0, - bias_b, - bias_h, - bias_sq, - bias_skv, - scaling_factor, - true, - dropout_probability, - qkv_layout, - o_format, - do_format, - dqkv_layout, - NVTE_QKV_Format_NOT_SET, - NVTE_QKV_Format_NOT_SET, - bias_type, - mask_type, - softmax_type, - window_size_left, - window_size_right, - bottom_right_diagonal, - deterministic, - tensorType, - cudnn_frontend::DataType_t::NOT_SET, - cudnn_frontend::DataType_t::NOT_SET, - cudnn_frontend::DataType_t::NOT_SET, - false, - }; - - namespace fe = cudnn_frontend; - using graph_and_tensors = - std::tuple, - std::shared_ptr, // q - std::shared_ptr, // k - std::shared_ptr, // v - std::shared_ptr, // o - std::shared_ptr, // dO - std::shared_ptr, // stats - std::shared_ptr, // attn_scale - std::shared_ptr, // dQ - std::shared_ptr, // dK - std::shared_ptr, // dV - std::shared_ptr, // bias - std::shared_ptr, // dBias - std::shared_ptr, // softmax_offset - std::shared_ptr, // d_softmax_offset - std::shared_ptr, // seq_q - std::shared_ptr, // seq_kv - std::shared_ptr, // offset_q - std::shared_ptr, // offset_k - std::shared_ptr, // offset_v - std::shared_ptr, // offset_o - std::shared_ptr, // offset_stats - std::shared_ptr, // dropout_seed - std::shared_ptr>; // dropout_offset - - using CacheType = std::map; - static thread_local CacheType sdpa_f16_bprop_cache; - - // Get plan from cache if cache is available, otherwise create one - auto get_graph = [&](CacheType &cache, const FADescriptor_v1 &descriptor) -> graph_and_tensors { - // if hit, return - auto it = cache.find(descriptor); - if (it != cache.end()) { - auto graph = it->second; - return graph; - } - - // otherwise, build the op_graph and the plan. Then update cache - auto mha_graph = std::make_shared(); - mha_graph->set_io_data_type(tensorType) - .set_intermediate_data_type(fe::DataType_t::FLOAT) - .set_compute_data_type(fe::DataType_t::FLOAT); - - std::shared_ptr q, k, v, o, dO, stats, attn_scale; - std::shared_ptr bias, dBias, softmax_offset, d_softmax_offset, - seq_q, seq_kv; - std::shared_ptr offset_q, offset_k, offset_v, offset_o, - offset_stats; - std::shared_ptr dropout_seed, dropout_offset; - - std::vector q_stride(4); - std::vector k_stride(4); - std::vector v_stride(4); - std::vector o_stride(4); - generateMatrixStrides(b, h, s_q, s_kv, d_qk, q_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_Q_Matrix); - generateMatrixStrides(b, hg, s_q, s_kv, d_qk, k_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_K_Matrix); - generateMatrixStrides(b, hg, s_q, s_kv, d_v, v_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_V_Matrix); - generateMatrixStrides(b, h, s_q, s_kv, d_v, o_stride.data(), qkv_layout, - NVTE_QKV_Matrix::NVTE_O_Matrix); - - q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Q") - .set_dim({b, h, s_q, d_qk}) - .set_stride(q_stride)); - k = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("K") - .set_dim({b, hg, s_kv, d_qk}) - .set_stride(k_stride)); - v = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("V") - .set_dim({b, hg, s_kv, d_v}) - .set_stride(v_stride)); - o = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("O") - .set_dim({b, h, s_q, d_v}) - .set_stride(o_stride)); - dO = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("dO") - .set_dim({b, h, s_q, d_v}) - .set_stride(o_stride)); - if (is_ragged_q) { - offset_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_q") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - offset_o = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_o") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - q->set_ragged_offset(offset_q); - o->set_ragged_offset(offset_o); - dO->set_ragged_offset(offset_o); - } - if (is_ragged_kv) { - offset_k = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_k") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - offset_v = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_v") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - k->set_ragged_offset(offset_k); - v->set_ragged_offset(offset_v); - } - - stats = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("stats") - .set_dim({b, h, s_q, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - if (use_ragged_stats) { - offset_stats = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_stats") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - stats->set_stride({h * s_q, 1, h, 1}).set_ragged_offset(offset_stats); - } else { - stats->set_stride({h * s_q, s_q, 1, 1}); - } - - attn_scale = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("attn_scale") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_is_pass_by_value(true) - .set_data_type(fe::DataType_t::FLOAT)); - - fe::graph::SDPA_backward_attributes sdpa_backward_options; - sdpa_backward_options = fe::graph::SDPA_backward_attributes() - .set_name("flash_attention_backward") - .set_causal_mask(is_causal) - .set_causal_mask_bottom_right(is_bottom_right) - .set_attn_scale(attn_scale); - - if (use_ragged_stats) { - sdpa_backward_options.set_max_total_seq_len_q(s_q); - } - if (is_ragged_kv && cudnn_runtime_version >= 90600 && sm_arch_ != 120) { - sdpa_backward_options.set_max_total_seq_len_kv(s_kv); - } - - fe::DiagonalAlignment_t const &diagonal_alignment = - bottom_right_diagonal ? fe::DiagonalAlignment_t::BOTTOM_RIGHT - : fe::DiagonalAlignment_t::TOP_LEFT; - sdpa_backward_options.set_diagonal_alignment(diagonal_alignment); - - if (cudnn_runtime_version >= 90200 && window_size_left != -1) { - sdpa_backward_options.set_diagonal_band_left_bound(window_size_left + 1); - } - if (cudnn_runtime_version >= 90600 && window_size_right != -1) { - sdpa_backward_options.set_diagonal_band_right_bound(window_size_right); - } - - if (cudnn_runtime_version >= 90000) { - sdpa_backward_options.set_deterministic_algorithm(deterministic); - } - - sdpa_backward_options.set_alibi_mask(is_alibi); - - if (is_bias) { - bias = mha_graph->tensor( - fe::graph::Tensor_attributes() - .set_name("bias") - .set_dim({bias_b, bias_h, bias_sq, bias_skv}) - .set_stride({bias_h * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1})); - sdpa_backward_options.set_bias(bias); - // bias shapes [1, 1, s, s], [b, 1, s, s], [b, h, s, s], [1, h, s, s] are supported for dbias calculation - // bias shape [1, 1, 1, s] is not supported for dbias calculation as of cuDNN 9.18 - if (!((bias_b == 1) && (bias_h == 1) && (bias_sq == 1))) { - dBias = mha_graph->tensor( - fe::graph::Tensor_attributes() - .set_name("dBias") - .set_dim({bias_b, bias_h, bias_sq, bias_skv}) - .set_stride({bias_h * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1})); - sdpa_backward_options.set_dbias(dBias); - } - } - - if (is_padding) { - seq_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_q") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - seq_kv = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_kv") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_backward_options.set_padding_mask(is_padding) - .set_seq_len_q(seq_q) - .set_seq_len_kv(seq_kv); - } - - if (is_dropout) { - dropout_seed = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Seed") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - dropout_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Offset") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - sdpa_backward_options.set_dropout(dropout_probability, dropout_seed, dropout_offset); - } - - if (is_softmax_offset) { - softmax_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("softmax_offset") - .set_dim({1, h, 1, 1}) - .set_stride({h, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - sdpa_backward_options.set_sink_token(softmax_offset); - d_softmax_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("d_softmax_offset") - .set_dim({1, h, 1, 1}) - .set_stride({h, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - sdpa_backward_options.set_dsink_token(d_softmax_offset); - } - - auto [dQ, dK, dV] = mha_graph->sdpa_backward(q, k, v, o, dO, stats, sdpa_backward_options); - - dQ->set_output(true).set_dim({b, h, s_q, d_qk}).set_stride(q_stride); - dK->set_output(true).set_dim({b, hg, s_kv, d_qk}).set_stride(k_stride); - dV->set_output(true).set_dim({b, hg, s_kv, d_v}).set_stride(v_stride); - if (is_ragged_q) { - dQ->set_ragged_offset(offset_q); - } - if (is_ragged_kv) { - dK->set_ragged_offset(offset_k); - dV->set_ragged_offset(offset_v); - } - - std::tuple, // q - std::shared_ptr, // k - std::shared_ptr, // v - std::shared_ptr, // o - std::shared_ptr, // dO - std::shared_ptr, // stats - std::shared_ptr, // attn_scale - std::shared_ptr, // dQ - std::shared_ptr, // dK - std::shared_ptr> // dV - key_tensors_tuple = std::make_tuple(q, k, v, o, dO, stats, attn_scale, dQ, dK, dV); - auto bias_tuple = is_bias ? std::make_tuple(bias, dBias) : std::make_tuple(nullptr, nullptr); - auto softmax_offset_tuple = is_softmax_offset - ? std::make_tuple(softmax_offset, d_softmax_offset) - : std::make_tuple(nullptr, nullptr); - auto padding_tuple = - is_padding ? std::make_tuple(seq_q, seq_kv) : std::make_tuple(nullptr, nullptr); - auto offset_qo_tuple = - is_ragged_q ? std::make_tuple(offset_q, offset_o) : std::make_tuple(nullptr, nullptr); - auto offset_kv_tuple = - is_ragged_kv ? std::make_tuple(offset_k, offset_v) : std::make_tuple(nullptr, nullptr); - auto offset_s_tuple = - use_ragged_stats ? std::make_tuple(offset_stats) : std::make_tuple(nullptr); - auto dropout_tuple = is_dropout ? std::make_tuple(dropout_seed, dropout_offset) - : std::make_tuple(nullptr, nullptr); - - NVTE_CHECK_CUDNN_FE(mha_graph->validate()); - NVTE_CHECK_CUDNN_FE(mha_graph->build_operation_graph(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->create_execution_plans({fe::HeurMode_t::A})); - NVTE_CHECK_CUDNN_FE(mha_graph->check_support(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->build_plans(handle)); - - auto return_tuple = std::tuple_cat(std::make_tuple(mha_graph), key_tensors_tuple, bias_tuple, - softmax_offset_tuple, padding_tuple, offset_qo_tuple, - offset_kv_tuple, offset_s_tuple, dropout_tuple); - cache.insert({descriptor, return_tuple}); - - return return_tuple; - }; - - auto [mha_graph, q, k, v, o, dO, stats, attn_scale, dQ, dK, dV, bias, dBias, softmax_offset, - d_softmax_offset, seq_q, seq_kv, offset_q, offset_o, offset_k, offset_v, offset_stats, - dropout_seed, dropout_offset] = get_graph(sdpa_f16_bprop_cache, descriptor); - - // Exit to request upper level API to allocate memory if needed - // n.b. Care should be taken to align each of the added worksapce tensors to their type. - // We do this by adding padding at the end of each separate allocation. - auto plan_workspace_size = alignTo<16>(mha_graph->get_workspace_size()); - const size_t num_bytes_per_seqlen = alignTo<16>(b * sizeof(int32_t)); - const size_t actual_seqlen_workspace_size = is_padding ? 2 * num_bytes_per_seqlen : 0; - const size_t num_bytes_per_ragged_offset = - alignTo<16>(((b + 1) * typeToNumBits(ragged_offset_type)) / 8); - size_t seqlen_offsets_workspace_size = 0; - if (is_ragged_q || is_ragged_kv) { - size_t count = 2 * (static_cast(is_ragged_q) + static_cast(is_ragged_kv)); - if (use_ragged_stats) { - seqlen_offsets_workspace_size = (count + 1) * num_bytes_per_ragged_offset; - } else { - seqlen_offsets_workspace_size = count * num_bytes_per_ragged_offset; - } - } - if (workspace == nullptr) { - *workspace_size = - plan_workspace_size + actual_seqlen_workspace_size + seqlen_offsets_workspace_size; - return; - } - - // cuDNN stream check needs to be moved here to support dummy kernel calls with - // null streams for sizing the cuDNN workspace. - NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); - - // build variant pack - std::unordered_map, void *> variant_pack = { - {q, devPtrQ}, - {k, devPtrKTranspose}, - {v, devPtrVTranspose}, - {o, devPtrO}, - {dO, devPtrdO}, - {stats, devPtrSoftmaxStats}, - {attn_scale, &scaling_factor}, - {dQ, devPtrdQ}, - {dK, devPtrdK}, - {dV, devPtrdV}, - }; - - if (is_bias) { - variant_pack[bias] = devPtrBias; - if (dBias != nullptr) { - variant_pack[dBias] = devPtrdBias; - } - } - - if (is_padding) { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block - 1) / nthreads_per_block; - void *devActualSeqlenQ = static_cast(workspace) + plan_workspace_size; - void *devActualSeqlenKV = static_cast(devActualSeqlenQ) + num_bytes_per_seqlen; - cu_seqlens_to_actual_seqlens<<>>( - actual_b, b, static_cast(devPtrCuSeqlensQ), - static_cast(devPtrCuSeqlensKV), static_cast(devActualSeqlenQ), - static_cast(devActualSeqlenKV)); - NVTE_CHECK_CUDA(cudaGetLastError()); - variant_pack[seq_q] = devActualSeqlenQ; - variant_pack[seq_kv] = devActualSeqlenKV; - } - - if (is_ragged_q || is_ragged_kv) { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block) / nthreads_per_block; - void *devOffsets = - static_cast(workspace) + plan_workspace_size + actual_seqlen_workspace_size; - void *devOffsetsQ = nullptr; - void *devOffsetsO = nullptr; - if (is_ragged_q) { - devOffsetsQ = devOffsets; - devOffsetsO = static_cast(devOffsetsQ) + num_bytes_per_ragged_offset; - } - void *devOffsetsK = nullptr; - void *devOffsetsV = nullptr; - if (is_ragged_kv) { - devOffsetsK = static_cast(devOffsets) + - static_cast(is_ragged_q) * 2 * num_bytes_per_ragged_offset; - devOffsetsV = static_cast(devOffsetsK) + num_bytes_per_ragged_offset; - } - void *devOffsetsS = nullptr; - if (use_ragged_stats) { - devOffsetsS = static_cast(devOffsets) + - (static_cast(is_ragged_q) + static_cast(is_ragged_kv)) * 2 * - num_bytes_per_ragged_offset; - } - const RaggedOffsetMultipliers offset_mults(nvte_get_qkv_layout_group(qkv_layout), h, hg, d_qk, - d_v); - cu_seqlens_padded_to_offsets<<>>( - offset_mults, actual_b, b, static_cast(devPtrSeqOffsetsQ), - static_cast(devPtrSeqOffsetsKV), ragged_offset_type, devOffsetsQ, devOffsetsK, - devOffsetsV, devOffsetsO, devOffsetsS); - NVTE_CHECK_CUDA(cudaGetLastError()); - if (is_ragged_q) { - variant_pack[offset_q] = devOffsetsQ; - variant_pack[offset_o] = devOffsetsO; - } - if (is_ragged_kv) { - variant_pack[offset_k] = devOffsetsK; - variant_pack[offset_v] = devOffsetsV; - } - if (use_ragged_stats) { - variant_pack[offset_stats] = devOffsetsS; - } - } - - if (is_dropout) { - variant_pack[dropout_seed] = devPtrDropoutSeed; - variant_pack[dropout_offset] = devPtrDropoutOffset; - } - - if (is_softmax_offset) { - variant_pack[softmax_offset] = devPtrSoftmaxOffset; - variant_pack[d_softmax_offset] = devPtrdSoftmaxOffset; - } - - NVTE_CHECK_CUDNN_FE(mha_graph->execute(handle, variant_pack, workspace)); - } catch (cudnn_frontend::cudnnException &e) { - NVTE_ERROR(e.what()); - } -} -} // namespace fused_attn - -using namespace transformer_engine::fused_attn; -void fused_attn_arbitrary_seqlen_fwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, size_t num_pages_k, size_t num_pages_v, size_t page_size_k, - size_t page_size_v, size_t max_pages_per_seq_k, size_t max_pages_per_seq_v, bool is_training, - bool return_max_logit, float attn_scale, float p_dropout, NVTE_QKV_Layout qkv_layout, - NVTE_QKV_Format o_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, - NVTE_Softmax_Type softmax_type, int64_t window_size_left, int64_t window_size_right, - bool bottom_right_diagonal, const Tensor *input_Q, const Tensor *input_K, const Tensor *input_V, - const Tensor *input_Bias, const Tensor *input_SoftmaxOffset, Tensor *output_O, - NVTETensorPack *Aux_CTX_Tensors, const Tensor *cu_seqlens_q, const Tensor *cu_seqlens_kv, - const Tensor *cu_seqlens_q_padded, const Tensor *cu_seqlens_kv_padded, - const Tensor *page_table_k, const Tensor *page_table_v, const Tensor *rng_state, - Tensor *workspace, cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - - const auto QKV_type = input_Q->data.dtype; - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - void *devPtrQ = input_Q->data.dptr; - void *devPtrK = input_K->data.dptr; - void *devPtrV = input_V->data.dptr; - void *devPtrO = output_O->data.dptr; - void *devPtrS1 = nullptr; - void *devPtrS2 = nullptr; - void *devPtrBias = nullptr; - size_t bias_b = 0; - size_t bias_h = 0; - size_t bias_sq = 0; - size_t bias_skv = 0; - if ((bias_type != NVTE_Bias_Type::NVTE_NO_BIAS) && (bias_type != NVTE_Bias_Type::NVTE_ALIBI)) { - devPtrBias = input_Bias->data.dptr; - bias_b = input_Bias->data.shape[0]; - bias_h = input_Bias->data.shape[1]; - bias_sq = input_Bias->data.shape[2]; - bias_skv = input_Bias->data.shape[3]; - } - void *devPtrSoftmaxOffset = nullptr; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - devPtrSoftmaxOffset = input_SoftmaxOffset->data.dptr; - } - - const int device_id = cuda::current_device(); - const int sm_arch_ = cuda::sm_arch(device_id); - - void *devPtrCuSeqlensQ = cu_seqlens_q->data.dptr; - void *devPtrCuSeqlensKV = cu_seqlens_kv->data.dptr; - void *devPtrSeqOffsetsQ = cu_seqlens_q_padded->data.dptr; - void *devPtrSeqOffsetsKV = cu_seqlens_kv_padded->data.dptr; - void *devPtrPageTableK = page_table_k ? page_table_k->data.dptr : nullptr; - void *devPtrPageTableV = page_table_v ? page_table_v->data.dptr : nullptr; - - size_t max_batch_size = 0; - size_t max_tokens_q = 0; - size_t max_tokens_kv = 0; - if (q_format == NVTE_QKV_Format::NVTE_THD || kv_format == NVTE_QKV_Format::NVTE_THD) { - max_batch_size = get_max_batch_size(batch); - } - if (q_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_q = get_max_tokens(num_tokens_q); - } - if (kv_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_kv = get_max_tokens(num_tokens_kv); - } - - size_t i = 0; - if (Aux_CTX_Tensors->size == 0) { - const auto cudnn_runtime_version = cudnnGetVersion(); - - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_S->data.dptr = nullptr; - // sm120 does not use ragged stats: the graph declares a dense - // [b, h, s_q, 1] stats tensor, so allocate to match (same as Max below). - if ((q_format == NVTE_QKV_Format::NVTE_THD && cudnn_runtime_version >= 90600) && - (sm_arch_ != 120)) { - output_S->data.shape = {num_tokens_q, num_attn_heads, 1}; - } else { - output_S->data.shape = {batch, num_attn_heads, max_seqlen_q, 1}; - } - output_S->data.dtype = DType::kFloat32; - - if (return_max_logit) { - Tensor *output_Max = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_Max->data.dptr = nullptr; - if ((q_format == NVTE_QKV_Format::NVTE_THD && cudnn_runtime_version >= 90600) && - (sm_arch_ != 120)) { - output_Max->data.shape = {num_tokens_q, num_attn_heads, 1}; - } else { - output_Max->data.shape = {batch, num_attn_heads, max_seqlen_q, 1}; - } - output_Max->data.dtype = DType::kFloat32; - } - - Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_rng_state->data.dptr = nullptr; - output_rng_state->data.shape = {2}; - output_rng_state->data.dtype = DType::kInt64; - - if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { - Tensor *output_bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_bias->data.dptr = nullptr; - output_bias->data.shape = {bias_b, bias_h, bias_sq, bias_skv}; - output_bias->data.dtype = QKV_type; - } - - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - Tensor *output_softmax_offset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_softmax_offset->data.dptr = nullptr; - output_softmax_offset->data.shape = {1, num_attn_heads, 1, 1}; - output_softmax_offset->data.dtype = DType::kFloat32; - } - - Aux_CTX_Tensors->size = i; - } else if (Aux_CTX_Tensors->size >= 2) { - Tensor *output_S = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - devPtrS1 = output_S->data.dptr; - - if (return_max_logit) { - Tensor *output_Max = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - devPtrS2 = output_Max->data.dptr; - } - Tensor *output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_rng_state->data.dptr = rng_state->data.dptr; - if ((bias_type != NVTE_NO_BIAS) && (bias_type != NVTE_ALIBI)) { - Tensor *output_bias = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_bias->data.dptr = devPtrBias; - } - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - Tensor *output_softmax_offset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_softmax_offset->data.dptr = devPtrSoftmaxOffset; - } - } else { - NVTE_ERROR("Unexpected Aux_CTX_Tensors->size."); - } - - void *devPtrDropoutSeed = rng_state->data.dptr; - void *devPtrDropoutOffset = - reinterpret_cast(reinterpret_cast(rng_state->data.dptr) + 1); - - size_t workspace_size = 0; - - fused_attn_arbitrary_seqlen_fwd_impl( - batch, num_attn_heads, num_gqa_groups, max_seqlen_q, max_seqlen_kv, head_dim_qk, head_dim_v, - max_batch_size, max_tokens_q, max_tokens_kv, num_pages_k, num_pages_v, page_size_k, - page_size_v, max_pages_per_seq_k, max_pages_per_seq_v, bias_b, bias_h, bias_sq, bias_skv, - is_training, return_max_logit, attn_scale, p_dropout, qkv_layout, o_format, bias_type, - mask_type, softmax_type, window_size_left, window_size_right, bottom_right_diagonal, devPtrQ, - devPtrK, devPtrV, devPtrBias, devPtrSoftmaxOffset, devPtrS1, devPtrS2, devPtrO, - devPtrDropoutSeed, devPtrDropoutOffset, devPtrCuSeqlensQ, devPtrCuSeqlensKV, devPtrPageTableK, - devPtrPageTableV, devPtrSeqOffsetsQ, devPtrSeqOffsetsKV, get_cudnn_fe_dtype(QKV_type), - workspace->data.dptr, &workspace_size, stream, handle); - - if (workspace_size > 0) { - if (workspace->data.dptr == nullptr) { - workspace->data.shape = {workspace_size}; - workspace->data.dtype = DType::kByte; - return; - } - } else if (workspace_size == 0) { - workspace->data.shape = {1}; - workspace->data.dtype = DType::kByte; - return; - } else { - NVTE_ERROR("Unexpected workspace_size."); - } -} - -void fused_attn_arbitrary_seqlen_bwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, float attn_scale, float p_dropout, NVTE_QKV_Layout qkv_layout, - NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal, - bool deterministic, const Tensor *input_Q, const Tensor *input_K, const Tensor *input_V, - const Tensor *input_O, const Tensor *input_dO, const Tensor *input_Bias, - const Tensor *input_SoftmaxOffset, Tensor *output_S, Tensor *output_dQ, Tensor *output_dK, - Tensor *output_dV, Tensor *output_dBias, Tensor *output_dSoftmaxOffset, - const Tensor *cu_seqlens_q, const Tensor *cu_seqlens_kv, const Tensor *cu_seqlens_q_padded, - const Tensor *cu_seqlens_kv_padded, const Tensor *rng_state, Tensor *workspace, - cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - const auto QKV_type = input_Q->data.dtype; - void *devPtrQ = input_Q->data.dptr; - void *devPtrK = input_K->data.dptr; - void *devPtrV = input_V->data.dptr; - void *devPtrO = input_O->data.dptr; - void *devPtrdO = input_dO->data.dptr; - void *devPtrBias = nullptr; - void *devPtrdBias = nullptr; - size_t bias_b = 0; - size_t bias_h = 0; - size_t bias_sq = 0; - size_t bias_skv = 0; - if ((bias_type != NVTE_Bias_Type::NVTE_NO_BIAS) && (bias_type != NVTE_Bias_Type::NVTE_ALIBI)) { - devPtrBias = input_Bias->data.dptr; - devPtrdBias = output_dBias->data.dptr; - bias_b = output_dBias->data.shape[0]; - bias_h = output_dBias->data.shape[1]; - bias_sq = output_dBias->data.shape[2]; - bias_skv = output_dBias->data.shape[3]; - } - - size_t max_batch_size = 0; - size_t max_tokens_q = 0; - size_t max_tokens_kv = 0; - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - if (q_format == NVTE_QKV_Format::NVTE_THD || kv_format == NVTE_QKV_Format::NVTE_THD) { - max_batch_size = get_max_batch_size(batch); - } - if (q_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_q = get_max_tokens(num_tokens_q); - } - if (kv_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_kv = get_max_tokens(num_tokens_kv); - } - - void *devPtrdQ = output_dQ->data.dptr; - void *devPtrdK = output_dK->data.dptr; - void *devPtrdV = output_dV->data.dptr; - void *devPtrSoftmaxStats = nullptr; - devPtrSoftmaxStats = output_S->data.dptr; - void *devPtrSoftmaxOffset = nullptr; - void *devPtrdSoftmaxOffset = nullptr; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - devPtrSoftmaxOffset = input_SoftmaxOffset->data.dptr; - devPtrdSoftmaxOffset = output_dSoftmaxOffset->data.dptr; - } - - void *devPtrCuSeqlensQ = cu_seqlens_q->data.dptr; - void *devPtrCuSeqlensKV = cu_seqlens_kv->data.dptr; - void *devPtrSeqOffsetsQ = cu_seqlens_q_padded->data.dptr; - void *devPtrSeqOffsetsKV = cu_seqlens_kv_padded->data.dptr; - - void *devPtrDropoutSeed = rng_state->data.dptr; - void *devPtrDropoutOffset = - reinterpret_cast(reinterpret_cast(rng_state->data.dptr) + 1); - - size_t workspace_size = 0; - - fused_attn_arbitrary_seqlen_bwd_impl( - batch, num_attn_heads, num_gqa_groups, max_seqlen_q, max_seqlen_kv, head_dim_qk, head_dim_v, - max_batch_size, max_tokens_q, max_tokens_kv, bias_b, bias_h, bias_sq, bias_skv, attn_scale, - p_dropout, qkv_layout, o_format, do_format, dqkv_layout, bias_type, mask_type, softmax_type, - window_size_left, window_size_right, bottom_right_diagonal, deterministic, devPtrQ, devPtrK, - devPtrV, devPtrO, devPtrSoftmaxStats, devPtrBias, devPtrSoftmaxOffset, devPtrdQ, devPtrdK, - devPtrdV, devPtrdO, devPtrdBias, devPtrdSoftmaxOffset, devPtrDropoutSeed, devPtrDropoutOffset, - devPtrCuSeqlensQ, devPtrCuSeqlensKV, devPtrSeqOffsetsQ, devPtrSeqOffsetsKV, - get_cudnn_fe_dtype(QKV_type), workspace->data.dptr, &workspace_size, stream, handle); - - if (workspace_size > 0) { - if (workspace->data.dptr == nullptr) { - workspace->data.shape = {workspace_size}; - workspace->data.dtype = DType::kByte; - return; - } - } else if (workspace_size == 0) { - workspace->data.shape = {1}; - workspace->data.dtype = DType::kByte; - return; - } else { - NVTE_ERROR("Unexpected workspace_size."); - } -} -} // namespace transformer_engine diff --git a/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.h b/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.h deleted file mode 100644 index 8f79b5bb4a8..00000000000 --- a/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.h +++ /dev/null @@ -1,52 +0,0 @@ -/************************************************************************* - * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * - * See LICENSE for license information. - ************************************************************************/ - -/*! \file fused_attn_arbitrary_seqlen.h - * \brief Functions for fused attention with seqlen > 512 - */ - -#ifndef TRANSFORMER_ENGINE_COMMON_FUSED_ATTN_FUSED_ATTN_ARBITRARY_SEQLEN_H_ -#define TRANSFORMER_ENGINE_COMMON_FUSED_ATTN_FUSED_ATTN_ARBITRARY_SEQLEN_H_ - -#include - -#include "common/common.h" -#include "transformer_engine/fused_attn.h" - -namespace transformer_engine { -void fused_attn_arbitrary_seqlen_fwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, size_t num_pages_k, size_t num_pages_v, size_t page_size_k, - size_t page_size_v, size_t max_pages_per_seq_k, size_t max_pages_per_seq_v, bool is_training, - bool return_max_logit, float attn_scale, float p_dropout, NVTE_QKV_Layout qkv_layout, - NVTE_QKV_Format o_format, NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, - NVTE_Softmax_Type softmax_type, int64_t window_size_left, int64_t window_size_right, - bool bottom_right_diagonal, const Tensor *input_Q, const Tensor *input_K, const Tensor *input_V, - const Tensor *input_Bias, const Tensor *input_SoftmaxOffset, Tensor *output_O, - NVTETensorPack *Aux_CTX_Tensors, const Tensor *cu_seqlens_q, const Tensor *cu_seqlens_kv, - const Tensor *cu_seqlens_q_padded, const Tensor *cu_seqlens_kv_padded, - const Tensor *page_table_k, const Tensor *page_table_v, const Tensor *rng_state, - Tensor *workspace, cudaStream_t stream, cudnnHandle_t handle); - -void fused_attn_arbitrary_seqlen_bwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, float attn_scale, float p_dropout, NVTE_QKV_Layout qkv_layout, - NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal, - bool deterministic, const Tensor *input_Q, const Tensor *input_K, const Tensor *input_V, - const Tensor *input_O, const Tensor *input_dO, const Tensor *input_Bias, - const Tensor *input_SoftmaxOffset, Tensor *output_S, Tensor *output_dQ, Tensor *output_dK, - Tensor *output_dV, Tensor *output_dBias, Tensor *output_dSoftmaxOffset, - const Tensor *cu_seqlens_q, const Tensor *cu_seqlens_kv, const Tensor *cu_seqlens_q_padded, - const Tensor *cu_seqlens_kv_padded, const Tensor *rng_state, Tensor *workspace, - cudaStream_t stream, cudnnHandle_t handle); - -} // namespace transformer_engine - -#endif // TRANSFORMER_ENGINE_COMMON_FUSED_ATTN_FUSED_ATTN_ARBITRARY_SEQLEN_H_ diff --git a/transformer_engine/common/fused_attn/fused_attn_fp8.cu b/transformer_engine/common/fused_attn/fused_attn_fp8.cu deleted file mode 100644 index 2ef0ac33938..00000000000 --- a/transformer_engine/common/fused_attn/fused_attn_fp8.cu +++ /dev/null @@ -1,1711 +0,0 @@ -/************************************************************************* - * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * - * See LICENSE for license information. - ************************************************************************/ - -#include "../common.h" -#include "../cudnn_utils.h" -#include "../util/cuda_runtime.h" -#include "../util/system.h" -#include "fused_attn_fp8.h" -#include "utils.h" - -namespace transformer_engine { -namespace fused_attn { - -using namespace transformer_engine; - -constexpr size_t kFP8THDRaggedCudnnVersion = 92300; - -// fused attention FWD FP8 with FE 1.0+ -void fused_attn_fp8_fwd_impl( - int64_t b, int64_t h, int64_t hg, int64_t s_q, int64_t s_kv, int64_t d_qk, int64_t d_v, - int64_t max_b, int64_t max_t_q, int64_t max_t_kv, bool is_training, float scaling_factor, - float dropout_probability, NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - int64_t window_size_left, int64_t window_size_right, bool bottom_right_diagonal, void* devPtrQ, - void* devPtrK, void* devPtrV, void* devPtrSoftmaxOffset, void* devPtrM, void* devPtrO, - void* devPtrDescaleQ, void* devPtrDescaleK, void* devPtrDescaleV, void* devPtrDescaleS, - void* devPtrScaleS, void* devPtrScaleO, void* devPtrAmaxO, void* devPtrAmaxS, - void* devPtrcuSeqlensQ, void* devPtrcuSeqlensKV, void* devPtrSeqOffsetsQ, - void* devPtrSeqOffsetsKV, void* devPtrDropoutSeed, void* devPtrDropoutOffset, - cudnn_frontend::DataType_t qkv_tensor_type, cudnn_frontend::DataType_t o_tensor_type, - NVTEScalingMode scaling_mode, NVTE_QKV_Format qkv_scale_inv_format, void* workspace, - size_t* workspace_size, cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - const auto cudnn_runtime_version = cudnnGetVersion(); - bool is_bias = (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS); - bool is_alibi = (bias_type == NVTE_Bias_Type::NVTE_ALIBI); - bool is_causal = ((mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - bool is_padding = ((mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - bool is_dropout = (is_training && dropout_probability != 0.0f); - bool is_softmax_offset = (softmax_type != NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX); - auto bias_b = b; - auto bias_h = h; - auto bias_sq = s_q; - auto bias_skv = s_kv; - NVTE_CHECK(~is_bias, "FP8 fused attention does not support pre/post_scale_bias yet!"); - NVTE_CHECK(~is_alibi, "FP8 fused attention does not support ALiBi yet!"); - bool is_delayed_scaling = (scaling_mode == NVTE_DELAYED_TENSOR_SCALING) && - (o_tensor_type == cudnn_frontend::DataType_t::FP8_E4M3 || - o_tensor_type == cudnn_frontend::DataType_t::FP8_E5M2); - bool is_current_scaling = (scaling_mode == NVTE_DELAYED_TENSOR_SCALING) && - (o_tensor_type == cudnn_frontend::DataType_t::HALF || - o_tensor_type == cudnn_frontend::DataType_t::BFLOAT16); - bool is_mxfp8 = (scaling_mode == NVTE_MXFP8_1D_SCALING) && - (o_tensor_type == cudnn_frontend::DataType_t::HALF || - o_tensor_type == cudnn_frontend::DataType_t::BFLOAT16); - NVTE_CHECK( - is_delayed_scaling || is_current_scaling || is_mxfp8, - "FP8 fused attention only supports FP8DelayedScaling or FP8CurrentScaling or MXFP8 recipes!"); - NVTE_CHECK(!is_mxfp8 || cudnn_runtime_version >= 92100, - "MXFP8 fused attention requires cuDNN 9.21.0 or later!"); - - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - bool is_ragged_q = (q_format == NVTE_QKV_Format::NVTE_THD); - bool is_ragged_kv = (kv_format == NVTE_QKV_Format::NVTE_THD); - const int device_id = cuda::current_device(); - const int sm_arch_ = cuda::sm_arch(device_id); - bool use_ragged_stats = - is_ragged_q && cudnn_runtime_version >= kFP8THDRaggedCudnnVersion && sm_arch_ != 120; - - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); - const DType ragged_offset_type = DType::kInt64; - - // Newer versions of cuDNN SDPA can accept sequence lengths directly as a cumulative - // tensor. Take advantage of this if possible to avoid the actual-seqlen conversion; - // THD inputs still use their separate ragged-offset tensors. - const bool use_cu_seqlens_directly = - // Frontend 1.26 supports fp8+cu_seqlens (for the C++ API). - // Note: For the Python API, 1.27 is required. - CUDNN_FRONTEND_VERSION >= 12600 && - // The frontend gates cu_seq_len support on min(compile-time, runtime) cuDNN - // version, so we'll do the same. - (CUDNN_VERSION >= 92500 && cudnn_runtime_version >= 92500) && - // This extra restriction is needed because cuDNN frontend doesn't yet allow - // the combination of dropout and stats generation for the fprop unified engine, - // so any such request would always get routed to the old composite SDPA engine - // (which doesn't support cu_seqlens). Remove this restriction when possible. - !is_dropout; - - int64_t actual_b = b; - if ((is_ragged_q || is_ragged_kv) && cudnn_runtime_version >= kFP8THDRaggedCudnnVersion) { - NVTE_CHECK(is_padding, "Ragged QKV input requires padding or padding_causal mask!"); - if (sm_arch_ != 120) { - // cuDNN reads the user's [actual_b+1] cu_seqlens buffers directly, so a quantized - // batch dimension would read out of bounds on the direct path. - if (!use_cu_seqlens_directly) { - b = max_b; - } - s_q = is_ragged_q ? max_t_q : s_q; - s_kv = is_ragged_kv ? max_t_kv : s_kv; - } - } - - try { - FADescriptor_v1 descriptor{b, - h, - hg, - s_q, - s_kv, - d_qk, - d_v, - 0, - 0, - 0, - 0, - 0, - 0, - bias_b, - bias_h, - bias_sq, - bias_skv, - scaling_factor, - is_training, - dropout_probability, - qkv_layout, - o_format, - NVTE_QKV_Format_NOT_SET, - NVTE_QKV_Layout_NOT_SET, - qkv_scale_inv_format, - NVTE_QKV_Format_NOT_SET, - bias_type, - mask_type, - softmax_type, - window_size_left, - window_size_right, - bottom_right_diagonal, - true, - qkv_tensor_type, - o_tensor_type, - cudnn_frontend::DataType_t::NOT_SET, - cudnn_frontend::DataType_t::NOT_SET, - false}; - - namespace fe = cudnn_frontend; - using graph_and_tensors = - std::tuple, - std::shared_ptr, // Q - std::shared_ptr, // K - std::shared_ptr, // V - std::shared_ptr, // descale_q - std::shared_ptr, // descale_k - std::shared_ptr, // descale_v - std::shared_ptr, // descale_s - std::shared_ptr, // scale_s - std::shared_ptr, // scale_o - std::shared_ptr, // attn_scale - std::shared_ptr, // O - std::shared_ptr, // amax_s - std::shared_ptr, // amax_o - std::shared_ptr, // Stats - std::shared_ptr, // bias - std::shared_ptr, // softmax_offset - std::shared_ptr, // seq_q - std::shared_ptr, // seq_kv - std::shared_ptr, // offset_q - std::shared_ptr, // offset_k - std::shared_ptr, // offset_v - std::shared_ptr, // offset_o - std::shared_ptr, // offset_stats - std::shared_ptr, // dropout_seed - std::shared_ptr>; // dropout_offset - - using CacheType = std::map; - static thread_local CacheType sdpa_fp8_fprop_cache; - - // Get plan from cache if cache is available, otherwise create one - auto get_graph = [&](CacheType& cache, const FADescriptor_v1& descriptor) -> graph_and_tensors { - // if hit, return - auto it = cache.find(descriptor); - if (it != cache.end()) { - auto graph = it->second; - return graph; - } - - // otherwise, build the op_graph and the plan. Then update cache - auto mha_graph = std::make_shared(); - mha_graph->set_io_data_type(qkv_tensor_type) - .set_intermediate_data_type(fe::DataType_t::FLOAT) - .set_compute_data_type(fe::DataType_t::FLOAT); - - std::shared_ptr Q, K, V, attn_scale; - std::shared_ptr descale_q, descale_k, descale_v; - std::shared_ptr descale_s, scale_s, scale_o; - std::shared_ptr bias, softmax_offset, seq_q, seq_kv; - std::shared_ptr offset_q, offset_k, offset_v, offset_o, - offset_stats; - std::shared_ptr dropout_seed, dropout_offset; - - // Q, K, V, attn_scale - std::vector q_strides(4), k_strides(4), v_strides(4); - generateMatrixStridesWithLayout(b, h, hg, s_q, s_kv, d_qk, d_v, q_strides.data(), - k_strides.data(), v_strides.data(), qkv_layout); - Q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Q") - .set_dim({b, h, s_q, d_qk}) - .set_stride(q_strides) - .set_data_type(qkv_tensor_type)); - if (is_ragged_q) { - offset_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_q") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - Q->set_ragged_offset(offset_q); - } - K = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("K") - .set_dim({b, hg, s_kv, d_qk}) - .set_stride(k_strides) - .set_data_type(qkv_tensor_type)); - V = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("V") - .set_dim({b, hg, s_kv, d_v}) - .set_stride(v_strides) - .set_data_type(qkv_tensor_type)); - if (is_ragged_kv) { - offset_k = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_k") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - offset_v = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_v") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - K->set_ragged_offset(offset_k); - V->set_ragged_offset(offset_v); - } - attn_scale = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("attn_scale") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_is_pass_by_value(true) - .set_data_type(fe::DataType_t::FLOAT)); - - // Descale_q, Descale_k, Descale_v, Descale_s, Scale_s, Scale_o - if (is_delayed_scaling || is_current_scaling) { - descale_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_q") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - descale_k = mha_graph->tensor_like(descale_q, "Descale_q"); - descale_v = mha_graph->tensor_like(descale_q, "Descale_v"); - descale_s = mha_graph->tensor_like(descale_q, "Descale_s"); - scale_s = mha_graph->tensor_like(descale_q, "Scale_s"); - if (is_delayed_scaling) { - scale_o = mha_graph->tensor_like(descale_q, "Scale_o"); - } - if (is_current_scaling) { - scale_o = mha_graph->tensor(1.0f); - } - } else if (is_mxfp8) { - NVTE_QKV_Format q_scale_inv_format = (qkv_scale_inv_format != NVTE_QKV_Format_NOT_SET) - ? qkv_scale_inv_format - : nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_scale_inv_format = (qkv_scale_inv_format != NVTE_QKV_Format_NOT_SET) - ? qkv_scale_inv_format - : nvte_get_kv_format(qkv_layout); - std::vector q_scale_strides(4); - std::vector k_scale_strides(4); - std::vector v_scale_strides(4); - auto padded = pad_s_d_for_mxfp8(s_q, s_kv, d_qk, d_v); - generateMatrixStridesWithFormat(b, h, padded.s_q_padded, padded.d_qk_scale_padded, - q_scale_strides.data(), q_scale_inv_format); - generateMatrixStridesWithFormat(b, hg, padded.s_kv_padded, padded.d_qk_scale_padded, - k_scale_strides.data(), kv_scale_inv_format); - generateMatrixStridesWithFormat(b, hg, padded.s_kv_scale_padded, padded.d_v_padded, - v_scale_strides.data(), kv_scale_inv_format); - descale_q = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_q") - .set_dim({b, h, padded.s_q_padded, padded.d_qk_scale_padded}) - .set_stride(q_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_k = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_k") - .set_dim({b, hg, padded.s_kv_padded, padded.d_qk_scale_padded}) - .set_stride(k_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_v = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_v") - .set_dim({b, hg, padded.s_kv_scale_padded, padded.d_v_padded}) - .set_stride(v_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - } - - fe::graph::SDPA_fp8_attributes sdpa_options; - sdpa_options = fe::graph::SDPA_fp8_attributes() - .set_name("sdpa_fp8") - .set_generate_stats(true) - .set_causal_mask(is_causal) - .set_attn_scale(attn_scale); - - fe::DiagonalAlignment_t const& diagonal_alignment = - bottom_right_diagonal ? fe::DiagonalAlignment_t::BOTTOM_RIGHT - : fe::DiagonalAlignment_t::TOP_LEFT; - sdpa_options.set_diagonal_alignment(diagonal_alignment); - - if (cudnn_runtime_version >= 92100) { - if (window_size_left != -1) { - sdpa_options.set_diagonal_band_left_bound(window_size_left + 1); - } - if (window_size_right != -1) { - sdpa_options.set_diagonal_band_right_bound(window_size_right); - } - } - - // sdpa_options.set_alibi_mask(is_alibi); - // if (is_bias) { - // bias = mha_graph->tensor(fe::graph::Tensor_attributes() - // .set_name("bias") - // .set_dim({bias_b, bias_h, bias_sq, bias_skv}) - // .set_stride({bias_h * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1})); - // sdpa_options.set_bias(bias); - // } - - if (is_padding) { - if (use_cu_seqlens_directly) { - // seq_q/seq_kv keep their tuple slots but hold (b+1)-shaped cu_seqlen tensors. - seq_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("cu_seq_len_q") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - seq_kv = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("cu_seq_len_kv") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_options.set_padding_mask(is_padding) - .set_cu_seq_len_q(seq_q) - .set_cu_seq_len_kv(seq_kv); - // cu_seq_len (and the ragged offset multiplier) are unified-engine-only. - // Pin the implementation so an unsupported config fails with the unified - // engine's specific error instead of auto-selection's generic failure. - sdpa_options.set_implementation(fe::AttentionImplementation_t::UNIFIED); - } else { - seq_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_q") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - seq_kv = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_kv") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_options.set_padding_mask(is_padding).set_seq_len_q(seq_q).set_seq_len_kv(seq_kv); - } - } - - if (is_dropout) { - dropout_seed = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Seed") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - dropout_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Offset") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - sdpa_options.set_dropout(dropout_probability, dropout_seed, dropout_offset); - } - - if (is_softmax_offset) { - softmax_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("softmax_offset") - .set_dim({1, h, 1, 1}) - .set_stride({h, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - sdpa_options.set_sink_token(softmax_offset); - } - - std::shared_ptr O, Stats, amax_s, amax_o; - if (is_delayed_scaling || is_current_scaling) { - auto outputs = mha_graph->sdpa_fp8(Q, K, V, descale_q, descale_k, descale_v, descale_s, - scale_s, scale_o, sdpa_options); - O = outputs[0]; - Stats = outputs[1]; - amax_s = outputs[2]; - amax_o = outputs[3]; - amax_s->set_output(true) - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT); - } else if (is_mxfp8) { - auto outputs = mha_graph->sdpa_fp8(Q, K, V, descale_q, descale_k, descale_v, sdpa_options); - O = outputs[0]; - Stats = outputs[1]; - amax_o = outputs[2]; - } - - std::vector o_strides(4); - generateMatrixStridesWithFormat(b, h, s_q, d_v, o_strides.data(), o_format); - O->set_output(true) - .set_dim({b, h, s_q, d_v}) - .set_stride(o_strides) - .set_data_type(o_tensor_type); - if (is_ragged_q) { - offset_o = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_o") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - O->set_ragged_offset(offset_o); - } - amax_o->set_output(!is_mxfp8) - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT); - - if (use_ragged_stats) { - offset_stats = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_stats") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - } - Stats->set_output(true).set_data_type(fe::DataType_t::FLOAT).set_dim({b, h, s_q, 1}); - if (use_ragged_stats) { - Stats->set_stride({h * s_q, 1, h, 1}).set_ragged_offset(offset_stats); - } else { - Stats->set_stride({h * s_q, s_q, 1, 1}); - } - - std::tuple, // Q - std::shared_ptr, // K - std::shared_ptr, // V - std::shared_ptr, // descale_q - std::shared_ptr, // descale_k - std::shared_ptr, // descale_v - std::shared_ptr, // descale_s - std::shared_ptr, // scale_s - std::shared_ptr, // scale_o - std::shared_ptr, // attn_scale - std::shared_ptr, // O - std::shared_ptr, // amax_s - std::shared_ptr> // amax_o - key_tensors_tuple = - is_mxfp8 ? std::make_tuple(Q, K, V, descale_q, descale_k, descale_v, nullptr, nullptr, - nullptr, attn_scale, O, nullptr, amax_o) - : std::make_tuple(Q, K, V, descale_q, descale_k, descale_v, descale_s, - scale_s, scale_o, attn_scale, O, amax_s, amax_o); - auto Stats_tuple = std::make_tuple(Stats); - auto bias_tuple = is_bias ? std::make_tuple(bias) : std::make_tuple(nullptr); - auto softmax_offset_tuple = - is_softmax_offset ? std::make_tuple(softmax_offset) : std::make_tuple(nullptr); - auto padding_tuple = - is_padding ? std::make_tuple(seq_q, seq_kv) : std::make_tuple(nullptr, nullptr); - auto offset_q_tuple = is_ragged_q ? std::make_tuple(offset_q) : std::make_tuple(nullptr); - auto offset_kv_tuple = - is_ragged_kv ? std::make_tuple(offset_k, offset_v) : std::make_tuple(nullptr, nullptr); - auto offset_o_tuple = is_ragged_q ? std::make_tuple(offset_o) : std::make_tuple(nullptr); - auto offset_s_tuple = - use_ragged_stats ? std::make_tuple(offset_stats) : std::make_tuple(nullptr); - auto dropout_tuple = is_dropout ? std::make_tuple(dropout_seed, dropout_offset) - : std::make_tuple(nullptr, nullptr); - - NVTE_CHECK_CUDNN_FE(mha_graph->validate()); - NVTE_CHECK_CUDNN_FE(mha_graph->build_operation_graph(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->create_execution_plans({fe::HeurMode_t::A})); - NVTE_CHECK_CUDNN_FE(mha_graph->check_support(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->build_plans(handle)); - auto return_tuple = - std::tuple_cat(std::make_tuple(mha_graph), key_tensors_tuple, Stats_tuple, bias_tuple, - softmax_offset_tuple, padding_tuple, offset_q_tuple, offset_kv_tuple, - offset_o_tuple, offset_s_tuple, dropout_tuple); - cache.insert({descriptor, return_tuple}); - - return return_tuple; - }; - - auto [mha_graph, Q, K, V, descale_q, descale_k, descale_v, descale_s, scale_s, scale_o, - attn_scale, O, amax_s, amax_o, Stats, bias, softmax_offset, seq_q, seq_kv, offset_q, - offset_k, offset_v, offset_o, offset_stats, dropout_seed, dropout_offset] = - get_graph(sdpa_fp8_fprop_cache, descriptor); - - auto plan_workspace_size = alignTo<16>(mha_graph->get_workspace_size()); - const size_t num_bytes_per_seqlen = alignTo<16>(b * sizeof(int32_t)); - const size_t actual_seqlen_workspace_size = - (is_padding && !use_cu_seqlens_directly) ? 2 * num_bytes_per_seqlen : 0; - const size_t num_bytes_per_ragged_offset = - alignTo<16>(((b + 1) * typeToNumBits(ragged_offset_type)) / 8); - size_t seqlen_offsets_workspace_size = 0; - if (is_ragged_q || is_ragged_kv) { - size_t count = 2 * (static_cast(is_ragged_q) + static_cast(is_ragged_kv)); - if (use_ragged_stats) { - seqlen_offsets_workspace_size = (count + 1) * num_bytes_per_ragged_offset; - } else { - seqlen_offsets_workspace_size = count * num_bytes_per_ragged_offset; - } - } - - if (workspace == nullptr) { - *workspace_size = - plan_workspace_size + actual_seqlen_workspace_size + seqlen_offsets_workspace_size; - return; - } - - // cuDNN stream check needs to be moved here to support dummy kernel calls with - // null streams for sizing the cuDNN workspace. - NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); - - // Build variant pack - std::unordered_map, void*> variant_pack = { - {Q, devPtrQ}, - {K, devPtrK}, - {V, devPtrV}, - {descale_q, devPtrDescaleQ}, - {descale_k, devPtrDescaleK}, - {descale_v, devPtrDescaleV}, - {attn_scale, &scaling_factor}, - {O, devPtrO}, - {Stats, devPtrM}}; - - if (is_delayed_scaling) { - variant_pack[scale_o] = devPtrScaleO; - } - if (is_delayed_scaling || is_current_scaling) { - variant_pack[descale_s] = devPtrDescaleS; - variant_pack[scale_s] = devPtrScaleS; - variant_pack[amax_s] = devPtrAmaxS; - variant_pack[amax_o] = devPtrAmaxO; - } - - /* if (is_bias) { - variant_pack[bias] = devPtrBias; - } */ - - if (is_padding) { - if (use_cu_seqlens_directly) { - variant_pack[seq_q] = devPtrcuSeqlensQ; - variant_pack[seq_kv] = devPtrcuSeqlensKV; - } else { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block - 1) / nthreads_per_block; - void* devActualSeqlenQ = static_cast(workspace) + plan_workspace_size; - void* devActualSeqlenKV = static_cast(devActualSeqlenQ) + num_bytes_per_seqlen; - cu_seqlens_to_actual_seqlens<<>>( - actual_b, b, static_cast(devPtrcuSeqlensQ), - static_cast(devPtrcuSeqlensKV), static_cast(devActualSeqlenQ), - static_cast(devActualSeqlenKV)); - NVTE_CHECK_CUDA(cudaGetLastError()); - variant_pack[seq_q] = devActualSeqlenQ; - variant_pack[seq_kv] = devActualSeqlenKV; - } - } - - if (is_ragged_q || is_ragged_kv) { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block) / nthreads_per_block; - void* devOffsets = - static_cast(workspace) + plan_workspace_size + actual_seqlen_workspace_size; - void* devOffsetsQ = nullptr; - void* devOffsetsO = nullptr; - if (is_ragged_q) { - devOffsetsQ = devOffsets; - devOffsetsO = static_cast(devOffsetsQ) + num_bytes_per_ragged_offset; - } - void* devOffsetsK = nullptr; - void* devOffsetsV = nullptr; - if (is_ragged_kv) { - devOffsetsK = static_cast(devOffsets) + - static_cast(is_ragged_q) * 2 * num_bytes_per_ragged_offset; - devOffsetsV = static_cast(devOffsetsK) + num_bytes_per_ragged_offset; - } - void* devOffsetsS = nullptr; - if (use_ragged_stats) { - devOffsetsS = static_cast(devOffsets) + - (static_cast(is_ragged_q) + static_cast(is_ragged_kv)) * 2 * - num_bytes_per_ragged_offset; - } - const RaggedOffsetMultipliers offset_mults(layout_group, h, hg, d_qk, d_v); - cu_seqlens_padded_to_offsets<<>>( - offset_mults, actual_b, b, static_cast(devPtrSeqOffsetsQ), - static_cast(devPtrSeqOffsetsKV), ragged_offset_type, devOffsetsQ, devOffsetsK, - devOffsetsV, devOffsetsO, devOffsetsS); - NVTE_CHECK_CUDA(cudaGetLastError()); - if (is_ragged_q) { - variant_pack[offset_q] = devOffsetsQ; - variant_pack[offset_o] = devOffsetsO; - } - if (is_ragged_kv) { - variant_pack[offset_k] = devOffsetsK; - variant_pack[offset_v] = devOffsetsV; - } - if (use_ragged_stats) { - variant_pack[offset_stats] = devOffsetsS; - } - } - - if (is_dropout) { - variant_pack[dropout_seed] = devPtrDropoutSeed; - variant_pack[dropout_offset] = devPtrDropoutOffset; - } - - if (is_softmax_offset) { - variant_pack[softmax_offset] = devPtrSoftmaxOffset; - } - - NVTE_CHECK_CUDNN_FE(mha_graph->execute(handle, variant_pack, workspace)); - } catch (cudnn_frontend::cudnnException& e) { - NVTE_ERROR(e.what()); - } -} // NOLINT(readability/fn_size) - -// fused attention BWD FP8 with FE 1.0+ -void fused_attn_fp8_bwd_impl( - int64_t b, int64_t h, int64_t hg, int64_t s_q, int64_t s_kv, int64_t d_qk, int64_t d_v, - int64_t max_b, int64_t max_t_q, int64_t max_t_kv, float scaling_factor, - float dropout_probability, NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, NVTE_Bias_Type bias_type, - NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal, bool deterministic, void* devPtrQ, - void* devPtrK, void* devPtrV, void* devPtrM, void* devPtrO, void* devPtrdO, - void* devPtrSoftmaxOffset, void* devPtrdQ, void* devPtrdK, void* devPtrdV, - void* devPtrdSoftmaxOffset, void* devPtrDescaleQ, void* devPtrDescaleK, void* devPtrDescaleV, - void* devPtrDescaleO, void* devPtrDescaledO, void* devPtrDescaleS, void* devPtrDescaledP, - void* devPtrScaleS, void* devPtrScaledP, void* devPtrScaledQ, void* devPtrScaledK, - void* devPtrScaledV, void* devPtrAmaxdP, void* devPtrAmaxdQ, void* devPtrAmaxdK, - void* devPtrAmaxdV, void* devPtrQ_t, void* devPtrK_t, void* devPtrdO_f16, void* devPtrdO_t, - void* devPtrDescaleQ_t, void* devPtrDescaleK_t, void* devPtrDescaledO_t, void* devPtrcuSeqlensQ, - void* devPtrcuSeqlensKV, void* devPtrSeqOffsetsQ, void* devPtrSeqOffsetsKV, - void* devPtrDropoutSeed, void* devPtrDropoutOffset, cudnn_frontend::DataType_t qkv_tensor_type, - cudnn_frontend::DataType_t o_tensor_type, cudnn_frontend::DataType_t do_tensor_type, - cudnn_frontend::DataType_t dqkv_tensor_type, NVTEScalingMode scaling_mode, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, void* workspace, - size_t* workspace_size, cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - const auto cudnn_runtime_version = cudnnGetVersion(); - bool is_bias = (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS); - bool is_alibi = (bias_type == NVTE_Bias_Type::NVTE_ALIBI); - bool is_causal = ((mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - bool is_padding = ((mask_type == NVTE_Mask_Type::NVTE_PADDING_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK) || - (mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK)); - bool is_dropout = (dropout_probability != 0.0f); - bool is_softmax_offset = (softmax_type != NVTE_Softmax_Type::NVTE_VANILLA_SOFTMAX); - auto bias_b = b; - auto bias_h = h; - auto bias_sq = s_q; - auto bias_skv = s_kv; - NVTE_CHECK(~is_bias, "FP8 fused attention does not support pre/post_scale_bias yet!"); - NVTE_CHECK(~is_alibi, "FP8 fused attention does not support ALiBi yet!"); - bool is_delayed_scaling = (scaling_mode == NVTE_DELAYED_TENSOR_SCALING) && - (dqkv_tensor_type == cudnn_frontend::DataType_t::FP8_E4M3 || - dqkv_tensor_type == cudnn_frontend::DataType_t::FP8_E5M2); - bool is_current_scaling = (scaling_mode == NVTE_DELAYED_TENSOR_SCALING) && - (dqkv_tensor_type == cudnn_frontend::DataType_t::HALF || - dqkv_tensor_type == cudnn_frontend::DataType_t::BFLOAT16); - bool is_mxfp8 = (scaling_mode == NVTE_MXFP8_1D_SCALING) && - (dqkv_tensor_type == cudnn_frontend::DataType_t::HALF || - dqkv_tensor_type == cudnn_frontend::DataType_t::BFLOAT16); - NVTE_CHECK( - is_delayed_scaling || is_current_scaling || is_mxfp8, - "FP8 fused attention only supports FP8DelayedScaling or FP8CurrentScaling or MXFP8 recipes!"); - NVTE_CHECK(!is_mxfp8 || cudnn_runtime_version >= 92100, - "MXFP8 fused attention requires cuDNN 9.21.0 or later!"); - - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - bool is_ragged_q = (q_format == NVTE_QKV_Format::NVTE_THD); - bool is_ragged_kv = (kv_format == NVTE_QKV_Format::NVTE_THD); - const int device_id = cuda::current_device(); - const int sm_arch_ = cuda::sm_arch(device_id); - bool use_ragged_stats = - is_ragged_q && cudnn_runtime_version >= kFP8THDRaggedCudnnVersion && sm_arch_ != 120; - - NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout); - const DType ragged_offset_type = DType::kInt64; - - int64_t actual_b = b; - if ((is_ragged_q || is_ragged_kv) && cudnn_runtime_version >= kFP8THDRaggedCudnnVersion) { - NVTE_CHECK(is_padding, "Ragged QKV input requires padding or padding_causal mask!"); - if (sm_arch_ != 120) { - b = max_b; - s_q = is_ragged_q ? max_t_q : s_q; - s_kv = is_ragged_kv ? max_t_kv : s_kv; - } - } - - bool is_O_in_F16 = (o_tensor_type == cudnn_frontend::DataType_t::HALF || - o_tensor_type == cudnn_frontend::DataType_t::BFLOAT16); - - try { - FADescriptor_v1 descriptor{b, - h, - hg, - s_q, - s_kv, - d_qk, - d_v, - 0, - 0, - 0, - 0, - 0, - 0, - bias_b, - bias_h, - bias_sq, - bias_skv, - scaling_factor, - true, - dropout_probability, - qkv_layout, - o_format, - do_format, - dqkv_layout, - qkv_scale_inv_format, - do_scale_inv_format, - bias_type, - mask_type, - softmax_type, - window_size_left, - window_size_right, - bottom_right_diagonal, - deterministic, - qkv_tensor_type, - o_tensor_type, - do_tensor_type, - dqkv_tensor_type, - false}; - - namespace fe = cudnn_frontend; - using graph_and_tensors = - std::tuple, - std::shared_ptr, // Q - std::shared_ptr, // Q_t - std::shared_ptr, // K - std::shared_ptr, // K_t - std::shared_ptr, // V - std::shared_ptr, // O - std::shared_ptr, // Stats - std::shared_ptr, // dO - std::shared_ptr, // dO_t - std::shared_ptr, // dO_f16 - std::shared_ptr, // attn_scale - std::shared_ptr, // descale_q - std::shared_ptr, // descale_q_t - std::shared_ptr, // descale_k - std::shared_ptr, // descale_k_t - std::shared_ptr, // descale_v - std::shared_ptr, // descale_o - std::shared_ptr, // descale_dO - std::shared_ptr, // descale_dO_t - std::shared_ptr, // descale_s - std::shared_ptr, // descale_dP - std::shared_ptr, // scale_dQ - std::shared_ptr, // scale_dK - std::shared_ptr, // scale_dV - std::shared_ptr, // scale_s - std::shared_ptr, // scale_dP - std::shared_ptr, // dQ - std::shared_ptr, // dK - std::shared_ptr, // dV - std::shared_ptr, // amax_dQ - std::shared_ptr, // amax_dK - std::shared_ptr, // amax_dV - std::shared_ptr, // amax_dP - std::shared_ptr, // bias - std::shared_ptr, // dBias - std::shared_ptr, // softmax_offset - std::shared_ptr, // d_softmax_offset - std::shared_ptr, // seq_q - std::shared_ptr, // seq_kv - std::shared_ptr, // offset_q - std::shared_ptr, // offset_k - std::shared_ptr, // offset_v - std::shared_ptr, // offset_o - std::shared_ptr, // offset_stats - std::shared_ptr, // dropout_seed - std::shared_ptr>; // dropout_offset - - using CacheType = std::map; - static thread_local CacheType sdpa_fp8_bprop_cache; - - // Get plan from cache if cache is available, otherwise create one - auto get_graph = [&](CacheType& cache, const FADescriptor_v1& descriptor) -> graph_and_tensors { - // if hit, return - auto it = cache.find(descriptor); - if (it != cache.end()) { - auto graph = it->second; - return graph; - } - - // otherwise, build the op_graph and the plan. Then update cache - auto mha_graph = std::make_shared(); - - mha_graph->set_io_data_type(qkv_tensor_type) - .set_intermediate_data_type(fe::DataType_t::FLOAT) - .set_compute_data_type(fe::DataType_t::FLOAT); - - std::shared_ptr Q, Q_t, K, K_t, V, O, dO, dO_t, dO_f16, Stats, - attn_scale; - std::shared_ptr descale_q, descale_q_t, descale_k, descale_k_t, - descale_v; - std::shared_ptr descale_s, descale_o; - std::shared_ptr descale_dP, descale_dO, descale_dO_t; - std::shared_ptr scale_s, scale_dP; - std::shared_ptr scale_dQ, scale_dK, scale_dV; - std::shared_ptr bias, dBias, softmax_offset, d_softmax_offset; - std::shared_ptr seq_q, seq_kv; - std::shared_ptr offset_q, offset_k, offset_v, offset_o, - offset_stats; - std::shared_ptr dropout_seed, dropout_offset; - - // Q, K, V, O, dO, stats, attn_scale - std::vector q_strides(4), k_strides(4), v_strides(4), o_strides(4), dO_strides(4); - generateMatrixStridesWithLayout(b, h, hg, s_q, s_kv, d_qk, d_v, q_strides.data(), - k_strides.data(), v_strides.data(), qkv_layout); - generateMatrixStridesWithFormat(b, h, s_q, d_v, o_strides.data(), o_format); - generateMatrixStridesWithFormat(b, h, s_q, d_v, dO_strides.data(), do_format); - Q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Q") - .set_dim({b, h, s_q, d_qk}) - .set_stride(q_strides) - .set_data_type(qkv_tensor_type)); - if (is_ragged_q) { - offset_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_q") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - offset_o = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_o") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - Q->set_ragged_offset(offset_q); - } - K = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("K") - .set_dim({b, hg, s_kv, d_qk}) - .set_stride(k_strides) - .set_data_type(qkv_tensor_type)); - V = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("V") - .set_dim({b, hg, s_kv, d_v}) - .set_stride(v_strides) - .set_data_type(qkv_tensor_type)); - if (is_ragged_kv) { - offset_k = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_k") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - offset_v = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_v") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - K->set_ragged_offset(offset_k); - V->set_ragged_offset(offset_v); - } - O = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("O") - .set_dim({b, h, s_q, d_v}) - .set_stride(o_strides) - .set_data_type(o_tensor_type)); - if (is_ragged_q) { - O->set_ragged_offset(offset_o); - } - dO = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("dO") - .set_dim({b, h, s_q, d_v}) - .set_stride(dO_strides) - .set_data_type(do_tensor_type)); - if (is_ragged_q) { - dO->set_ragged_offset(offset_o); - } - if (use_ragged_stats) { - offset_stats = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("offset_stats") - .set_dim({b + 1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(get_cudnn_fe_dtype(ragged_offset_type))); - } - Stats = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Stats") - .set_dim({b, h, s_q, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - if (use_ragged_stats) { - Stats->set_stride({h * s_q, 1, h, 1}).set_ragged_offset(offset_stats); - } else { - Stats->set_stride({h * s_q, s_q, 1, 1}); - } - attn_scale = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("attn_scale") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_is_pass_by_value(true) - .set_data_type(fe::DataType_t::FLOAT)); - - // Descale_q, Descale_k, Descale_v, Descale_s, Scale_s, Descale_dP, Scale_dP, Descale_o, Descale_dO, Scale_dQ, Scale_dK, Scale_dV - if (is_delayed_scaling || is_current_scaling) { - descale_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_q") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - descale_k = mha_graph->tensor_like(descale_q, "Descale_q"); - descale_v = mha_graph->tensor_like(descale_q, "Descale_v"); - descale_s = mha_graph->tensor_like(descale_q, "Descale_s"); - scale_s = mha_graph->tensor_like(descale_q, "Scale_s"); - descale_dP = mha_graph->tensor_like(descale_q, "Descale_dP"); - scale_dP = mha_graph->tensor_like(descale_q, "Scale_dP"); - if (is_current_scaling && is_O_in_F16) { - descale_o = mha_graph->tensor(1.0f); - } else { - descale_o = mha_graph->tensor_like(descale_q, "Descale_O"); - } - descale_dO = mha_graph->tensor_like(descale_q, "Descale_dO"); - if (is_delayed_scaling) { - scale_dQ = mha_graph->tensor_like(descale_q, "Scale_dQ"); - scale_dK = mha_graph->tensor_like(descale_q, "Scale_dK"); - scale_dV = mha_graph->tensor_like(descale_q, "Scale_dV"); - } - if (is_current_scaling) { - scale_dQ = mha_graph->tensor(1.0f); - scale_dK = mha_graph->tensor(1.0f); - scale_dV = mha_graph->tensor(1.0f); - } - } else if (is_mxfp8) { - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - NVTE_QKV_Format q_scale_inv_format = - (qkv_scale_inv_format != NVTE_QKV_Format_NOT_SET) ? qkv_scale_inv_format : q_format; - NVTE_QKV_Format kv_scale_inv_format = - (qkv_scale_inv_format != NVTE_QKV_Format_NOT_SET) ? qkv_scale_inv_format : kv_format; - NVTE_QKV_Format do_scale_format_ = - (do_scale_inv_format != NVTE_QKV_Format_NOT_SET) ? do_scale_inv_format : do_format; - // Q_t, K_t, dO_t, dO_f16 - std::vector q_t_strides(4), k_t_strides(4), dO_t_strides(4); - generateMatrixStridesWithFormat(b, h, s_q, d_qk, q_t_strides.data(), q_format); - generateMatrixStridesWithFormat(b, hg, s_kv, d_qk, k_t_strides.data(), kv_format); - generateMatrixStridesWithFormat(b, h, s_q, d_v, dO_t_strides.data(), do_format); - Q_t = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Q_t") - .set_dim({b, h, s_q, d_qk}) - .set_stride(q_t_strides) - .set_data_type(qkv_tensor_type)); - K_t = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("K_t") - .set_dim({b, hg, s_kv, d_qk}) - .set_stride(k_t_strides) - .set_data_type(qkv_tensor_type)); - dO_t = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("dO_t") - .set_dim({b, h, s_q, d_v}) - .set_stride(dO_t_strides) - .set_data_type(do_tensor_type)); - dO_f16 = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("dO_f16") - .set_dim({b, h, s_q, d_v}) - .set_stride(dO_strides) - .set_data_type(o_tensor_type)); - // Descale_q, Descale_q_t, Descale_k, Descale_k_t, Descale_v, Descale_dO, Descale_dO_t - auto padded = pad_s_d_for_mxfp8(s_q, s_kv, d_qk, d_v); - std::vector q_scale_strides(4), q_t_scale_strides(4), k_scale_strides(4), - k_t_scale_strides(4), v_scale_strides(4), dO_scale_strides(4), dO_t_scale_strides(4); - generateMatrixStridesWithFormat(b, h, padded.s_q_padded, padded.d_qk_scale_padded, - q_scale_strides.data(), q_scale_inv_format); - generateMatrixStridesWithFormat(b, h, padded.s_q_scale_padded, padded.d_qk_padded, - q_t_scale_strides.data(), q_scale_inv_format); - generateMatrixStridesWithFormat(b, hg, padded.s_kv_padded, padded.d_qk_scale_padded, - k_scale_strides.data(), kv_scale_inv_format); - generateMatrixStridesWithFormat(b, hg, padded.s_kv_scale_padded, padded.d_qk_padded, - k_t_scale_strides.data(), kv_scale_inv_format); - generateMatrixStridesWithFormat(b, hg, padded.s_kv_padded, padded.d_v_scale_padded, - v_scale_strides.data(), kv_scale_inv_format); - generateMatrixStridesWithFormat(b, h, padded.s_q_padded, padded.d_v_scale_padded, - dO_scale_strides.data(), do_scale_format_); - generateMatrixStridesWithFormat(b, h, padded.s_q_scale_padded, padded.d_v_padded, - dO_t_scale_strides.data(), do_scale_format_); - descale_q = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_q") - .set_dim({b, h, padded.s_q_padded, padded.d_qk_scale_padded}) - .set_stride(q_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_q_t = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_q_t") - .set_dim({b, h, padded.s_q_scale_padded, padded.d_qk_padded}) - .set_stride(q_t_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_k = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_k") - .set_dim({b, hg, padded.s_kv_padded, padded.d_qk_scale_padded}) - .set_stride(k_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_k_t = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_k_t") - .set_dim({b, hg, padded.s_kv_scale_padded, padded.d_qk_padded}) - .set_stride(k_t_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_v = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_v") - .set_dim({b, hg, padded.s_kv_padded, padded.d_v_scale_padded}) - .set_stride(v_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_dO = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_dO") - .set_dim({b, h, padded.s_q_padded, padded.d_v_scale_padded}) - .set_stride(dO_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - descale_dO_t = - mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Descale_dO_t") - .set_dim({b, h, padded.s_q_scale_padded, padded.d_v_padded}) - .set_stride(dO_t_scale_strides) - .set_data_type(fe::DataType_t::FP8_E8M0) - .set_reordering_type(fe::TensorReordering_t::F8_128x4)); - } - - fe::graph::SDPA_fp8_backward_attributes sdpa_backward_options; - sdpa_backward_options = fe::graph::SDPA_fp8_backward_attributes() - .set_name("sdpa_fp8_backward") - .set_causal_mask(is_causal) - .set_attn_scale(attn_scale); - - fe::DiagonalAlignment_t const& diagonal_alignment = - bottom_right_diagonal ? fe::DiagonalAlignment_t::BOTTOM_RIGHT - : fe::DiagonalAlignment_t::TOP_LEFT; - sdpa_backward_options.set_diagonal_alignment(diagonal_alignment); - - if (cudnn_runtime_version >= 92100) { - if (window_size_left != -1) { - sdpa_backward_options.set_diagonal_band_left_bound(window_size_left + 1); - } - if (window_size_right != -1) { - sdpa_backward_options.set_diagonal_band_right_bound(window_size_right); - } - } - - // sdpa_backward_options.set_alibi_mask(is_alibi); - - // if (is_bias) { - // bias = mha_graph->tensor(fe::graph::Tensor_attributes() - // .set_name("bias") - // .set_dim({bias_b, bias_h, bias_sq, bias_skv}) - // .set_stride({bias_h * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1})); - // dBias = mha_graph->tensor(fe::graph::Tensor_attributes() - // .set_name("dBias") - // .set_dim({bias_b, bias_h, bias_sq, bias_skv}) - // .set_stride({bias_h * bias_sq * bias_skv, bias_sq * bias_skv, bias_skv, 1})); - // sdpa_backward_options.set_bias(bias); - // bias shapes [1, 1, s, s], [b, 1, s, s], [b, h, s, s], [1, h, s, s] are supported for dbias calculation - // bias shape [1, 1, 1, s] is not supported for dbias calculation as of cuDNN 9.18 - // if (!((bias_b == 1) && (bias_h == 1) && (bias_sq == 1))) { - // sdpa_backward_options.set_dbias(dBias); - // } - // } - - if (cudnn_runtime_version >= 91900) { - sdpa_backward_options.set_deterministic_algorithm(deterministic); - } - - if (is_padding) { - seq_q = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_q") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - seq_kv = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("seq_kv") - .set_dim({b, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT32)); - sdpa_backward_options.set_padding_mask(is_padding) - .set_seq_len_q(seq_q) - .set_seq_len_kv(seq_kv); - } - - if (is_dropout) { - dropout_seed = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Seed") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - dropout_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("Offset") - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::INT64)); - sdpa_backward_options.set_dropout(dropout_probability, dropout_seed, dropout_offset); - } - - if (is_softmax_offset) { - softmax_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("softmax_offset") - .set_dim({1, h, 1, 1}) - .set_stride({h, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - sdpa_backward_options.set_sink_token(softmax_offset); - d_softmax_offset = mha_graph->tensor(fe::graph::Tensor_attributes() - .set_name("d_softmax_offset") - .set_dim({1, h, 1, 1}) - .set_stride({h, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT)); - sdpa_backward_options.set_dsink_token(d_softmax_offset); - } - - std::shared_ptr dQ, dK, dV, amax_dQ, amax_dK, amax_dV, amax_dP; - if (is_delayed_scaling || is_current_scaling) { - std::tie(dQ, dK, dV, amax_dQ, amax_dK, amax_dV, amax_dP) = - std::apply([](const auto&... elems) { return std::make_tuple(elems...); }, - mha_graph->sdpa_fp8_backward(Q, K, V, O, dO, Stats, descale_q, descale_k, - descale_v, descale_o, descale_dO, descale_s, - descale_dP, scale_s, scale_dQ, scale_dK, - scale_dV, scale_dP, sdpa_backward_options)); - } else if (is_mxfp8) { - std::tie(dQ, dK, dV, amax_dQ, amax_dK, amax_dV) = std::apply( - [](const auto&... elems) { return std::make_tuple(elems...); }, - mha_graph->sdpa_fp8_backward(Q, Q_t, K, K_t, V, O, dO_f16, dO, dO_t, Stats, descale_q, - descale_q_t, descale_k, descale_k_t, descale_v, descale_dO, - descale_dO_t, sdpa_backward_options)); - } - std::vector dq_strides(4), dk_strides(4), dv_strides(4); - generateMatrixStridesWithLayout(b, h, hg, s_q, s_kv, d_qk, d_v, dq_strides.data(), - dk_strides.data(), dv_strides.data(), dqkv_layout); - dQ->set_output(true) - .set_dim({b, h, s_q, d_qk}) - .set_stride(dq_strides) - .set_data_type(dqkv_tensor_type); - if (is_ragged_q) { - dQ->set_ragged_offset(offset_q); - } - dK->set_output(true) - .set_dim({b, hg, s_kv, d_qk}) - .set_stride(dk_strides) - .set_data_type(dqkv_tensor_type); - dV->set_output(true) - .set_dim({b, hg, s_kv, d_v}) - .set_stride(dv_strides) - .set_data_type(dqkv_tensor_type); - if (is_ragged_kv) { - dK->set_ragged_offset(offset_k); - dV->set_ragged_offset(offset_v); - } - amax_dQ->set_output(!is_mxfp8) - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT); - amax_dK->set_output(!is_mxfp8) - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT); - amax_dV->set_output(!is_mxfp8) - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT); - if (is_delayed_scaling || is_current_scaling) { - amax_dP->set_output(true) - .set_dim({1, 1, 1, 1}) - .set_stride({1, 1, 1, 1}) - .set_data_type(fe::DataType_t::FLOAT); - } - - std::tuple, // Q - std::shared_ptr, // K - std::shared_ptr, // V - std::shared_ptr, // O - std::shared_ptr, // Stats - std::shared_ptr, // dO - std::shared_ptr, // attn_scale - std::shared_ptr, // descale_q - std::shared_ptr, // descale_k - std::shared_ptr, // descale_v - std::shared_ptr, // descale_o - std::shared_ptr, // descale_dO - std::shared_ptr, // descale_s - std::shared_ptr, // descale_dP - std::shared_ptr, // scale_dQ - std::shared_ptr, // scale_dK - std::shared_ptr, // scale_dV - std::shared_ptr, // scale_s - std::shared_ptr, // scale_dP - std::shared_ptr, // dQ - std::shared_ptr, // dK - std::shared_ptr, // dV - std::shared_ptr, // amax_dQ - std::shared_ptr, // amax_dK - std::shared_ptr, // amax_dV - std::shared_ptr> // amax_dP - key_tensors_tuple = std::make_tuple( - Q, K, V, O, Stats, dO, attn_scale, descale_q, descale_k, descale_v, descale_o, - descale_dO, descale_s, descale_dP, scale_s, scale_dQ, scale_dK, scale_dV, scale_dP, - dQ, dK, dV, amax_dQ, amax_dK, amax_dV, amax_dP); - auto mxfp8_tensors_tuple = - is_mxfp8 ? std::make_tuple(Q_t, K_t, dO_f16, dO_t, descale_q_t, descale_k_t, descale_dO_t) - : std::make_tuple(nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr); - auto bias_tuple = is_bias ? std::make_tuple(bias, dBias) : std::make_tuple(nullptr, nullptr); - auto softmax_offset_tuple = is_softmax_offset - ? std::make_tuple(softmax_offset, d_softmax_offset) - : std::make_tuple(nullptr, nullptr); - auto padding_tuple = - is_padding ? std::make_tuple(seq_q, seq_kv) : std::make_tuple(nullptr, nullptr); - auto offset_q_tuple = is_ragged_q ? std::make_tuple(offset_q) : std::make_tuple(nullptr); - auto offset_kv_tuple = - is_ragged_kv ? std::make_tuple(offset_k, offset_v) : std::make_tuple(nullptr, nullptr); - auto offset_o_tuple = is_ragged_q ? std::make_tuple(offset_o) : std::make_tuple(nullptr); - auto offset_s_tuple = - use_ragged_stats ? std::make_tuple(offset_stats) : std::make_tuple(nullptr); - auto dropout_tuple = is_dropout ? std::make_tuple(dropout_seed, dropout_offset) - : std::make_tuple(nullptr, nullptr); - - NVTE_CHECK_CUDNN_FE(mha_graph->validate()); - NVTE_CHECK_CUDNN_FE(mha_graph->build_operation_graph(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->create_execution_plans({fe::HeurMode_t::A})); - NVTE_CHECK_CUDNN_FE(mha_graph->check_support(handle)); - NVTE_CHECK_CUDNN_FE(mha_graph->build_plans(handle)); - - auto return_tuple = - std::tuple_cat(std::make_tuple(mha_graph), key_tensors_tuple, mxfp8_tensors_tuple, - bias_tuple, softmax_offset_tuple, padding_tuple, offset_q_tuple, - offset_kv_tuple, offset_o_tuple, offset_s_tuple, dropout_tuple); - cache.insert({descriptor, return_tuple}); - - return return_tuple; - }; - auto [mha_graph, Q, K, V, O, Stats, dO, attn_scale, descale_q, descale_k, descale_v, descale_o, - descale_dO, descale_s, descale_dP, scale_s, scale_dQ, scale_dK, scale_dV, scale_dP, dQ, - dK, dV, amax_dQ, amax_dK, amax_dV, amax_dP, Q_t, K_t, dO_f16, dO_t, descale_q_t, - descale_k_t, descale_dO_t, bias, dBias, softmax_offset, d_softmax_offset, seq_q, seq_kv, - offset_q, offset_k, offset_v, offset_o, offset_stats, dropout_seed, dropout_offset] = - get_graph(sdpa_fp8_bprop_cache, descriptor); - - auto plan_workspace_size = alignTo<16>(mha_graph->get_workspace_size()); - const size_t num_bytes_per_seqlen = alignTo<16>(b * sizeof(int32_t)); - const size_t actual_seqlen_workspace_size = is_padding ? 2 * num_bytes_per_seqlen : 0; - const size_t num_bytes_per_ragged_offset = - alignTo<16>(((b + 1) * typeToNumBits(ragged_offset_type)) / 8); - size_t seqlen_offsets_workspace_size = 0; - if (is_ragged_q || is_ragged_kv) { - size_t count = 2 * (static_cast(is_ragged_q) + static_cast(is_ragged_kv)); - if (use_ragged_stats) { - seqlen_offsets_workspace_size = (count + 1) * num_bytes_per_ragged_offset; - } else { - seqlen_offsets_workspace_size = count * num_bytes_per_ragged_offset; - } - } - - if (workspace == nullptr) { - *workspace_size = - plan_workspace_size + actual_seqlen_workspace_size + seqlen_offsets_workspace_size; - return; - } - - // cuDNN stream check needs to be moved here to support dummy kernel calls with - // null streams for sizing the cuDNN workspace. - NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); - - // build variant pack - std::unordered_map, void*> variant_pack = { - {Q, devPtrQ}, - {K, devPtrK}, - {V, devPtrV}, - {O, devPtrO}, - {Stats, devPtrM}, - {dO, devPtrdO}, - {attn_scale, &scaling_factor}, - {descale_q, devPtrDescaleQ}, - {descale_k, devPtrDescaleK}, - {descale_v, devPtrDescaleV}, - {descale_dO, devPtrDescaledO}, - {dQ, devPtrdQ}, - {dK, devPtrdK}, - {dV, devPtrdV}, - }; - if (is_delayed_scaling || is_current_scaling) { - variant_pack[descale_s] = devPtrDescaleS; - variant_pack[descale_dP] = devPtrDescaledP; - variant_pack[scale_s] = devPtrScaleS; - variant_pack[scale_dP] = devPtrScaledP; - variant_pack[amax_dP] = devPtrAmaxdP; - variant_pack[amax_dQ] = devPtrAmaxdQ; - variant_pack[amax_dK] = devPtrAmaxdK; - variant_pack[amax_dV] = devPtrAmaxdV; - } - if (is_delayed_scaling || (is_current_scaling && !is_O_in_F16)) { - variant_pack[descale_o] = devPtrDescaleO; - } - if (is_delayed_scaling) { - variant_pack[scale_dQ] = devPtrScaledQ; - variant_pack[scale_dK] = devPtrScaledK; - variant_pack[scale_dV] = devPtrScaledV; - } - if (is_mxfp8) { - variant_pack[Q_t] = devPtrQ_t; - variant_pack[K_t] = devPtrK_t; - variant_pack[dO_f16] = devPtrdO_f16; - variant_pack[dO_t] = devPtrdO_t; - variant_pack[descale_q_t] = devPtrDescaleQ_t; - variant_pack[descale_k_t] = devPtrDescaleK_t; - variant_pack[descale_dO_t] = devPtrDescaledO_t; - } - - /* if (is_bias) { - variant_pack[bias] = devPtrBias; - if ((bias_b == 1) && (bias_h == h)) { - variant_pack[dBias] = devPtrdBias; - } else { - variant_pack[dBias] = nullptr; - } - } */ - - if (is_padding) { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block - 1) / nthreads_per_block; - void* devActualSeqlenQ = static_cast(workspace) + plan_workspace_size; - void* devActualSeqlenKV = static_cast(devActualSeqlenQ) + num_bytes_per_seqlen; - cu_seqlens_to_actual_seqlens<<>>( - actual_b, b, static_cast(devPtrcuSeqlensQ), - static_cast(devPtrcuSeqlensKV), static_cast(devActualSeqlenQ), - static_cast(devActualSeqlenKV)); - NVTE_CHECK_CUDA(cudaGetLastError()); - variant_pack[seq_q] = devActualSeqlenQ; - variant_pack[seq_kv] = devActualSeqlenKV; - } - - if (is_ragged_q || is_ragged_kv) { - constexpr size_t nthreads_per_block = 128; - const size_t grid = (b + nthreads_per_block) / nthreads_per_block; - void* devOffsets = - static_cast(workspace) + plan_workspace_size + actual_seqlen_workspace_size; - void* devOffsetsQ = nullptr; - void* devOffsetsO = nullptr; - if (is_ragged_q) { - devOffsetsQ = devOffsets; - devOffsetsO = static_cast(devOffsetsQ) + num_bytes_per_ragged_offset; - } - void* devOffsetsK = nullptr; - void* devOffsetsV = nullptr; - if (is_ragged_kv) { - devOffsetsK = static_cast(devOffsets) + - static_cast(is_ragged_q) * 2 * num_bytes_per_ragged_offset; - devOffsetsV = static_cast(devOffsetsK) + num_bytes_per_ragged_offset; - } - void* devOffsetsS = nullptr; - if (use_ragged_stats) { - devOffsetsS = static_cast(devOffsets) + - (static_cast(is_ragged_q) + static_cast(is_ragged_kv)) * 2 * - num_bytes_per_ragged_offset; - } - const RaggedOffsetMultipliers offset_mults(layout_group, h, hg, d_qk, d_v); - cu_seqlens_padded_to_offsets<<>>( - offset_mults, actual_b, b, static_cast(devPtrSeqOffsetsQ), - static_cast(devPtrSeqOffsetsKV), ragged_offset_type, devOffsetsQ, devOffsetsK, - devOffsetsV, devOffsetsO, devOffsetsS); - NVTE_CHECK_CUDA(cudaGetLastError()); - if (is_ragged_q) { - variant_pack[offset_q] = devOffsetsQ; - variant_pack[offset_o] = devOffsetsO; - } - if (is_ragged_kv) { - variant_pack[offset_k] = devOffsetsK; - variant_pack[offset_v] = devOffsetsV; - } - if (use_ragged_stats) { - variant_pack[offset_stats] = devOffsetsS; - } - } - - if (is_dropout) { - variant_pack[dropout_seed] = devPtrDropoutSeed; - variant_pack[dropout_offset] = devPtrDropoutOffset; - } - - if (is_softmax_offset) { - variant_pack[softmax_offset] = devPtrSoftmaxOffset; - variant_pack[d_softmax_offset] = devPtrdSoftmaxOffset; - } - - NVTE_CHECK_CUDNN_FE(mha_graph->execute(handle, variant_pack, workspace)); - } catch (cudnn_frontend::cudnnException& e) { - NVTE_ERROR(e.what()); - } -} // NOLINT(readability/fn_size) - -} // namespace fused_attn - -// fused attention FWD FP8 with separate Q, K, V -void fused_attn_fp8_fwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, bool is_training, float attn_scale, float p_dropout, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format qkv_scale_inv_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - size_t window_size_left, size_t window_size_right, bool bottom_right_diagonal, - const Tensor* input_Q, const Tensor* input_K, const Tensor* input_V, - const Tensor* input_SoftmaxOffset, Tensor* input_output_S, Tensor* output_O, - NVTETensorPack* Aux_CTX_Tensors, const Tensor* cu_seqlens_q, const Tensor* cu_seqlens_kv, - const Tensor* cu_seqlens_q_padded, const Tensor* cu_seqlens_kv_padded, const Tensor* rng_state, - Tensor* workspace, cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - void *devPtrQ = nullptr, *devPtrK = nullptr, *devPtrV = nullptr; - void *devPtrDescaleQ = nullptr, *devPtrDescaleK = nullptr, *devPtrDescaleV = nullptr; - void *devPtrO = nullptr, *devPtrAmaxO = nullptr, *devPtrScaleO = nullptr; - void *devPtrAmaxS = nullptr, *devPtrScaleS = nullptr, *devPtrDescaleS = nullptr; - devPtrQ = input_Q->data.dptr; - devPtrDescaleQ = input_Q->scale_inv.dptr; - devPtrK = input_K->data.dptr; - devPtrDescaleK = input_K->scale_inv.dptr; - devPtrO = output_O->data.dptr; - if (input_Q->scaling_mode == NVTE_DELAYED_TENSOR_SCALING) { - devPtrV = input_V->data.dptr; - devPtrDescaleV = input_V->scale_inv.dptr; - devPtrScaleO = output_O->scale.dptr; - devPtrAmaxS = input_output_S->amax.dptr; - devPtrScaleS = input_output_S->scale.dptr; - devPtrDescaleS = input_output_S->scale_inv.dptr; - devPtrAmaxO = output_O->amax.dptr; - } else if (input_Q->scaling_mode == NVTE_MXFP8_1D_SCALING) { - devPtrV = input_V->columnwise_data.dptr; - devPtrDescaleV = input_V->columnwise_scale_inv.dptr; - } - void* devPtrSoftmaxOffset = nullptr; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - devPtrSoftmaxOffset = input_SoftmaxOffset->data.dptr; - } - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - const auto cudnn_runtime_version = cudnnGetVersion(); - const int sm_arch_ = cuda::sm_arch(cuda::current_device()); - - void* devPtrSeqOffsetsQ = cu_seqlens_q_padded->data.dptr; - void* devPtrSeqOffsetsKV = cu_seqlens_kv_padded->data.dptr; - - size_t max_batch_size = 0; - size_t max_tokens_q = 0; - size_t max_tokens_kv = 0; - if (q_format == NVTE_QKV_Format::NVTE_THD || kv_format == NVTE_QKV_Format::NVTE_THD) { - max_batch_size = fused_attn::get_max_batch_size(batch); - } - if (q_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_q = fused_attn::get_max_tokens(num_tokens_q); - } - if (kv_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_kv = fused_attn::get_max_tokens(num_tokens_kv); - } - - void* devPtrM = nullptr; - if (Aux_CTX_Tensors->size == 0) { - int i = 0; - Tensor* output_M = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_M->data.dptr = nullptr; - // SM120 uses dense stats in the graph, so its allocation must remain [b, h, s_q, 1]. - if (q_format == NVTE_QKV_Format::NVTE_THD && - cudnn_runtime_version >= fused_attn::kFP8THDRaggedCudnnVersion && sm_arch_ != 120) { - output_M->data.shape = {num_tokens_q, num_attn_heads, 1}; - } else { - output_M->data.shape = {batch, num_attn_heads, max_seqlen_q, 1}; - } - output_M->data.dtype = DType::kFloat32; - Tensor* output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_rng_state->data.dptr = nullptr; - output_rng_state->data.shape = {2}; - output_rng_state->data.dtype = DType::kInt64; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - Tensor* output_softmax_offset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_softmax_offset->data.dptr = nullptr; - output_softmax_offset->data.shape = {1, num_attn_heads, 1, 1}; - output_softmax_offset->data.dtype = DType::kFloat32; - } - Aux_CTX_Tensors->size = i; - } else if (Aux_CTX_Tensors->size >= 2) { - int i = 0; - Tensor* output_M = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - devPtrM = output_M->data.dptr; - Tensor* output_rng_state = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_rng_state->data.dptr = rng_state->data.dptr; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - Tensor* output_softmax_offset = convertNVTETensorCheck(Aux_CTX_Tensors->tensors[i++]); - output_softmax_offset->data.dptr = devPtrSoftmaxOffset; - } - } else { - NVTE_ERROR("Unexpected Aux_CTX_Tensors->size."); - } - - void* devPtrcuSeqlensQ = - reinterpret_cast(reinterpret_cast(cu_seqlens_q->data.dptr)); - void* devPtrcuSeqlensKV = - reinterpret_cast(reinterpret_cast(cu_seqlens_kv->data.dptr)); - void* devPtrDropoutSeed = - reinterpret_cast(reinterpret_cast(rng_state->data.dptr)); - void* devPtrDropoutOffset = - reinterpret_cast(reinterpret_cast(rng_state->data.dptr) + 1); - - const DType QKV_type = input_Q->data.dtype; - const DType O_type = output_O->data.dtype; - size_t workspace_size = 0; - - NVTE_QKV_Format qkv_format = nvte_get_qkv_format(qkv_layout); - if ((qkv_format == NVTE_QKV_Format::NVTE_BSHD) || (qkv_format == NVTE_QKV_Format::NVTE_SBHD) || - (qkv_format == NVTE_QKV_Format::NVTE_BHSD) || (qkv_format == NVTE_QKV_Format::NVTE_THD)) { - fused_attn::fused_attn_fp8_fwd_impl( - batch, num_attn_heads, num_gqa_groups, max_seqlen_q, max_seqlen_kv, head_dim_qk, head_dim_v, - max_batch_size, max_tokens_q, max_tokens_kv, is_training, attn_scale, p_dropout, qkv_layout, - o_format, bias_type, mask_type, softmax_type, window_size_left, window_size_right, - bottom_right_diagonal, devPtrQ, devPtrK, devPtrV, devPtrSoftmaxOffset, devPtrM, devPtrO, - devPtrDescaleQ, devPtrDescaleK, devPtrDescaleV, devPtrDescaleS, devPtrScaleS, devPtrScaleO, - devPtrAmaxO, devPtrAmaxS, devPtrcuSeqlensQ, devPtrcuSeqlensKV, devPtrSeqOffsetsQ, - devPtrSeqOffsetsKV, devPtrDropoutSeed, devPtrDropoutOffset, get_cudnn_fe_dtype(QKV_type), - get_cudnn_fe_dtype(O_type), input_Q->scaling_mode, qkv_scale_inv_format, - workspace->data.dptr, &workspace_size, stream, handle); - } else { - NVTE_ERROR("FP8 fused attention only supports qkv_format=BSHD, SBHD, BHSD, or THD.\n"); - } - - if (workspace_size > 0) { - if (workspace->data.dptr == nullptr) { - workspace->data.shape = {workspace_size}; - workspace->data.dtype = DType::kByte; - return; - } - } else if (workspace_size == 0) { - workspace->data.shape = {1}; - workspace->data.dtype = DType::kByte; - return; - } -} -// fused attention BWD FP8 with separate Q, K, V -void fused_attn_fp8_bwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, float attn_scale, float p_dropout, NVTE_QKV_Layout qkv_layout, - NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - size_t window_size_left, size_t window_size_right, bool bottom_right_diagonal, - bool deterministic, const Tensor* input_Q, const Tensor* input_K, const Tensor* input_V, - const Tensor* input_O, const Tensor* input_dO, const Tensor* input_dO_f16, - const Tensor* input_M, const Tensor* input_S, const Tensor* input_SoftmaxOffset, - Tensor* input_output_dP, const Tensor* output_dQ, const Tensor* output_dK, - const Tensor* output_dV, Tensor* output_dSoftmaxOffset, const Tensor* cu_seqlens_q, - const Tensor* cu_seqlens_kv, const Tensor* cu_seqlens_q_padded, - const Tensor* cu_seqlens_kv_padded, const Tensor* rng_state, Tensor* workspace, - cudaStream_t stream, cudnnHandle_t handle) { - using namespace transformer_engine; - void* devPtrQ = input_Q->data.dptr; - void* devPtrK = input_K->data.dptr; - void* devPtrV = input_V->data.dptr; - void* devPtrDescaleQ = input_Q->scale_inv.dptr; - void* devPtrDescaleK = input_K->scale_inv.dptr; - void* devPtrDescaleV = input_V->scale_inv.dptr; - void *devPtrQ_t = nullptr, *devPtrK_t = nullptr, *devPtrDescaleQ_t = nullptr, - *devPtrDescaleK_t = nullptr; - if (input_Q->scaling_mode == NVTE_MXFP8_1D_SCALING) { - devPtrQ_t = input_Q->columnwise_data.dptr; - devPtrDescaleQ_t = input_Q->columnwise_scale_inv.dptr; - devPtrK_t = input_K->columnwise_data.dptr; - devPtrDescaleK_t = input_K->columnwise_scale_inv.dptr; - } - - void* devPtrO = input_O->data.dptr; - const DType O_type = input_O->data.dtype; - void* devPtrDescaleO = nullptr; - if (O_type == DType::kFloat8E4M3 || O_type == DType::kFloat8E5M2) { - devPtrDescaleO = input_O->scale_inv.dptr; - } - void* devPtrdO = input_dO->data.dptr; - void* devPtrDescaledO = input_dO->scale_inv.dptr; - void *devPtrdO_t = nullptr, *devPtrdO_f16 = nullptr, *devPtrDescaledO_t = nullptr; - if (input_dO->scaling_mode == NVTE_MXFP8_1D_SCALING) { - devPtrdO_t = input_dO->columnwise_data.dptr; - devPtrdO_f16 = input_dO_f16->data.dptr; - devPtrDescaledO_t = input_dO->columnwise_scale_inv.dptr; - } - - void* devPtrM = input_M->data.dptr; - - void *devPtrScaleS = nullptr, *devPtrDescaleS = nullptr, *devPtrAmaxdP = nullptr, - *devPtrScaledP = nullptr, *devPtrDescaledP = nullptr; - if (input_Q->scaling_mode == NVTE_DELAYED_TENSOR_SCALING) { - devPtrScaleS = input_S->scale.dptr; - devPtrDescaleS = input_S->scale_inv.dptr; - devPtrAmaxdP = input_output_dP->amax.dptr; - devPtrScaledP = input_output_dP->scale.dptr; - devPtrDescaledP = input_output_dP->scale_inv.dptr; - } - - void* devPtrSoftmaxOffset = nullptr; - void* devPtrdSoftmaxOffset = nullptr; - if (softmax_type != NVTE_VANILLA_SOFTMAX) { - devPtrSoftmaxOffset = input_SoftmaxOffset->data.dptr; - devPtrdSoftmaxOffset = output_dSoftmaxOffset->data.dptr; - } - - void* devPtrdQ = output_dQ->data.dptr; - void* devPtrdK = output_dK->data.dptr; - void* devPtrdV = output_dV->data.dptr; - void *devPtrAmaxdQ = nullptr, *devPtrAmaxdK = nullptr, *devPtrAmaxdV = nullptr, - *devPtrScaledQ = nullptr, *devPtrScaledK = nullptr, *devPtrScaledV = nullptr; - if (input_Q->scaling_mode == NVTE_DELAYED_TENSOR_SCALING) { - devPtrAmaxdQ = output_dQ->amax.dptr; - devPtrAmaxdK = output_dK->amax.dptr; - devPtrAmaxdV = output_dV->amax.dptr; - devPtrScaledQ = output_dQ->scale.dptr; - devPtrScaledK = output_dK->scale.dptr; - devPtrScaledV = output_dV->scale.dptr; - } - - NVTE_QKV_Format q_format = nvte_get_q_format(qkv_layout); - NVTE_QKV_Format kv_format = nvte_get_kv_format(qkv_layout); - - void* devPtrSeqOffsetsQ = cu_seqlens_q_padded->data.dptr; - void* devPtrSeqOffsetsKV = cu_seqlens_kv_padded->data.dptr; - - size_t max_batch_size = 0; - size_t max_tokens_q = 0; - size_t max_tokens_kv = 0; - if (q_format == NVTE_QKV_Format::NVTE_THD || kv_format == NVTE_QKV_Format::NVTE_THD) { - max_batch_size = fused_attn::get_max_batch_size(batch); - } - if (q_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_q = fused_attn::get_max_tokens(num_tokens_q); - } - if (kv_format == NVTE_QKV_Format::NVTE_THD) { - max_tokens_kv = fused_attn::get_max_tokens(num_tokens_kv); - } - - void* devPtrcuSeqlensQ = - reinterpret_cast(reinterpret_cast(cu_seqlens_q->data.dptr)); - void* devPtrcuSeqlensKV = - reinterpret_cast(reinterpret_cast(cu_seqlens_kv->data.dptr)); - void* devPtrDropoutSeed = - reinterpret_cast(reinterpret_cast(rng_state->data.dptr)); - void* devPtrDropoutOffset = - reinterpret_cast(reinterpret_cast(rng_state->data.dptr) + 1); - - const DType QKV_type = input_Q->data.dtype; - const DType dO_type = input_dO->data.dtype; - const DType dQKV_type = output_dQ->data.dtype; - size_t workspace_size = 0; - - NVTE_QKV_Format dqkv_format = nvte_get_qkv_format(dqkv_layout); - if ((dqkv_format == NVTE_QKV_Format::NVTE_BSHD) || (dqkv_format == NVTE_QKV_Format::NVTE_SBHD) || - (dqkv_format == NVTE_QKV_Format::NVTE_BHSD) || (dqkv_format == NVTE_QKV_Format::NVTE_THD)) { - fused_attn::fused_attn_fp8_bwd_impl( - batch, num_attn_heads, num_gqa_groups, max_seqlen_q, max_seqlen_kv, head_dim_qk, head_dim_v, - max_batch_size, max_tokens_q, max_tokens_kv, attn_scale, p_dropout, qkv_layout, o_format, - do_format, dqkv_layout, bias_type, mask_type, softmax_type, window_size_left, - window_size_right, bottom_right_diagonal, deterministic, devPtrQ, devPtrK, devPtrV, devPtrM, - devPtrO, devPtrdO, devPtrSoftmaxOffset, devPtrdQ, devPtrdK, devPtrdV, devPtrdSoftmaxOffset, - devPtrDescaleQ, devPtrDescaleK, devPtrDescaleV, devPtrDescaleO, devPtrDescaledO, - devPtrDescaleS, devPtrDescaledP, devPtrScaleS, devPtrScaledP, devPtrScaledQ, devPtrScaledK, - devPtrScaledV, devPtrAmaxdP, devPtrAmaxdQ, devPtrAmaxdK, devPtrAmaxdV, devPtrQ_t, devPtrK_t, - devPtrdO_f16, devPtrdO_t, devPtrDescaleQ_t, devPtrDescaleK_t, devPtrDescaledO_t, - devPtrcuSeqlensQ, devPtrcuSeqlensKV, devPtrSeqOffsetsQ, devPtrSeqOffsetsKV, - devPtrDropoutSeed, devPtrDropoutOffset, get_cudnn_fe_dtype(QKV_type), - get_cudnn_fe_dtype(O_type), get_cudnn_fe_dtype(dO_type), get_cudnn_fe_dtype(dQKV_type), - input_dO->scaling_mode, qkv_scale_inv_format, do_scale_inv_format, workspace->data.dptr, - &workspace_size, stream, handle); - } else { - NVTE_ERROR("FP8 fused attention only supports dqkv_format=BSHD, SBHD, BHSD, or THD.\n"); - } - - if (workspace_size > 0) { - if (workspace->data.dptr == nullptr) { - workspace->data.shape = {workspace_size}; - workspace->data.dtype = DType::kByte; - return; - } - } else if (workspace_size == 0) { - workspace->data.shape = {1}; - workspace->data.dtype = DType::kByte; - return; - } -} -} // namespace transformer_engine diff --git a/transformer_engine/common/fused_attn/fused_attn_fp8.h b/transformer_engine/common/fused_attn/fused_attn_fp8.h deleted file mode 100644 index 5b3b2fff147..00000000000 --- a/transformer_engine/common/fused_attn/fused_attn_fp8.h +++ /dev/null @@ -1,46 +0,0 @@ -/************************************************************************* - * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * - * See LICENSE for license information. - ************************************************************************/ - -/*! \file fused_attn_fp8.h - * \brief Functions for fused attention for FP8 - */ - -#include "transformer_engine/fused_attn.h" -#include "transformer_engine/transformer_engine.h" - -namespace transformer_engine { -// fused attention FWD FP8 with separate Q, K, V -void fused_attn_fp8_fwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, bool is_training, float attn_scale, float p_dropout, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, NVTE_QKV_Format qkv_scale_inv_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - size_t window_size_left, size_t window_size_right, bool bottom_right_diagonal, - const Tensor *input_Q, const Tensor *input_K, const Tensor *input_V, - const Tensor *input_SoftmaxOffset, Tensor *input_output_S, Tensor *output_O, - NVTETensorPack *Aux_CTX_Tensors, const Tensor *cu_seqlens_q, const Tensor *cu_seqlens_kv, - const Tensor *cu_seqlens_q_padded, const Tensor *cu_seqlens_kv_padded, const Tensor *rng_state, - Tensor *workspace, cudaStream_t stream, cudnnHandle_t handle); - -// fused attention BWD FP8 with separate Q, K, V -void fused_attn_fp8_bwd( - size_t batch, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, size_t num_tokens_q, - size_t num_tokens_kv, float attn_scale, float p_dropout, NVTE_QKV_Layout qkv_layout, - NVTE_QKV_Format o_format, NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type mask_type, NVTE_Softmax_Type softmax_type, - size_t window_size_left, size_t window_size_right, bool bottom_right_diagonal, - bool deterministic, const Tensor *input_Q, const Tensor *input_K, const Tensor *input_V, - const Tensor *input_O, const Tensor *input_dO, const Tensor *input_dO_f16, - const Tensor *input_M, const Tensor *input_S, const Tensor *input_SoftmaxOffset, - Tensor *input_output_dP, const Tensor *output_dQ, const Tensor *output_dK, - const Tensor *output_dV, Tensor *output_dSoftmaxOffset, const Tensor *cu_seqlens_q, - const Tensor *cu_seqlens_kv, const Tensor *cu_seqlens_q_padded, - const Tensor *cu_seqlens_kv_padded, const Tensor *rng_state, Tensor *workspace, - cudaStream_t stream, cudnnHandle_t handle); -} // namespace transformer_engine diff --git a/transformer_engine/common/fused_attn/utils.cu b/transformer_engine/common/fused_attn/utils.cu deleted file mode 100644 index 9b54a64cbe7..00000000000 --- a/transformer_engine/common/fused_attn/utils.cu +++ /dev/null @@ -1,619 +0,0 @@ -/************************************************************************* - * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * - * See LICENSE for license information. - ************************************************************************/ - -#include -#include - -#include "../common.h" -#include "../cudnn_utils.h" -#include "transformer_engine/fused_attn.h" -#include "utils.h" - -namespace transformer_engine { -namespace fused_attn { - -using namespace transformer_engine; - -// get matrix strides based on matrix type -void generateMatrixStrides(int64_t b, int64_t h, int64_t s_q, int64_t s_kv, int64_t d, - int64_t *strideA, NVTE_QKV_Layout layout, NVTE_QKV_Matrix matrix) { - constexpr int batch_dim_idx = 0; - constexpr int head_dim_idx = 1; - constexpr int seqlen_dim_idx = 2; - constexpr int hidden_dim_idx = 3; - - constexpr int seqlen_transpose_dim_idx = 3; - constexpr int hidden_transpose_dim_idx = 2; - - constexpr int seqlen_q_dim_idx = 2; - constexpr int seqlen_kv_dim_idx = 3; - - switch (layout) { - case NVTE_QKV_Layout::NVTE_SB3HD: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = 3 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * 3 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = 3 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = b * 3 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_SBH3D: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = 3 * h * d; - strideA[head_dim_idx] = 3 * d; - strideA[seqlen_dim_idx] = b * 3 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = 3 * h * d; - strideA[head_dim_idx] = 3 * d; - strideA[seqlen_transpose_dim_idx] = b * 3 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_SB2HD: - if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = 2 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * 2 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = 2 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = b * 2 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_SBH2D: - if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = 2 * h * d; - strideA[head_dim_idx] = 2 * d; - strideA[seqlen_dim_idx] = b * 2 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = 2 * h * d; - strideA[head_dim_idx] = 2 * d; - strideA[seqlen_transpose_dim_idx] = b * 2 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_SBHD_SBHD_SBHD: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = b * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BS3HD: - case NVTE_QKV_Layout::NVTE_T3HD: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = s_q * 3 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = 3 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = s_q * 3 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = 3 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix) { - strideA[batch_dim_idx] = s_q * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BSH3D: - case NVTE_QKV_Layout::NVTE_TH3D: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = s_q * 3 * h * d; - strideA[head_dim_idx] = 3 * d; - strideA[seqlen_dim_idx] = 3 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = s_q * 3 * h * d; - strideA[head_dim_idx] = 3 * d; - strideA[seqlen_transpose_dim_idx] = 3 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix) { - strideA[batch_dim_idx] = s_q * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BSHD_BS2HD: - case NVTE_QKV_Layout::NVTE_THD_T2HD: - if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = s_kv * 2 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = 2 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = s_kv * 2 * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = 2 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = s_q * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BSHD_BSH2D: - case NVTE_QKV_Layout::NVTE_THD_TH2D: - if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = s_kv * 2 * h * d; - strideA[head_dim_idx] = 2 * d; - strideA[seqlen_dim_idx] = 2 * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = s_kv * 2 * h * d; - strideA[head_dim_idx] = 2 * d; - strideA[seqlen_transpose_dim_idx] = 2 * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = s_q * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BSHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_THD_THD_THD: - case NVTE_QKV_Layout::NVTE_THD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_BSHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_THD_BSHD_BSHD: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = s_q * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = s_kv * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = s_kv * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = h * d; - strideA[hidden_transpose_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_SBHD_BSHD_BSHD: - if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = s_kv * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = s_kv * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BSHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_THD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_BSHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_THD_SBHD_SBHD: - if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = b * h * d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_transpose_dim_idx] = b * h * d; - strideA[hidden_transpose_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = s_q * h * d; - strideA[head_dim_idx] = d; - strideA[seqlen_dim_idx] = h * d; - strideA[hidden_dim_idx] = 1; - } - break; - case NVTE_QKV_Layout::NVTE_BHSD_BHSD_BHSD: - if ((matrix == NVTE_QKV_Matrix::NVTE_Q_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_O_Matrix)) { - strideA[batch_dim_idx] = h * s_q * d; - strideA[head_dim_idx] = s_q * d; - strideA[seqlen_dim_idx] = d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix)) { - strideA[batch_dim_idx] = h * s_kv * d; - strideA[head_dim_idx] = s_kv * d; - strideA[seqlen_dim_idx] = d; - strideA[hidden_dim_idx] = 1; - } else if ((matrix == NVTE_QKV_Matrix::NVTE_K_Matrix_Transpose) || - (matrix == NVTE_QKV_Matrix::NVTE_V_Matrix_Transpose)) { - strideA[batch_dim_idx] = h * s_kv * d; - strideA[head_dim_idx] = s_kv * d; - strideA[seqlen_transpose_dim_idx] = d; - strideA[hidden_transpose_dim_idx] = 1; - } - break; - } - - if (matrix == NVTE_QKV_Matrix::NVTE_S_Matrix) { - strideA[seqlen_kv_dim_idx] = 1; - strideA[seqlen_q_dim_idx] = s_kv; - strideA[head_dim_idx] = s_q * s_kv; - strideA[batch_dim_idx] = h * s_q * s_kv; - } -} - -bool allowAllConfig(cudnnBackendDescriptor_t engine_config) { - (void)engine_config; - return false; -} - -cudnn_frontend::Tensor tensor_create(cudnnDataType_t type, int64_t id, int64_t const *dim, - int64_t const *stride, bool is_virtual, bool is_value) { - int nbDims = 4; - auto tensor_created = - cudnn_frontend::TensorBuilder() - .setDim(nbDims, dim) - .setStride(nbDims, stride) - .setId(id) - .setAlignment(16) // 16B alignment is needed to run a tensor core engine - .setDataType(type) - .setVirtual(is_virtual) - .setByValue(is_value) - .build(); - return tensor_created; -} - -cudnn_frontend::Tensor tensor_create_with_offset( - cudnnDataType_t type, int64_t id, int64_t const *dim, int64_t const *stride, bool is_virtual, - bool is_value, std::shared_ptr raggedOffset) { - int nbDims = 4; - auto tensor_created = - cudnn_frontend::TensorBuilder() - .setDim(nbDims, dim) - .setStride(nbDims, stride) - .setId(id) - .setAlignment(16) // 16B alignment is needed to run a tensor core engine - .setDataType(type) - .setVirtual(is_virtual) - .setByValue(is_value) - .setRaggedOffset(raggedOffset) - .build(); - return tensor_created; -} - -cudnn_frontend::PointWiseDesc pw_desc_create(cudnnDataType_t type, cudnnPointwiseMode_t mode) { - auto pw_desc_created = - cudnn_frontend::PointWiseDescBuilder().setMode(mode).setComputeType(type).build(); - return pw_desc_created; -} - -cudnn_frontend::Operation unary_pw_op_create(cudnn_frontend::Tensor const &xDesc, - cudnn_frontend::Tensor const &yDesc, - cudnn_frontend::PointWiseDesc const &pwDesc) { - auto pw_op_created = - cudnn_frontend::OperationBuilder(CUDNN_BACKEND_OPERATION_POINTWISE_DESCRIPTOR) - .setxDesc(xDesc) - .setyDesc(yDesc) - .setpwDesc(pwDesc) - .build(); - return pw_op_created; -} - -cudnn_frontend::Operation binary_pw_op_create(cudnn_frontend::Tensor const &xDesc, - cudnn_frontend::Tensor const &bDesc, - cudnn_frontend::Tensor const &yDesc, - cudnn_frontend::PointWiseDesc const &pwDesc) { - auto pw_op_created = - cudnn_frontend::OperationBuilder(CUDNN_BACKEND_OPERATION_POINTWISE_DESCRIPTOR) - .setxDesc(xDesc) - .setbDesc(bDesc) - .setyDesc(yDesc) - .setpwDesc(pwDesc) - .build(); - return pw_op_created; -} - -cudnn_frontend::Operation ternary_pw_op_create(cudnn_frontend::Tensor const &xDesc, - cudnn_frontend::Tensor const &bDesc, - cudnn_frontend::Tensor const &tDesc, - cudnn_frontend::Tensor const &yDesc, - cudnn_frontend::PointWiseDesc const &pwDesc) { - auto pw_op_created = - cudnn_frontend::OperationBuilder(CUDNN_BACKEND_OPERATION_POINTWISE_DESCRIPTOR) - .setxDesc(xDesc) - .setbDesc(bDesc) - .settDesc(tDesc) - .setyDesc(yDesc) - .setpwDesc(pwDesc) - .build(); - return pw_op_created; -} - -// convert cu_seqlens to actual_seqlens -__global__ void cu_seqlens_to_actual_seqlens(int64_t actual_b, int64_t max_b, - int32_t const *const q_cu_seqlens, - int32_t const *const kv_cu_seqlens, int32_t *q_seqlens, - int32_t *kv_seqlens) { - size_t tid = blockIdx.x * blockDim.x + threadIdx.x; - if (tid < actual_b) { - q_seqlens[tid] = q_cu_seqlens[tid + 1] - q_cu_seqlens[tid]; - kv_seqlens[tid] = kv_cu_seqlens[tid + 1] - kv_cu_seqlens[tid]; - } else if (tid < max_b) { - q_seqlens[tid] = 0; - kv_seqlens[tid] = 0; - } -} - -// convert cu_seqlens_padded to offsets -template -__device__ void cu_seqlens_padded_to_offsets_impl( - const RaggedOffsetMultipliers &mults, int64_t actual_b, int64_t max_b, - const int32_t *cu_seqlens_q_padded, const int32_t *cu_seqlens_kv_padded, OFFSETS_T *offsets_q, - OFFSETS_T *offsets_k, OFFSETS_T *offsets_v, OFFSETS_T *offsets_o, OFFSETS_T *offsets_s) { - size_t tid = blockIdx.x * blockDim.x + threadIdx.x; - auto cu_seqlens_id = min(tid, actual_b); - if (tid <= max_b) { - if (offsets_s != nullptr) { - offsets_s[tid] = mults.stats * cu_seqlens_q_padded[cu_seqlens_id]; - } - if (offsets_q != nullptr && offsets_o != nullptr) { - offsets_q[tid] = mults.q * cu_seqlens_q_padded[cu_seqlens_id]; - offsets_o[tid] = mults.o * cu_seqlens_q_padded[cu_seqlens_id]; - } - if (offsets_k != nullptr && offsets_v != nullptr) { - const int32_t *cu_seqlens_kv_src = - mults.kv_from_q ? cu_seqlens_q_padded : cu_seqlens_kv_padded; - offsets_k[tid] = mults.k * cu_seqlens_kv_src[cu_seqlens_id]; - offsets_v[tid] = mults.v * cu_seqlens_kv_src[cu_seqlens_id]; - } - } -} - -__global__ void cu_seqlens_padded_to_offsets(RaggedOffsetMultipliers mults, int64_t actual_b, - int64_t max_b, const int32_t *cu_seqlens_q_padded, - const int32_t *cu_seqlens_kv_padded, - DType offset_dtype, void *offsets_q, void *offsets_k, - void *offsets_v, void *offsets_o, void *offsets_s) { - if (offset_dtype == DType::kInt32) { - cu_seqlens_padded_to_offsets_impl( - mults, actual_b, max_b, cu_seqlens_q_padded, cu_seqlens_kv_padded, - reinterpret_cast(offsets_q), reinterpret_cast(offsets_k), - reinterpret_cast(offsets_v), reinterpret_cast(offsets_o), - reinterpret_cast(offsets_s)); - } else { - assert(offset_dtype == DType::kInt64 && "expect int64"); - cu_seqlens_padded_to_offsets_impl( - mults, actual_b, max_b, cu_seqlens_q_padded, cu_seqlens_kv_padded, - reinterpret_cast(offsets_q), reinterpret_cast(offsets_k), - reinterpret_cast(offsets_v), reinterpret_cast(offsets_o), - reinterpret_cast(offsets_s)); - } -} - -DType get_ragged_offset_dtype(NVTE_QKV_Layout_Group layout_group, int64_t num_attn_heads, - int64_t num_gqa_groups, int64_t max_seqlen_q, int64_t max_seqlen_kv, - int64_t head_dim_qk, int64_t head_dim_v) { - std::array offsets_qkvo{}; - switch (layout_group) { - case NVTE_QKV_Layout_Group::NVTE_HD_HD_HD: - case NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD: - offsets_qkvo[0] = num_attn_heads * head_dim_qk * max_seqlen_q; - offsets_qkvo[1] = num_gqa_groups * head_dim_qk * max_seqlen_kv; - offsets_qkvo[2] = num_gqa_groups * head_dim_v * max_seqlen_kv; - break; - case NVTE_QKV_Layout_Group::NVTE_3HD: - case NVTE_QKV_Layout_Group::NVTE_H3D: - offsets_qkvo[0] = 3 * num_attn_heads * head_dim_qk * max_seqlen_q; - offsets_qkvo[1] = offsets_qkvo[0]; - offsets_qkvo[2] = offsets_qkvo[0]; - break; - case NVTE_QKV_Layout_Group::NVTE_HD_2HD: - case NVTE_QKV_Layout_Group::NVTE_HD_H2D: - offsets_qkvo[0] = num_attn_heads * head_dim_qk * max_seqlen_q; - offsets_qkvo[1] = 2 * num_gqa_groups * head_dim_qk * max_seqlen_kv; - offsets_qkvo[2] = offsets_qkvo[1]; - break; - } - - offsets_qkvo[3] = num_attn_heads * head_dim_qk * max_seqlen_q; - - size_t max_offset = *std::max_element(offsets_qkvo.begin(), offsets_qkvo.end()); - if (max_offset > std::numeric_limits::max()) { - return DType::kInt64; - } - - return DType::kInt32; -} - -// quantize batch size -size_t get_max_batch_size(size_t batch_size) { - size_t max_b = batch_size; - size_t log2_b = ceil(log2(batch_size)); - // batch size is expected to be 10s-100s - // b = 1, ..., 32 -> max_b = 32 - // b = 33, ..., 512 -> max_b = next power of 2 - // b = 513, ... -> max_b = increment by 512 - if (log2_b <= 5) { - max_b = 32; - } else if (log2_b <= 9) { - max_b = pow(2, log2_b); - } else { - max_b = (batch_size + 511) / 512 * 512; - } - return max_b; -} - -// quantize token count -size_t get_max_tokens(size_t num_tokens) { - // token count is expected to be 1k's-100k's - // t = 0, ..., 1024 -> max_t = 1024 - // t = 1025, ..., 32k -> max_t = next power of 2 - // t = 32k+1, ... -> max_t = increment by 32k - size_t log2_t = ceil(log2(num_tokens)); - size_t max_t = 0; - if (log2_t <= 10) { - max_t = 1024; - } else if (log2_t <= 15) { - max_t = pow(2, log2_t); - } else { - max_t = (num_tokens + 32767) / 32768 * 32768; - } - return max_t; -} - -__global__ void populate_rng_state_kernel(int64_t *rng_state_dst, const int64_t *const seed, - int64_t offset) { - int tid = blockIdx.x * blockDim.x + threadIdx.x; - if (tid > 0) return; - rng_state_dst[0] = seed[0]; - rng_state_dst[1] = offset; -} - -__global__ void get_runtime_num_segments_kernel(int32_t *cu_seqlen, size_t len, uint32_t *out) { - int tid = blockDim.x * blockIdx.x + threadIdx.x; - if (tid >= len) return; - - if (cu_seqlen[tid] > 0) { - // atomicAdd only support 32 bits dtype - atomicAdd(out, 1); - } -} - -void PopulateRngStateAsync(void *rng_state_dst, const void *seed, size_t q_max_seqlen, - size_t kv_max_seqlen, NVTE_Fused_Attn_Backend backend, - cudaStream_t stream) { - size_t increment = 0; - if (backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) { - increment = 16; - } else { - constexpr int threads_per_cta = 128; - increment = (q_max_seqlen * kv_max_seqlen + threads_per_cta - 1) / threads_per_cta; - } - auto offset = FusedAttnOffsetManager::Instance().GetAndUpdateOffset(increment); - populate_rng_state_kernel<<<1, 1, 0, stream>>>(reinterpret_cast(rng_state_dst), - reinterpret_cast(seed), offset); - NVTE_CHECK_CUDA(cudaGetLastError()); -} - -uint32_t GetRuntimeNumSegments(void *cu_seqlen, void *workspace, size_t len, cudaStream_t stream) { - // workspace size requires 4 bytes - uint32_t *dout = static_cast(workspace); - uint32_t hout{}; - NVTE_CHECK_CUDA(cudaMemsetAsync(dout, 0, sizeof(uint32_t), stream)); - constexpr int threads = 128; - const int blocks = (len - 1) / threads + 1; - get_runtime_num_segments_kernel<<>>(static_cast(cu_seqlen), - len, dout); - NVTE_CHECK_CUDA(cudaGetLastError()); - NVTE_CHECK_CUDA(cudaMemcpyAsync(&hout, dout, sizeof(uint32_t), cudaMemcpyDeviceToHost, stream)); - NVTE_CHECK_CUDA(cudaStreamSynchronize(stream)); - return hout; -} - -__global__ void extract_seed_and_offset(int64_t *rng_state_ptr, bool captured, int64_t *seed_ptr, - uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, - uint32_t offset_intragraph) { - if (captured) { - rng_state_ptr[0] = *seed_ptr; - rng_state_ptr[1] = static_cast(*offset_ptr + static_cast(offset_intragraph)); - } else { - rng_state_ptr[0] = static_cast(seed_val); - rng_state_ptr[1] = static_cast(offset_val); - } -} - -} // namespace fused_attn -} // namespace transformer_engine - -void nvte_extract_seed_and_offset(int64_t *rng_state_ptr, int captured, int64_t *seed_ptr, - uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, - uint32_t offset_intragraph, cudaStream_t stream) { - NVTE_API_CALL(nvte_extract_seed_and_offset); - using namespace transformer_engine; - - fused_attn::extract_seed_and_offset<<<1, 1, 0, stream>>>( - rng_state_ptr, captured, seed_ptr, seed_val, offset_ptr, offset_val, offset_intragraph); - NVTE_CHECK_CUDA(cudaGetLastError()); -} diff --git a/transformer_engine/common/fused_attn/utils.h b/transformer_engine/common/fused_attn/utils.h deleted file mode 100644 index 1864f9417d3..00000000000 --- a/transformer_engine/common/fused_attn/utils.h +++ /dev/null @@ -1,397 +0,0 @@ -/************************************************************************* - * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. - * - * See LICENSE for license information. - ************************************************************************/ - -#ifndef TRANSFORMER_ENGINE_FUSED_ATTN_UTILS_H_ -#define TRANSFORMER_ENGINE_FUSED_ATTN_UTILS_H_ - -#include -#include -#include - -#include -#include - -#include "../common.h" -#include "transformer_engine/fused_attn.h" -#include "transformer_engine/transformer_engine.h" - -namespace transformer_engine { -namespace fused_attn { - -using namespace transformer_engine; - -enum NVTE_QKV_Matrix { - NVTE_Q_Matrix = 0, // queries - NVTE_K_Matrix = 1, // keys - NVTE_K_Matrix_Transpose = 2, // keys transposed - NVTE_V_Matrix = 3, // values - NVTE_V_Matrix_Transpose = 4, // values transposed - NVTE_S_Matrix = 5, // output of GEMM1 - NVTE_O_Matrix = 6, // final output -}; - -// Padded sizes for MXFP8 layout (s_q/s_kv/d_qk/d_v and their scaled dimensions) -struct MXFP8PaddedSizes { - int64_t s_q_padded; - int64_t s_kv_padded; - int64_t s_q_scale; - int64_t s_kv_scale; - int64_t s_q_scale_padded; - int64_t s_kv_scale_padded; - int64_t d_qk_padded; - int64_t d_v_padded; - int64_t d_qk_scale; - int64_t d_v_scale; - int64_t d_qk_scale_padded; - int64_t d_v_scale_padded; -}; - -// Pad s and d for MXFP8 quantization -inline MXFP8PaddedSizes pad_s_d_for_mxfp8(int64_t s_q, int64_t s_kv, int64_t d_qk, int64_t d_v) { - constexpr int64_t block_size = 32; - MXFP8PaddedSizes p; - p.s_q_padded = DIVUP_TO_MULTIPLE(s_q, 128); - p.s_kv_padded = DIVUP_TO_MULTIPLE(s_kv, 128); - p.s_q_scale = DIVUP(s_q, block_size); - p.s_kv_scale = DIVUP(s_kv, block_size); - p.s_q_scale_padded = DIVUP_TO_MULTIPLE(p.s_q_scale, 4); - p.s_kv_scale_padded = DIVUP_TO_MULTIPLE(p.s_kv_scale, 4); - p.d_qk_padded = DIVUP_TO_MULTIPLE(d_qk, 128); - p.d_v_padded = DIVUP_TO_MULTIPLE(d_v, 128); - p.d_qk_scale = DIVUP(d_qk, block_size); - p.d_v_scale = DIVUP(d_v, block_size); - p.d_qk_scale_padded = DIVUP_TO_MULTIPLE(p.d_qk_scale, 4); - p.d_v_scale_padded = DIVUP_TO_MULTIPLE(p.d_v_scale, 4); - return p; -} - -// Get matrix strides for a 4D tensor [batch_size, num_heads, sequence_len, head_dim] given a QKV format. -// strides must point to at least 4 int64_t elements. -inline void generateMatrixStridesWithFormat(int64_t b, int64_t h, int64_t s, int64_t d, - int64_t *strides, NVTE_QKV_Format format) { - constexpr int b_dim = 0; - constexpr int h_dim = 1; - constexpr int s_dim = 2; - constexpr int d_dim = 3; - - switch (format) { - case NVTE_QKV_Format::NVTE_BSHD: - case NVTE_QKV_Format::NVTE_THD: - strides[b_dim] = s * h * d; - strides[h_dim] = d; - strides[s_dim] = h * d; - strides[d_dim] = 1; - break; - case NVTE_QKV_Format::NVTE_SBHD: - strides[b_dim] = h * d; - strides[h_dim] = d; - strides[s_dim] = b * h * d; - strides[d_dim] = 1; - break; - case NVTE_QKV_Format::NVTE_BHSD: - strides[b_dim] = h * s * d; - strides[h_dim] = s * d; - strides[s_dim] = d; - strides[d_dim] = 1; - break; - default: - NVTE_CHECK(false, "Invalid format."); - break; - } -} - -// get matrix strides based on layout and matrix type -inline void generateMatrixStridesWithLayout(int64_t b, int64_t h, int64_t hg, int64_t s_q, - int64_t s_kv, int64_t d_qk, int64_t d_v, - int64_t *q_strides, int64_t *k_strides, - int64_t *v_strides, NVTE_QKV_Layout layout) { - constexpr int b_dim = 0; - constexpr int h_dim = 1; - constexpr int s_dim = 2; - constexpr int d_dim = 3; - const NVTE_QKV_Format q_format = nvte_get_q_format(layout); - const NVTE_QKV_Format kv_format = nvte_get_kv_format(layout); - - switch (layout) { - case NVTE_QKV_Layout::NVTE_SB3HD: - q_strides[b_dim] = 3 * h * d_qk; - q_strides[h_dim] = d_qk; - q_strides[s_dim] = b * 3 * h * d_qk; - q_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - k_strides[i] = v_strides[i] = q_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_SBH3D: - q_strides[b_dim] = 3 * h * d_qk; - q_strides[h_dim] = 3 * d_qk; - q_strides[s_dim] = b * 3 * h * d_qk; - q_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - k_strides[i] = v_strides[i] = q_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_SB2HD: - generateMatrixStridesWithFormat(b, h, s_q, d_qk, q_strides, q_format); - k_strides[b_dim] = 2 * hg * d_qk; - k_strides[h_dim] = d_qk; - k_strides[s_dim] = b * 2 * hg * d_qk; - k_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - v_strides[i] = k_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_SBH2D: - generateMatrixStridesWithFormat(b, h, s_q, d_qk, q_strides, q_format); - k_strides[b_dim] = 2 * hg * d_qk; - k_strides[h_dim] = 2 * d_qk; - k_strides[s_dim] = b * 2 * hg * d_qk; - k_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - v_strides[i] = k_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_BS3HD: - case NVTE_QKV_Layout::NVTE_T3HD: - q_strides[b_dim] = s_q * 3 * h * d_qk; - q_strides[h_dim] = d_qk; - q_strides[s_dim] = 3 * h * d_qk; - q_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - k_strides[i] = v_strides[i] = q_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_BSH3D: - case NVTE_QKV_Layout::NVTE_TH3D: - q_strides[b_dim] = s_q * 3 * h * d_qk; - q_strides[h_dim] = 3 * d_qk; - q_strides[s_dim] = 3 * h * d_qk; - q_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - k_strides[i] = v_strides[i] = q_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_BSHD_BS2HD: - case NVTE_QKV_Layout::NVTE_THD_T2HD: - generateMatrixStridesWithFormat(b, h, s_q, d_qk, q_strides, q_format); - k_strides[b_dim] = s_kv * 2 * hg * d_qk; - k_strides[h_dim] = d_qk; - k_strides[s_dim] = 2 * hg * d_qk; - k_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - v_strides[i] = k_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_BSHD_BSH2D: - case NVTE_QKV_Layout::NVTE_THD_TH2D: - generateMatrixStridesWithFormat(b, h, s_q, d_qk, q_strides, q_format); - k_strides[b_dim] = s_kv * 2 * hg * d_qk; - k_strides[h_dim] = 2 * d_qk; - k_strides[s_dim] = 2 * hg * d_qk; - k_strides[d_dim] = 1; - for (int i = 0; i < 4; i++) { - v_strides[i] = k_strides[i]; - } - break; - case NVTE_QKV_Layout::NVTE_SBHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_SBHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_BSHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_THD_THD_THD: - case NVTE_QKV_Layout::NVTE_THD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_BSHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_THD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_SBHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_SBHD_BSHD_BSHD: - case NVTE_QKV_Layout::NVTE_BSHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_THD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_BSHD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_Paged_KV_THD_SBHD_SBHD: - case NVTE_QKV_Layout::NVTE_BHSD_BHSD_BHSD: - generateMatrixStridesWithFormat(b, h, s_q, d_qk, q_strides, q_format); - generateMatrixStridesWithFormat(b, hg, s_kv, d_qk, k_strides, kv_format); - generateMatrixStridesWithFormat(b, hg, s_kv, d_v, v_strides, kv_format); - break; - default: - NVTE_CHECK(false, "Invalid layout."); - break; - } -} - -void generateMatrixStrides(int64_t b, int64_t h, int64_t s_q, int64_t s_kv, int64_t d, - int64_t *strideA, NVTE_QKV_Layout layout, NVTE_QKV_Matrix matrix); - -bool allowAllConfig(cudnnBackendDescriptor_t engine_config); - -cudnn_frontend::Tensor tensor_create(cudnnDataType_t type, int64_t id, int64_t const *dim, - int64_t const *stride, bool is_virtual, bool is_value); - -cudnn_frontend::Tensor tensor_create_with_offset( - cudnnDataType_t type, int64_t id, int64_t const *dim, int64_t const *stride, bool is_virtual, - bool is_value, std::shared_ptr raggedOffset); - -cudnn_frontend::PointWiseDesc pw_desc_create(cudnnDataType_t type, cudnnPointwiseMode_t mode); - -cudnn_frontend::Operation unary_pw_op_create(cudnn_frontend::Tensor const &xDesc, - cudnn_frontend::Tensor const &yDesc, - cudnn_frontend::PointWiseDesc const &pwDesc); - -cudnn_frontend::Operation binary_pw_op_create(cudnn_frontend::Tensor const &xDesc, - cudnn_frontend::Tensor const &bDesc, - cudnn_frontend::Tensor const &yDesc, - cudnn_frontend::PointWiseDesc const &pwDesc); - -cudnn_frontend::Operation ternary_pw_op_create(cudnn_frontend::Tensor const &xDesc, - cudnn_frontend::Tensor const &bDesc, - cudnn_frontend::Tensor const &tDesc, - cudnn_frontend::Tensor const &yDesc, - cudnn_frontend::PointWiseDesc const &pwDesc); - -struct FADescriptor_v1 { - std::int64_t b; - std::int64_t h; - std::int64_t hg; - std::int64_t s_q; - std::int64_t s_kv; - std::int64_t d_qk; - std::int64_t d_v; - std::int64_t num_pages_k; - std::int64_t num_pages_v; - std::int64_t page_size_k; - std::int64_t page_size_v; - std::int64_t max_pages_per_seq_k; - std::int64_t max_pages_per_seq_v; - std::int64_t bias_b; - std::int64_t bias_h; - std::int64_t bias_sq; - std::int64_t bias_skv; - float attnScale; - bool isTraining; - float dropoutProbability; - NVTE_QKV_Layout qkv_layout; - NVTE_QKV_Format o_format; - NVTE_QKV_Format do_format; - NVTE_QKV_Layout dqkv_layout; - NVTE_QKV_Format qkv_scale_inv_format; - NVTE_QKV_Format do_scale_inv_format; - NVTE_Bias_Type bias_type; - NVTE_Mask_Type mask_type; - NVTE_Softmax_Type softmax_type; - std::int64_t window_size_left; - std::int64_t window_size_right; - bool bottom_right_diagonal; - bool deterministic; - cudnn_frontend::DataType_t qkv_tensor_type; - cudnn_frontend::DataType_t o_tensor_type; - cudnn_frontend::DataType_t do_tensor_type; - cudnn_frontend::DataType_t dqkv_tensor_type; - bool return_max_logit; - - bool operator<(const FADescriptor_v1 &rhs) const { - return std::tie(b, h, hg, s_q, s_kv, d_qk, d_v, num_pages_k, num_pages_v, page_size_k, - page_size_v, max_pages_per_seq_k, max_pages_per_seq_v, bias_b, bias_h, bias_sq, - bias_skv, attnScale, isTraining, dropoutProbability, qkv_layout, o_format, - do_format, dqkv_layout, qkv_scale_inv_format, do_scale_inv_format, mask_type, - softmax_type, window_size_left, window_size_right, bottom_right_diagonal, - deterministic, bias_type, qkv_tensor_type, o_tensor_type, do_tensor_type, - dqkv_tensor_type, return_max_logit) < - std::tie(rhs.b, rhs.h, rhs.hg, rhs.s_q, rhs.s_kv, rhs.d_qk, rhs.d_v, rhs.num_pages_k, - rhs.num_pages_v, rhs.page_size_k, rhs.page_size_v, rhs.max_pages_per_seq_k, - rhs.max_pages_per_seq_v, rhs.bias_b, rhs.bias_h, rhs.bias_sq, rhs.bias_skv, - rhs.attnScale, rhs.isTraining, rhs.dropoutProbability, rhs.qkv_layout, - rhs.o_format, rhs.do_format, rhs.dqkv_layout, rhs.qkv_scale_inv_format, - rhs.do_scale_inv_format, rhs.mask_type, rhs.softmax_type, rhs.window_size_left, - rhs.window_size_right, rhs.bottom_right_diagonal, rhs.deterministic, - rhs.bias_type, rhs.qkv_tensor_type, rhs.o_tensor_type, rhs.do_tensor_type, - rhs.dqkv_tensor_type, rhs.return_max_logit); - } -}; - -// Per-tensor scale factors relating cu_seqlens_padded (token units) to tensor-element -// ragged offsets, as a function of the QKV layout group. Single source of truth shared -// by the cu_seqlens_padded_to_offsets conversion kernel and the direct-seqlens path -// (which passes them to cuDNN as ragged offset multipliers). -struct RaggedOffsetMultipliers { - RaggedOffsetMultipliers(NVTE_QKV_Layout_Group layout_group, int64_t h, int64_t hg, int64_t d_qk, - int64_t d_v) - : q(h * d_qk), k(hg * d_qk), v(hg * d_v), o(h * d_v), stats(h), kv_from_q(false) { - switch (layout_group) { - case NVTE_QKV_Layout_Group::NVTE_3HD: - case NVTE_QKV_Layout_Group::NVTE_H3D: - q = k = v = 3 * h * d_qk; - kv_from_q = true; - break; - case NVTE_QKV_Layout_Group::NVTE_HD_2HD: - case NVTE_QKV_Layout_Group::NVTE_HD_H2D: - k = v = 2 * hg * d_qk; - break; - default: - break; - } - } - - int64_t q; - int64_t k; - int64_t v; - int64_t o; - int64_t stats; - // K/V offsets scale the Q-side cu_seqlens_padded (interleaved QKV layouts) - bool kv_from_q; -}; - -__global__ void cu_seqlens_to_actual_seqlens(int64_t actual_b, int64_t max_b, - int32_t const *const q_cu_seqlens, - int32_t const *const kv_cu_seqlens, int32_t *q_seqlens, - int32_t *kv_seqlens); - -__global__ void cu_seqlens_padded_to_offsets(RaggedOffsetMultipliers mults, int64_t actual_b, - int64_t max_b, const int32_t *cu_seqlens_q_padded, - const int32_t *cu_seqlens_kv_padded, - DType offset_dtype, void *offsets_q, void *offsets_k, - void *offsets_v, void *offsets_o, void *offsets_s); - -DType get_ragged_offset_dtype(NVTE_QKV_Layout_Group layout_group, int64_t num_attn_heads, - int64_t num_gqa_groups, int64_t max_seqlen_q, int64_t max_seqlen_kv, - int64_t head_dim_qk, int64_t head_dim_v); - -size_t get_max_batch_size(size_t batch_size); -size_t get_max_tokens(size_t num_tokens); - -class FusedAttnOffsetManager { - public: - static FusedAttnOffsetManager &Instance() { - static thread_local FusedAttnOffsetManager instance; - return instance; - } - - size_t GetAndUpdateOffset(size_t increment) { - size_t ret = offset_; - offset_ += increment; - return ret; - } - - FusedAttnOffsetManager(FusedAttnOffsetManager const &) = delete; - void operator=(FusedAttnOffsetManager const &) = delete; - - private: - FusedAttnOffsetManager() {} - size_t offset_ = 0; -}; - -__global__ void populate_rng_state_kernel(int64_t *rng_state_dst, const int64_t *const seed, - int64_t offset); - -__global__ void get_runtime_num_segments_kernel(int32_t *cu_seqlen, size_t len, uint32_t *out); - -void PopulateRngStateAsync(void *rng_state_dst, const void *const seed, size_t q_max_seqlen, - size_t kv_max_seqlen, NVTE_Fused_Attn_Backend backend, - cudaStream_t stream); - -uint32_t GetRuntimeNumSegments(void *cu_seqlen, void *workspace, size_t len, cudaStream_t stream); - -} // namespace fused_attn -} // namespace transformer_engine - -#endif diff --git a/transformer_engine/common/include/transformer_engine/fused_attn.h b/transformer_engine/common/include/transformer_engine/fused_attn.h index 4a7fd55f9d2..5c2474700c9 100644 --- a/transformer_engine/common/include/transformer_engine/fused_attn.h +++ b/transformer_engine/common/include/transformer_engine/fused_attn.h @@ -5,11 +5,11 @@ ************************************************************************/ /*! \file fused_attn.h - * \brief Enums and functions for fused attention. + * \brief Attention enums and framework support functions. */ -#ifndef TRANSFORMER_ENGINE_FUSED_ATTN_FP8_H_ -#define TRANSFORMER_ENGINE_FUSED_ATTN_FP8_H_ +#ifndef TRANSFORMER_ENGINE_FUSED_ATTN_H_ +#define TRANSFORMER_ENGINE_FUSED_ATTN_H_ #include "stdint.h" #include "transformer_engine.h" @@ -194,229 +194,6 @@ NVTE_QKV_Format nvte_get_q_format(NVTE_QKV_Layout qkv_layout); */ NVTE_QKV_Format nvte_get_kv_format(NVTE_QKV_Layout qkv_layout); -/*! \brief Get fused attention backend based on input parameters. - * - * \param[in] is_training Whether the model is in training mode. - * \param[in] q_dtype The data type of Tensor Q. - * \param[in] kv_dtype The data type of Tensors K, V. - * \param[in] qkv_layout The layout of Tensors Q, K, V. - * \param[in] bias_type The attention bias type. - * \param[in] attn_mask_type The attention mask type. - * \param[in] softmax_type The attention softmax type. - * \param[in] dropout The dropout probability. - * \param[in] num_attn_heads The number of heads in Q. - * \param[in] num_gqa_groups The number of heads in K, V. - * \param[in] max_seqlen_q The sequence length of Q. - * \param[in] max_seqlen_kv The sequence length of K, V. - * \param[in] head_dim_qk The head dimension of Q, K. - * \param[in] head_dim_v The head dimension of V. - * \param[in] window_size_left Sliding window size (the left half). - * \param[in] window_size_right Sliding window size (the right half). - * \param[in] return_max_logit Whether to produce Max along with Stats. - * \param[in] cuda_graph Whether cuda graph capture is enabled or not. - * \param[in] deterministic Whether determinism is required or not. - */ -NVTE_Fused_Attn_Backend nvte_get_fused_attn_backend( - bool is_training, NVTEDType q_dtype, NVTEDType kv_dtype, NVTE_QKV_Layout qkv_layout, - NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, - float dropout, size_t num_attn_heads, size_t num_gqa_groups, size_t max_seqlen_q, - size_t max_seqlen_kv, size_t head_dim_qk, size_t head_dim_v, int64_t window_size_left, - int64_t window_size_right, bool return_max_logit, bool cuda_graph, bool deterministic); - -/*! \brief Compute dot product attention with separate Q, K and V. - * - * Computes: - * - P = Q * Transpose(K) + Bias - * - S = ScaleMaskSoftmax(P) - * - D = Dropout(S) - * - O = D * Transpose(V) - * - * Notes: - * - * Tensors `cu_seqlens_q_padded` and `cu_seqlens_kv_padded` - * help identify the correct offsets of different sequences in tensors Q, K, V and O. - * When the QKV format (`nvte_get_qkv_format(qkv_layout)`) is `bshd` or `sbhd`, - * offset tensors are not used in the attention calculation and can be set to empty `NVTETensor`s. - * When the QKV format is `thd`, these tensors should follow the following rules. - * When there is no padding between sequences, the offset tensors should be equal to - * `cu_seqlens_q` and `cu_seqlens_kv` respectively. - * When there is padding between sequences, users are responsible to adjust the offsets as needed. - * For example, a tensor of 4 sequences `[a, PAD, b, b, c, PAD, PAD, d, d]` should have - * `cu_seqlens = [0, 1, 3, 4, 6]` and `cu_seqlens_padded= [0, 2, 4, 7, 9]`. - * - * \param[in] Q The Q tensor. - * \param[in] K The K tensor. - * \param[in] V The V tensor. - * \param[in] Bias The Bias tensor. - * \param[in] SoftmaxOffset The SoftmaxOffset tensor. - * \param[in,out] S The S tensor. - * \param[out] O The output O tensor. - * \param[out] Aux_CTX_Tensors Auxiliary output tensors when training, - * e.g. softmax stats, optional Max, rng_state. - * \param[in] cu_seqlens_q Cumulative sequence lengths for Q, [batch_size + 1]. - * \param[in] cu_seqlens_kv Cumulative sequence lengths for K and V, [batch_size + 1]. - * \param[in] cu_seqlens_q_padded Cumulative sequence offsets for Q, [batch_size + 1]. - * \param[in] cu_seqlens_kv_padded Cumulative sequence offsets for KV, [batch_size + 1]. - * \param[in] page_table_k Page table for K cache, [batch_size, max_pages_per_seq_k]. - * \param[in] page_table_v Page table for V cache, [batch_size, max_pages_per_seq_v]. - * \param[in] rng_state Seed and offset of CUDA random number generator. - * \param[in] max_seqlen_q Max sequence length used for computing for Q. - * it may be >= max(seqlen_q_i) for i=0,...batch_size-1. - * \param[in] max_seqlen_kv Max sequence length used for computing for K and V. - * it may be >= max(seqlen_kv_i) for i=0,...batch_size-1. - * \param[in] is_training Whether this is in training mode or inference. - * \param[in] return_max_logit Whether to produce Max along with Stats. - * \param[in] cuda_graph Whether cuda graph capture is enabled or not. - * \param[in] attn_scale Scaling factor for Q * K.T. - * \param[in] dropout Dropout probability. - * \param[in] qkv_layout QKV tensors' layout. - * \param[in] o_format Output format. - * \param[in] qkv_scale_inv_format Format of scale-inverse tensors for QKV; - * if NVTE_QKV_Format_NOT_SET, inferred from qkv_layout. - * \param[in] bias_type Bias type. - * \param[in] attn_mask_type Attention mask type. - * \param[in] softmax_type Attention softmax type. - * \param[in] window_size_left Sliding window size (the left half). - * \param[in] window_size_right Sliding window size (the right half). - * \param[in] bottom_right_diagonal Whether to align sliding window and ALiBi diagonal to the bottom right corner of the softmax matrix. - * \param[in] workspace Workspace tensor. - * \param[in] stream CUDA stream used for this operation. - */ -void nvte_fused_attn_fwd(const NVTETensor Q, const NVTETensor K, const NVTETensor V, - const NVTETensor Bias, const NVTETensor SoftmaxOffset, NVTETensor S, - NVTETensor O, NVTETensorPack *Aux_CTX_Tensors, - const NVTETensor cu_seqlens_q, const NVTETensor cu_seqlens_kv, - const NVTETensor cu_seqlens_q_padded, - const NVTETensor cu_seqlens_kv_padded, const NVTETensor page_table_k, - const NVTETensor page_table_v, const NVTETensor rng_state, - size_t max_seqlen_q, size_t max_seqlen_kv, bool is_training, - bool return_max_logit, bool cuda_graph, float attn_scale, float dropout, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_Bias_Type bias_type, - NVTE_Mask_Type attn_mask_type, NVTE_Softmax_Type softmax_type, - int64_t window_size_left, int64_t window_size_right, - bool bottom_right_diagonal, NVTETensor workspace, cudaStream_t stream); - -/*! \brief Compute the backward of the dot product attention with separate Q, K and V. - * - * Notes: - * - * Tensors `cu_seqlens_q_padded` and `cu_seqlens_kv_padded` - * help identify the correct offsets of different sequences in tensors Q, K, V and O. - * When the QKV format (`nvte_get_qkv_format(qkv_layout)`) is `bshd` or `sbhd`, - * offset tensors are not used in the attention calculation and can be set to empty `NVTETensor`s. - * When the QKV format is `thd`, these tensors should follow the following rules. - * When there is no padding between sequences, the offset tensors should be equal to - * `cu_seqlens_q` and `cu_seqlens_kv` respectively. - * When there is padding between sequences, users are responsible to adjust the offsets as needed. - * For example, a tensor of 4 sequences `[a, PAD, b, b, c, PAD, PAD, d, d]` should have - * `cu_seqlens = [0, 1, 3, 4, 6]` and `cu_seqlens_padded= [0, 2, 4, 7, 9]`. - * - * \param[in] Q The Q tensor. - * \param[in] K The K tensor. - * \param[in] V The V tensor. - * \param[in] O The O tensor from forward. - * \param[in] dO The gradient of the O tensor. - * \param[in] S The S tensor. - * \param[in,out] dP The gradient of the P tensor. - * \param[in] Aux_CTX_Tensors Auxiliary tensors from context when in training mode, - * e.g. softmax stats, optional Max, rng_state. - * \param[out] dQ The gradient of the Q tensor. - * \param[out] dK The gradient of the K tensor. - * \param[out] dV The gradient of the V tensor. - * \param[out] dBias The gradient of the Bias tensor. - * \param[out] dSoftmaxOffset The gradient of the SoftmaxOffset tensor. - * \param[in] cu_seqlens_q Cumulative sequence lengths for Q, [batch_size + 1]. - * \param[in] cu_seqlens_kv Cumulative sequence lengths for K and V, [batch_size + 1]. - * \param[in] cu_seqlens_q_padded Cumulative sequence offsets for Q, [batch_size + 1]. - * \param[in] cu_seqlens_kv_padded Cumulative sequence offsets for KV, [batch_size + 1]. - * \param[in] max_seqlen_q Max sequence length used for computing for Q. - * it may be >= max(seqlen_q_i) for i=0,...batch_size-1. - * \param[in] max_seqlen_kv Max sequence length used for computing for K and V. - * it may be >= max(seqlen_kv_i) for i=0,...batch_size-1. - * \param[in] attn_scale Scaling factor for Q * K.T. - * \param[in] dropout Dropout probability. - * \param[in] qkv_layout QKV tensors' layout. - * \param[in] o_format Output format. - * \param[in] do_format Output gradient's format. - * \param[in] dqkv_layout QKV gradient tensors' layout. - * \param[in] qkv_scale_inv_format Format of scale-inverse tensors for QKV; - * if NVTE_QKV_Format_NOT_SET, inferred from qkv_layout. - * \param[in] do_scale_inv_format Format of scale-inverse tensors for dO; - * if NVTE_QKV_Format_NOT_SET, inferred from the output layout. - * \param[in] bias_type Bias type. - * \param[in] attn_mask_type Attention mask type. - * \param[in] softmax_type Attention softmax type. - * \param[in] window_size_left Sliding window size (the left half). - * \param[in] window_size_right Sliding window size (the right half). - * \param[in] bottom_right_diagonal Whether to align sliding window and ALiBi diagonal to the bottom right corner of the softmax matrix. - * \param[in] deterministic Whether to execute with deterministic behaviours. - * \param[in] cuda_graph Whether cuda graph capture is enabled or not. - * \param[in] workspace Workspace tensor. - * \param[in] stream CUDA stream used for this operation. - */ -void nvte_fused_attn_bwd(const NVTETensor Q, const NVTETensor K, const NVTETensor V, - const NVTETensor O, const NVTETensor dO, const NVTETensor S, NVTETensor dP, - const NVTETensorPack *Aux_CTX_Tensors, NVTETensor dQ, NVTETensor dK, - NVTETensor dV, NVTETensor dBias, NVTETensor dSoftmaxOffset, - const NVTETensor cu_seqlens_q, const NVTETensor cu_seqlens_kv, - const NVTETensor cu_seqlens_q_padded, - const NVTETensor cu_seqlens_kv_padded, size_t max_seqlen_q, - size_t max_seqlen_kv, float attn_scale, float dropout, - NVTE_QKV_Layout qkv_layout, NVTE_QKV_Format o_format, - NVTE_QKV_Format do_format, NVTE_QKV_Layout dqkv_layout, - NVTE_QKV_Format qkv_scale_inv_format, NVTE_QKV_Format do_scale_inv_format, - NVTE_Bias_Type bias_type, NVTE_Mask_Type attn_mask_type, - NVTE_Softmax_Type softmax_type, int64_t window_size_left, - int64_t window_size_right, bool bottom_right_diagonal, bool deterministic, - bool cuda_graph, NVTETensor workspace, cudaStream_t stream); - -/*! \brief Update the RNG state with the seed and calculated offset. - * - * \warning This API is **experimental** and subject to change. - * - * \param[in] rng_state_dst RNG state to store seed and offset. - * \param[in] seed Seed for RNG state. - * \param[in] q_max_seqlen Max sequence length used for computing for Q. - * it may be >= max(seqlen_q_i) for i=0,...batch_size-1. - * \param[in] kv_max_seqlen Max sequence length used for computing for K and V. - * it may be >= max(seqlen_kv_i) for i=0,...batch_size-1. - * \param[in] backend Fused attention backend. - * \param[in] stream CUDA stream used for this operation. - */ -void nvte_populate_rng_state_async(NVTETensor rng_state_dst, const NVTETensor seed, - size_t q_max_seqlen, size_t kv_max_seqlen, - NVTE_Fused_Attn_Backend backend, cudaStream_t stream); - -/*! \brief Get KV format for a given QKV layout. - * - * \warning This API is **experimental** and subject to change. - * - * \param[in] cu_seqlens Cumulative sequence lengths, [batch_size + 1]. - * \param[in] workspace Workspace tensor. - * \param[in] len batch_size x sequence_length. - * \param[in] stream CUDA stream used for this operation. - */ -uint32_t nvte_get_runtime_num_segments(NVTETensor cu_seqlens, NVTETensor workspace, size_t len, - cudaStream_t stream); - -/*! \brief Set the seed and offset for RNG state. - * - * \warning This API is **experimental** and subject to change. - * - * \param[out] rng_state_ptr A size 2 array storing the RNG's seed and offset respectively. - * \param[in] captured Whether a CUDA graph is being captured. - * \param[in] seed_ptr Seed pointer. - * \param[in] seed_val Seed value. - * \param[in] offset_ptr Offset pointer. - * \param[in] offset_val Offset value. - * \param[in] offset_intragraph Intragraph offset in RNG states. For use with CUDA Graphs. - * \param[in] stream CUDA stream used for this operation. - */ -void nvte_extract_seed_and_offset(int64_t *rng_state_ptr, int captured, int64_t *seed_ptr, - uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, - uint32_t offset_intragraph, cudaStream_t stream); - /*! \brief Copy keys and values into the KV cache. * * \warning This API is **experimental** and subject to change. @@ -670,50 +447,6 @@ void nvte_multi_tensor_pad_last_dim(NVTETensor *inputs, NVTETensor *outputs, siz #ifdef __cplusplus } // extern "C" - -#include -#include -#include - -/*! \brief Parses a QKV tensor shape into canonical (b, h, s, d, t) dimensions - * and converts between QKV formats. - */ -class AttentionShape { - public: - inline AttentionShape(NVTE_QKV_Format fmt, const size_t *shape) : canonical_{} { - auto [ndim, order] = dim_order(fmt); - for (size_t i = 0; i < ndim; ++i) canonical_[order[i]] = shape[i]; - } - - size_t b() const { return canonical_[0]; } - size_t h() const { return canonical_[1]; } - size_t s() const { return canonical_[2]; } - size_t d() const { return canonical_[3]; } - size_t t() const { return canonical_[4]; } - - inline void to_format(NVTE_QKV_Format dst_fmt, size_t *dst_shape) const { - auto [ndim, order] = dim_order(dst_fmt); - for (size_t i = 0; i < ndim; ++i) dst_shape[i] = canonical_[order[i]]; - } - - private: - static inline std::pair> dim_order(NVTE_QKV_Format fmt) { - switch (fmt) { - case NVTE_QKV_Format::NVTE_BSHD: - return {4, {0, 2, 1, 3}}; // b s h d - case NVTE_QKV_Format::NVTE_SBHD: - return {4, {2, 0, 1, 3}}; // s b h d - case NVTE_QKV_Format::NVTE_BHSD: - return {4, {0, 1, 2, 3}}; // b h s d - case NVTE_QKV_Format::NVTE_THD: - return {3, {4, 1, 3, -1}}; // t h d - default: - return {0, {}}; - } - } - size_t canonical_[5] = {}; -}; - -#endif // __cplusplus - #endif + +#endif // TRANSFORMER_ENGINE_FUSED_ATTN_H_ diff --git a/transformer_engine/common/include/transformer_engine/utils.h b/transformer_engine/common/include/transformer_engine/utils.h index fda49dd5499..69ac9676f34 100644 --- a/transformer_engine/common/include/transformer_engine/utils.h +++ b/transformer_engine/common/include/transformer_engine/utils.h @@ -40,6 +40,24 @@ void nvte_copy_host_to_device_via_kernel(const void *host_ptr, void *device_ptr, void nvte_convert_pointers_to_tensor(const uint64_t *host_ptrs, NVTETensor output, int64_t count, cudaStream_t stream); +/*! \brief Extract a CUDA generator's seed and offset into a device RNG-state buffer. + * + * When a CUDA graph is being captured, the seed and offset are read from device pointers and + * the graph-local offset is added. Otherwise the provided host values are stored directly. + * + * \param[out] rng_state_ptr A two-element device array containing seed and offset. + * \param[in] captured Whether CUDA graph capture is active. + * \param[in] seed_ptr Device pointer to the seed used during capture. + * \param[in] seed_val Seed value used outside capture. + * \param[in] offset_ptr Device pointer to the offset used during capture. + * \param[in] offset_val Offset value used outside capture. + * \param[in] offset_intragraph Offset to add within a captured graph. + * \param[in] stream CUDA stream for the operation. + */ +void nvte_extract_seed_and_offset(int64_t *rng_state_ptr, int captured, int64_t *seed_ptr, + uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, + uint32_t offset_intragraph, cudaStream_t stream); + #ifdef __cplusplus } // extern "C" #endif diff --git a/transformer_engine/common/util/utils.cu b/transformer_engine/common/util/utils.cu index 39d82624632..ae3d1da8ed4 100644 --- a/transformer_engine/common/util/utils.cu +++ b/transformer_engine/common/util/utils.cu @@ -14,6 +14,23 @@ #include "../util/logging.h" namespace transformer_engine { +namespace extract_seed_and_offset { +namespace { + +__global__ void kernel(int64_t *rng_state_ptr, bool captured, int64_t *seed_ptr, + uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, + uint32_t offset_intragraph) { + if (captured) { + rng_state_ptr[0] = *seed_ptr; + rng_state_ptr[1] = static_cast(*offset_ptr + static_cast(offset_intragraph)); + } else { + rng_state_ptr[0] = static_cast(seed_val); + rng_state_ptr[1] = static_cast(offset_val); + } +} + +} // namespace +} // namespace extract_seed_and_offset namespace copy_host_to_device_via_kernel { namespace { @@ -80,3 +97,12 @@ void nvte_convert_pointers_to_tensor(const uint64_t *host_ptrs, NVTETensor outpu nvte_copy_host_to_device_via_kernel(host_ptrs, out_tensor->data.dptr, static_cast(count) * sizeof(uint64_t), stream); } + +void nvte_extract_seed_and_offset(int64_t *rng_state_ptr, int captured, int64_t *seed_ptr, + uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, + uint32_t offset_intragraph, cudaStream_t stream) { + NVTE_API_CALL(nvte_extract_seed_and_offset); + transformer_engine::extract_seed_and_offset::kernel<<<1, 1, 0, stream>>>( + rng_state_ptr, captured, seed_ptr, seed_val, offset_ptr, offset_val, offset_intragraph); + NVTE_CHECK_CUDA(cudaGetLastError()); +} diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index e65a4fd2a40..5f6d437de72 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -211,7 +211,7 @@ def _device_arch() -> int: def ragged_graph_batch_size(input_batch: int, max_segments_per_seq: int) -> int: - """Match TE common's cuDNN graph batch-size bucket for ragged attention.""" + """Preserve the legacy cuDNN graph batch-size bucket for ragged attention.""" batch = int(input_batch) * int(max_segments_per_seq) # Bucketing is part of cuDNN's ragged-stats layout, introduced in 9.6. # Older versions use dense stats and require the physical metadata extent. @@ -225,7 +225,7 @@ def ragged_graph_batch_size(input_batch: int, max_segments_per_seq: int) -> int: def _ragged_graph_token_count(tokens: int) -> int: - """Match TE common's cuDNN graph token-count bucket for ragged attention.""" + """Preserve the legacy cuDNN graph token-count bucket for ragged attention.""" tokens = int(tokens) if tokens <= 1024: return 1024 diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py index 2fb78a856a3..268af899104 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py @@ -5,8 +5,8 @@ """cuDNN SDPA capability selection for the PyTorch frontend. This is intentionally kept in Python alongside the Python cuDNN graph builder. -The conditions mirror ``nvte_get_fused_attn_backend`` in TE common, which is -still used by the other framework frontends. +The conditions preserve the compatibility policy of the removed TE-common +attention backend selector. """ from __future__ import annotations diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index ce87f78fa5a..36df9ad416a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -5,8 +5,8 @@ """PyTorch implementation of cuDNN-backed scaled dot-product attention. All cuDNN graph construction and execution in this module goes through the -``nvidia-cudnn-frontend`` Python API. TE common retains its C++ implementation -for other framework frontends, but PyTorch does not call it. +``nvidia-cudnn-frontend`` Python API rather than a TE-common attention +implementation. """ from __future__ import annotations diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index b15e53e6ee3..7015ff12e96 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -343,7 +343,7 @@ def wrap_mxfp8( blk = MXFP8_BLOCK_SCALING_SIZE # Both rowwise and columnwise Q are required: # - Forward QK^T uses rowwise - # - cuDNN backward (fused_attn_fp8_bwd_impl) requires columnwise for dK gradient + # - The cuDNN backward graph requires columnwise for the dK gradient quantizer = MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=True) return MXFP8Tensor( shape=(s, b, nh, d), From 3a4be7daa1fa88ff5cd04108e40ceb92f1332816 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Thu, 10 Sep 2026 22:34:14 +0000 Subject: [PATCH 15/36] Factor shared cuDNN attention helpers Centralize the FP16/BF16 support policy, graph lifecycle, ragged bucketing, mask normalization, and score-mod cache keys used by the JAX and PyTorch frontends. Preserve framework-specific capability differences explicitly and add regression coverage. Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fused_attn_score_mod.py | 21 + tests/test_common_attention_helpers.py | 242 ++++++++ .../common/attention/__init__.py | 5 + transformer_engine/common/attention/cudnn.py | 515 ++++++++++++++++++ .../common/attention/score_mod.py | 130 +++++ transformer_engine/common/cudnn_frontend.py | 58 ++ .../jax/cpp_extensions/cudnn_attention.py | 387 +++---------- .../jax/cpp_extensions/cudnn_graph.py | 36 +- .../jax/cpp_extensions/flex_attention.py | 131 +---- .../dot_product_attention/_cudnn_backend.py | 465 +++------------- .../dot_product_attention/_cudnn_graph.py | 39 +- .../dot_product_attention/cudnn_attention.py | 79 ++- .../dot_product_attention/flex_attention.py | 103 +--- 13 files changed, 1222 insertions(+), 989 deletions(-) create mode 100644 tests/test_common_attention_helpers.py create mode 100644 transformer_engine/common/attention/__init__.py create mode 100644 transformer_engine/common/attention/cudnn.py create mode 100644 transformer_engine/common/attention/score_mod.py create mode 100644 transformer_engine/common/cudnn_frontend.py diff --git a/tests/jax/test_fused_attn_score_mod.py b/tests/jax/test_fused_attn_score_mod.py index 4a00862d2d3..86596826007 100644 --- a/tests/jax/test_fused_attn_score_mod.py +++ b/tests/jax/test_fused_attn_score_mod.py @@ -731,6 +731,27 @@ def forward(self, _graph, score, _tensors): assert tex_attention._graph_cache_key("fwd", config_1, ()) is None +def test_fused_attn_score_mod_module_lambda_cache_keys_do_not_collide(): + """Different module-level lambdas must not reuse the same cuDNN graph.""" + score_mod_1 = lambda _graph, score, _tensors: score + score_mod_2 = lambda _graph, score, _tensors: score + score_mod_1.__module__ = __name__ + score_mod_2.__module__ = __name__ + score_mod_1.__qualname__ = "" + score_mod_2.__qualname__ = "" + + config_1, _, _ = make_fused_attn_score_mod_config( + score_mod_1, None, None, None, _CONFIG_TEST_SCALING_FACTOR, True + ) + config_2, _, _ = make_fused_attn_score_mod_config( + score_mod_2, None, None, None, _CONFIG_TEST_SCALING_FACTOR, True + ) + + assert config_1 != config_2 + assert tex_attention._graph_cache_key("fwd", config_1, ()) is not None + assert tex_attention._graph_cache_key("fwd", config_2, ()) is not None + + @pytest.mark.skipif(not _has_cudnn_frontend_python(), reason="cuDNN Python frontend is required") def test_fused_attn_score_mod_post_scale_bias_optional_bprop(): """Post-scale-bias score_mod matches the JAX reference without explicit bprop.""" diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py new file mode 100644 index 00000000000..0f9c3d3e41c --- /dev/null +++ b/tests/test_common_attention_helpers.py @@ -0,0 +1,242 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""GPU-independent tests for shared cuDNN attention helpers.""" + +from dataclasses import replace +from types import SimpleNamespace + +import pytest + +from transformer_engine.common.attention.cudnn import ( + AttentionLayout, + FusedAttentionConfig, + check_f16_fused_attention_support, + encode_cudnn_version, + normalize_attention_mask, + ragged_batch_bucket, + ragged_token_bucket, +) +from transformer_engine.common.attention.score_mod import ( + UNCACHEABLE_SCORE_MOD, + freeze_score_mod_cache_key, + score_mod_callback_cache_key, +) +from transformer_engine.common.cudnn_frontend import build_cudnn_graph, make_cudnn_graph + + +def _attention_config(**kwargs): + config = FusedAttentionConfig( + is_training=False, + q_dtype="float16", + kv_dtype="float16", + layout=AttentionLayout("bshd", "bshd", "bshd", "separate"), + bias_type="no_bias", + mask_type="no_mask", + softmax_type="vanilla", + dropout=0.0, + num_attn_heads=16, + num_gqa_groups=16, + max_seqlen_q=128, + max_seqlen_kv=256, + head_dim_qk=128, + head_dim_v=128, + window_size=(-1, -1), + return_max_logit=False, + cuda_graph=False, + deterministic=False, + cudnn_version=(9, 25, 0), + sm_arch=90, + ) + return replace(config, **kwargs) + + +def test_cudnn_version_encoding(): + assert encode_cudnn_version((8, 9, 7)) == 8907 + assert encode_cudnn_version((9, 25, 1)) == 92501 + + +@pytest.mark.parametrize( + "value, expected", + [(1, 32), (32, 32), (33, 64), (513, 1024), (1025, 1536)], +) +def test_ragged_batch_bucket(value, expected): + assert ragged_batch_bucket(value) == expected + + +@pytest.mark.parametrize( + "value, expected", + [(1, 1024), (1024, 1024), (1025, 2048), (32769, 65536), (65537, 98304)], +) +def test_ragged_token_bucket(value, expected): + assert ragged_token_bucket(value) == expected + + +def test_bottom_right_self_attention_normalizes_to_causal(): + mask = normalize_attention_mask( + causal=False, + bottom_right=True, + padding=False, + bottom_right_diagonal=True, + window_size=(-1, 0), + max_seqlen_q=128, + max_seqlen_kv=128, + ) + assert mask.causal + assert not mask.bottom_right + assert not mask.bottom_right_diagonal + + +def test_shared_f16_policy_basic_support_and_rejection(): + assert check_f16_fused_attention_support(_attention_config()).supported + unsupported = check_f16_fused_attention_support(_attention_config(sm_arch=75)) + assert not unsupported.supported + assert "architecture" in unsupported.reason + + +def test_shared_f16_policy_explicit_frontend_capabilities(): + alibi = _attention_config(bias_type="alibi", mask_type="causal") + assert not check_f16_fused_attention_support(alibi).supported + assert check_f16_fused_attention_support(replace(alibi, allow_alibi=True)).supported + + extended_causal = _attention_config( + mask_type="causal", + window_size=(128, 64), + sm_arch=100, + ) + assert not check_f16_fused_attention_support(extended_causal).supported + assert check_f16_fused_attention_support( + replace(extended_causal, allow_extended_causal_window=True) + ).supported + + modern_padding = _attention_config( + mask_type="padding", + dropout=0.1, + cudnn_version=(9, 7, 0), + ) + assert check_f16_fused_attention_support(modern_padding).supported + assert not check_f16_fused_attention_support( + replace(modern_padding, modern_mask_rules_override=True) + ).supported + + +def _not_an_array(_value): + return False + + +def test_score_mod_module_lambda_keys_do_not_collide(): + score_mod_0 = lambda _graph, score, _tensors: score + score_mod_1 = lambda _graph, score, _tensors: score + score_mod_0.__module__ = __name__ + score_mod_1.__module__ = __name__ + score_mod_0.__qualname__ = "" + score_mod_1.__qualname__ = "" + + key_0 = score_mod_callback_cache_key(score_mod_0, is_array=_not_an_array) + key_1 = score_mod_callback_cache_key(score_mod_1, is_array=_not_an_array) + assert key_0 is not UNCACHEABLE_SCORE_MOD + assert key_1 is not UNCACHEABLE_SCORE_MOD + assert key_0 != key_1 + + +def test_score_mod_bound_method_cache_policy(): + class Unkeyed: + def forward(self, _graph, score, _tensors): + return score + + class Keyed: + def score_mod_graph_cache_key(self): + return {"layers": [1, 2]} + + def forward(self, _graph, score, _tensors): + return score + + assert ( + score_mod_callback_cache_key(Unkeyed().forward, is_array=_not_an_array) + is UNCACHEABLE_SCORE_MOD + ) + assert score_mod_callback_cache_key( + Keyed().forward, is_array=_not_an_array + ) == score_mod_callback_cache_key(Keyed().forward, is_array=_not_an_array) + + +def test_score_mod_key_rejects_runtime_arrays(): + marker = object() + with pytest.raises(TypeError, match="must not include tensors"): + freeze_score_mod_cache_key( + {"nested": [marker]}, + is_array=lambda value: value is marker, + ) + + +class _FakeGraph: + def __init__(self, workspace_size=0, unsupported=False): + self.workspace_size = workspace_size + self.unsupported = unsupported + self.calls = [] + + def validate(self): + self.calls.append("validate") + + def build_operation_graph(self): + self.calls.append("build_operation_graph") + + def create_execution_plans(self, modes): + self.calls.append(("create_execution_plans", modes)) + + def check_support(self): + self.calls.append("check_support") + if self.unsupported: + raise _FakeCudnn.cudnnGraphNotSupportedError("unsupported") + + def build_plans(self, policy): + self.calls.append(("build_plans", policy)) + + def get_workspace_size(self): + return self.workspace_size + + +class _FakeCudnn: + class cudnnGraphNotSupportedError(Exception): + pass + + data_type = SimpleNamespace(FLOAT="float") + heur_mode = SimpleNamespace(A="a", FALLBACK="fallback") + build_plan_policy = SimpleNamespace(HEURISTICS_CHOICE="heuristics") + + def __init__(self): + self.graph_kwargs = None + + def pygraph(self, **kwargs): + self.graph_kwargs = kwargs + return "graph" + + +def test_shared_cudnn_graph_creation_and_finalization(): + cudnn = _FakeCudnn() + assert make_cudnn_graph(cudnn, "half", name="attention", handle=7) == "graph" + assert cudnn.graph_kwargs == { + "io_data_type": "half", + "intermediate_data_type": "float", + "compute_data_type": "float", + "name": "attention", + "handle": 7, + } + + graph = _FakeGraph(workspace_size=0) + assert build_cudnn_graph(cudnn, graph, description="attention") == 1 + assert graph.calls == [ + "validate", + "build_operation_graph", + ("create_execution_plans", ["a", "fallback"]), + "check_support", + ("build_plans", "heuristics"), + ] + + +def test_shared_cudnn_graph_support_error_has_context(): + with pytest.raises(RuntimeError, match="cuDNN test graph is not supported"): + build_cudnn_graph( + _FakeCudnn(), _FakeGraph(unsupported=True), description="test" + ) diff --git a/transformer_engine/common/attention/__init__.py b/transformer_engine/common/attention/__init__.py new file mode 100644 index 00000000000..da628165044 --- /dev/null +++ b/transformer_engine/common/attention/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Framework-independent attention helpers.""" diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py new file mode 100644 index 00000000000..7194d50f9de --- /dev/null +++ b/transformer_engine/common/attention/cudnn.py @@ -0,0 +1,515 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Framework-independent cuDNN attention policy and graph-shape helpers.""" + +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class AttentionLayout: + """Normalized layout properties used by the cuDNN support policy.""" + + qkv_format: str + q_format: str + kv_format: str + layout_group: str + is_qkvpacked: bool = False + + @property + def is_thd(self) -> bool: + """Return whether either Q or KV uses packed-token storage.""" + + return self.q_format == "thd" or self.kv_format == "thd" + + +@dataclass(frozen=True) +class FusedAttentionConfig: + """Normalized inputs to the shared FP16/BF16 cuDNN support policy.""" + + is_training: bool + q_dtype: str + kv_dtype: str + layout: AttentionLayout + bias_type: str + mask_type: str + softmax_type: str + dropout: float + num_attn_heads: int + num_gqa_groups: int + max_seqlen_q: int + max_seqlen_kv: int + head_dim_qk: int + head_dim_v: int + window_size: tuple[int, int] + return_max_logit: bool + cuda_graph: bool + deterministic: bool + cudnn_version: tuple[int, int, int] + sm_arch: int + allow_alibi: bool = False + allow_extended_causal_window: bool = False + modern_mask_rules_override: bool = False + + +@dataclass(frozen=True) +class FusedAttentionSupport: + """Result of checking a normalized attention configuration.""" + + supported: bool + reason: str = "" + warning: str | None = None + + +@dataclass(frozen=True) +class AttentionMask: + """Framework-neutral interpretation of mask and sliding-window options.""" + + causal: bool + bottom_right: bool + padding: bool + bottom_right_diagonal: bool + window_left: int + window_right: int + + +def encode_cudnn_version(version: tuple[int, int, int]) -> int: + """Encode a cuDNN backend version using its native integer convention.""" + + major, minor, patch = (int(part) for part in version) + magnitude = 1000 if major < 9 else 10000 + return major * magnitude + minor * 100 + patch + + +def round_up(value: int, multiple: int) -> int: + """Round ``value`` up to a positive multiple.""" + + return (int(value) + int(multiple) - 1) // int(multiple) * int(multiple) + + +def ragged_token_bucket(tokens: int) -> int: + """Return the cuDNN graph bucket for a packed-token extent.""" + + tokens = int(tokens) + if tokens <= 1024: + return 1024 + if tokens <= 32768: + return 1 << (tokens - 1).bit_length() + return round_up(tokens, 32768) + + +def ragged_batch_bucket(batch: int) -> int: + """Return the cuDNN graph bucket for a ragged batch extent.""" + + batch = int(batch) + if batch <= 32: + return 32 + if batch <= 512: + return 1 << (batch - 1).bit_length() + return round_up(batch, 512) + + +def normalize_attention_mask( + *, + causal: bool, + bottom_right: bool, + padding: bool, + bottom_right_diagonal: bool, + window_size: tuple[int, int], + max_seqlen_q: int, + max_seqlen_kv: int, +) -> AttentionMask: + """Normalize equivalent causal and bottom-right mask configurations.""" + + if bottom_right and max_seqlen_q == max_seqlen_kv and not padding: + causal = True + bottom_right = False + bottom_right_diagonal = False + return AttentionMask( + causal=bool(causal), + bottom_right=bool(bottom_right), + padding=bool(padding), + bottom_right_diagonal=bool(bottom_right_diagonal), + window_left=int(window_size[0]), + window_right=int(window_size[1]), + ) + + +def requires_64bit_ragged_offset( + layout: AttentionLayout, + num_attn_heads: int, + num_gqa_groups: int, + max_seqlen_q: int, + max_seqlen_kv: int, + head_dim_qk: int, + head_dim_v: int, +) -> bool: + """Return whether legacy THD element offsets can overflow signed int32.""" + + if layout.qkv_format != "thd": + return False + if layout.layout_group in ("3hd", "h3d", "qkv_packed"): + q = k = v = 3 * num_attn_heads * head_dim_qk * max_seqlen_q + elif layout.layout_group in ("hd_2hd", "hd_h2d", "kv_packed"): + q = num_attn_heads * head_dim_qk * max_seqlen_q + k = v = 2 * num_gqa_groups * head_dim_qk * max_seqlen_kv + else: + q = num_attn_heads * head_dim_qk * max_seqlen_q + k = num_gqa_groups * head_dim_qk * max_seqlen_kv + v = num_gqa_groups * head_dim_v * max_seqlen_kv + output = num_attn_heads * head_dim_qk * max_seqlen_q + return max(q, k, v, output) > 2**31 - 1 + + +def _unsupported(reason: str, warning: str | None = None) -> FusedAttentionSupport: + return FusedAttentionSupport(False, reason, warning) + + +def check_f16_fused_attention_support( + config: FusedAttentionConfig, +) -> FusedAttentionSupport: + """Check the shared FP16/BF16 cuDNN fused-attention compatibility policy.""" + + if config.q_dtype != config.kv_dtype: + return _unsupported("Q and KV must have the same data type") + if config.q_dtype not in ("float16", "bfloat16"): + return _unsupported("only FP16 and BF16 are supported") + + version = encode_cudnn_version(config.cudnn_version) + arch = int(config.sm_arch) + layout = config.layout + is_thd = layout.is_thd + is_training = bool(config.is_training) + sq = int(config.max_seqlen_q) + skv = int(config.max_seqlen_kv) + h = int(config.num_attn_heads) + hg = int(config.num_gqa_groups) + dqk = int(config.head_dim_qk) + dv = int(config.head_dim_v) + dropout = float(config.dropout) + bias = config.bias_type + mask = config.mask_type + softmax = config.softmax_type + left, right = config.window_size + + if version < 8900: + return _unsupported( + "cuDNN is older than 8.9.0", + "FP16/BF16 fused attention requires cuDNN 8.9.0 or newer", + ) + + architecture_ok = ( + (version < 8903 and arch in (80, 90)) + or (version >= 8903 and 80 <= arch < 100) + or (version >= 90700 and arch >= 100) + ) + if not architecture_ok: + return _unsupported("device architecture is not supported") + if version < 90000 and (sq % 64 or skv % 64): + return _unsupported("sequence lengths must be multiples of 64") + if version < 8907 and h != hg: + return _unsupported("GQA requires cuDNN 8.9.7 or newer") + if dqk % 8 or dv % 8: + return _unsupported("head dimensions must be multiples of 8") + + standard_dim = dqk <= 128 and dv <= 128 + hopper_large_dim = ( + dqk <= 256 + and dv <= 256 + and ( + (not is_training and arch == 90 and version >= 90100) + or (is_training and arch == 90 and version >= 90500) + ) + ) + blackwell_fwd_any_dim = ( + not is_training + and arch >= 100 + and version >= 90900 + and sq > 1 + and layout.layout_group != "paged_separate" + ) + generic_fwd_any_dim = ( + not is_training + and version >= 91002 + and ( + layout.layout_group == "paged_separate" + or sq > 1 + or (sq == 1 and mask not in ("causal", "padding_causal")) + ) + ) + blackwell_mla_bwd = ( + dqk == 192 and dv == 128 and is_training and arch >= 100 and version >= 91100 + ) + blackwell_d256_bwd = ( + dqk == 256 + and dv == 256 + and is_training + and 100 <= arch < 110 + and version >= (92500 if is_thd else 92300) + and layout.layout_group != "paged_separate" + and bias == "no_bias" + and dropout == 0.0 + and softmax == "vanilla" + and ( + (left == -1 and right == -1) + or ( + mask + in ( + "causal", + "padding_causal", + "causal_bottom_right", + "padding_causal_bottom_right", + ) + and right in (-1, 0) + ) + ) + ) + if not ( + standard_dim + or hopper_large_dim + or blackwell_fwd_any_dim + or generic_fwd_any_dim + or blackwell_mla_bwd + or blackwell_d256_bwd + ): + return _unsupported("head dimensions are not supported") + if ( + version >= 91100 + and is_training + and arch == 90 + and dqk >= 128 + and dv >= 128 + and (dqk, dv) != (192, 128) + and dqk != dv + ): + return _unsupported( + "this Hopper backward head-dimension combination is unsupported" + ) + + alibi_supported = ( + config.allow_alibi + and bias == "alibi" + and version >= 8906 + and arch >= 90 + and mask + not in ( + "no_mask", + "padding", + "padding_causal", + "padding_causal_bottom_right", + ) + ) + post_scale_bias_supported = bias == "post_scale_bias" and ( + (version >= 8906 and arch >= 90) or (version >= 90000 and arch >= 80) + ) + if bias != "no_bias" and not alibi_supported and not post_scale_bias_supported: + return _unsupported("attention bias is not supported") + + standard_format = layout.qkv_format in ("sbhd", "bshd") + basic_masks = mask in ("no_mask", "causal", "padding", "padding_causal") + mask_ok = version < 8906 and mask == "causal" + if version >= 8906 and standard_format and basic_masks: + mask_ok = True + if ( + version >= 90100 + and layout.qkv_format == "thd" + and mask + in ( + "padding", + "padding_causal", + ) + ): + mask_ok = True + if ( + version >= 90300 + and standard_format + and mask == "causal_bottom_right" + and sq % 64 == 0 + and skv % 64 == 0 + and sq <= skv + and bias == "no_bias" + and dropout == 0.0 + ): + mask_ok = True + if ( + version >= 90500 + and layout.layout_group == "paged_separate" + and ( + mask in ("padding", "padding_causal") + or ( + mask == "padding_causal_bottom_right" + and sq % 64 == 0 + and skv % 64 == 0 + and sq <= skv + ) + ) + and bias == "no_bias" + and dropout == 0.0 + ): + mask_ok = True + if ( + version >= 90600 + and mask == "padding_causal_bottom_right" + and sq % 64 == 0 + and skv % 64 == 0 + and sq <= skv + and bias == "no_bias" + and dropout == 0.0 + ): + mask_ok = True + if version >= 90700: + modern_mask_ok = ( + mask in ("no_mask", "causal") + or ( + mask in ("padding", "padding_causal", "padding_causal_bottom_right") + and bias == "no_bias" + and dropout == 0.0 + ) + or ( + mask in ("causal_bottom_right", "padding_causal_bottom_right") + and sq <= skv + ) + ) + mask_ok = ( + modern_mask_ok + if config.modern_mask_rules_override + else mask_ok or modern_mask_ok + ) + if not mask_ok: + return _unsupported("attention mask is not supported") + if mask in ("padding", "padding_causal") and bias == "post_scale_bias": + return _unsupported("post-scale bias cannot be combined with this padding mask") + + format_ok = ( + layout.qkv_format in ("sbhd", "bshd", "bhsd") + or ( + layout.qkv_format == "thd" + and arch >= 90 + and ((version >= 90100 and h == hg) or version >= 90600) + ) + or ( + layout.q_format in ("sbhd", "bshd", "bhsd", "thd") + and layout.kv_format in ("sbhd", "bshd", "bhsd", "thd") + and (layout.q_format != "thd" or arch >= 90) + and (layout.kv_format != "thd" or arch >= 90) + and version >= 90700 + ) + ) + if not format_ok: + return _unsupported("QKV format is not supported") + + pre_902_window = version < 90200 and left == -1 and right in (-1, 0) + v902_window = version >= 90200 and ( + (left == -1 and right == -1 and mask == "no_mask") + or ( + left >= -1 + and right == 0 + and ( + mask in ("no_mask", "causal") + or (mask == "causal_bottom_right" and sq == skv) + ) + and sq <= skv + and dropout == 0.0 + and bias == "no_bias" + and standard_format + ) + ) + bottom_right_swa_supported = ( + mask not in ("causal_bottom_right", "padding_causal_bottom_right") + or arch < 100 + or sq == skv + or version > 90700 + ) + v906_window = version >= 90600 and ( + (left == -1 and right in (-1, 0)) + or ( + left >= -1 + and right >= -1 + and ( + mask + in ( + "no_mask", + "padding", + "padding_causal", + "causal_bottom_right", + "padding_causal_bottom_right", + ) + or (config.allow_extended_causal_window and mask == "causal") + ) + ) + and sq <= skv + and bias == "no_bias" + and dropout == 0.0 + and bottom_right_swa_supported + ) + window_ok = pre_902_window or v902_window or v906_window + if not window_ok: + return _unsupported("sliding-window configuration is not supported") + + requires_i64 = requires_64bit_ragged_offset( + layout, + h, + hg, + sq, + skv, + dqk, + dv, + ) + if requires_i64 and version < 90500: + return _unsupported("ragged offsets require int64 support") + if version in (91000, 91001): + return _unsupported("cuDNN 9.10.0 and 9.10.1 have known SDPA issues") + if version < 91301 and softmax != "vanilla": + return _unsupported("this softmax type requires cuDNN 9.13.1 or newer") + if config.return_max_logit and version < 92100: + return _unsupported("returning max logits requires cuDNN 9.21 or newer") + if arch >= 100 and is_training: + if config.deterministic: + if version < 91801 or dropout != 0.0 or bias != "no_bias": + return _unsupported("deterministic Blackwell backward is not supported") + elif dropout != 0.0 and bias != "no_bias": + return _unsupported("Blackwell backward does not support dropout with bias") + + if ( + version == 91400 + and skv > 1024 + and left != -1 + and mask not in ("causal", "causal_bottom_right") + ): + return _unsupported( + "cuDNN 9.14.0 does not support this non-causal sliding window", + "This non-causal sliding-window configuration requires cuDNN > 9.14.0", + ) + if ( + version <= 91500 + and is_training + and standard_format + and skv % 128 != 0 + and config.cuda_graph + and mask not in ("padding", "padding_causal", "padding_causal_bottom_right") + ): + return _unsupported( + "this backward CUDA-graph configuration requires cuDNN 9.15.1", + "This backward CUDA-graph configuration requires cuDNN 9.15.1 or newer", + ) + if arch == 120: + if version < 91801: + return _unsupported( + "SM120 requires cuDNN 9.18.1", + "SM120 fused attention requires cuDNN 9.18.1 or newer", + ) + if config.deterministic and is_training: + return _unsupported( + "deterministic backward is not supported on SM120", + "Deterministic fused-attention backward is not supported on SM120", + ) + if is_thd and layout.is_qkvpacked: + return _unsupported( + "QKV-packed THD attention is not supported on SM120", + "T3HD/TH3D fused attention is not supported on SM120", + ) + + return FusedAttentionSupport(True) diff --git a/transformer_engine/common/attention/score_mod.py b/transformer_engine/common/attention/score_mod.py new file mode 100644 index 00000000000..5e9e770dcca --- /dev/null +++ b/transformer_engine/common/attention/score_mod.py @@ -0,0 +1,130 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Framework-independent cache-key policy for cuDNN score-modification graphs.""" + +from __future__ import annotations + +import inspect +from collections.abc import Callable, Mapping +from typing import Any + + +class UncacheableScoreModKey: + """Identity key for callbacks whose graph topology cannot be cached safely.""" + + def __hash__(self): + return id(self) + + def __eq__(self, other): + return self is other + + +UNCACHEABLE_SCORE_MOD = UncacheableScoreModKey() + + +def is_uncacheable_score_mod_key(key: Any) -> bool: + """Return whether a score-modification graph key disables caching.""" + + return isinstance(key, UncacheableScoreModKey) + + +def freeze_score_mod_cache_key(value: Any, *, is_array: Callable[[Any], bool]) -> Any: + """Convert a user-provided score-modification key into a hashable structure.""" + + if is_array(value): + raise TypeError( + "score_mod_graph_cache_key() must not include tensors. Pass runtime tensors " + "through score_mod_tensors or score_mod_bprop_tensors instead." + ) + if isinstance(value, Mapping): + items = ( + ( + freeze_score_mod_cache_key(key, is_array=is_array), + freeze_score_mod_cache_key(item, is_array=is_array), + ) + for key, item in value.items() + ) + return tuple(sorted(items, key=repr)) + if isinstance(value, (list, tuple)): + return tuple( + freeze_score_mod_cache_key(item, is_array=is_array) for item in value + ) + if isinstance(value, (set, frozenset)): + items = (freeze_score_mod_cache_key(item, is_array=is_array) for item in value) + return tuple(sorted(items, key=repr)) + try: + hash(value) + except TypeError as exc: + raise TypeError( + "score_mod_graph_cache_key() must return a hashable value or a nested " + "combination of mapping/list/tuple/set values." + ) from exc + return value + + +def _explicit_cache_key( + callback_owner: Any, *, is_array: Callable[[Any], bool] +) -> Any | None: + explicit_key = getattr(callback_owner, "score_mod_graph_cache_key", None) + if explicit_key is None: + return None + explicit_key = explicit_key() if callable(explicit_key) else explicit_key + return freeze_score_mod_cache_key(explicit_key, is_array=is_array) + + +def score_mod_callback_cache_key( + callback: Callable | None, + *, + is_array: Callable[[Any], bool], + uncacheable_key_factory: Callable[[], Any] | None = None, +) -> Any: + """Create a stable graph key for a score-modification callable. + + Stateful callables must provide ``score_mod_graph_cache_key``. Stateless named + functions use their qualified name, while lambdas use their code object so two + lambdas in the same module cannot collide. + """ + + def uncacheable_key(): + if uncacheable_key_factory is None: + return UNCACHEABLE_SCORE_MOD + return uncacheable_key_factory() + + if callback is None: + return None + self_obj = getattr(callback, "__self__", None) + func_obj = getattr(callback, "__func__", None) + if self_obj is not None and func_obj is not None: + explicit_key = _explicit_cache_key(self_obj, is_array=is_array) + if explicit_key is None: + return uncacheable_key() + return ( + "bound_method", + type(self_obj), + func_obj.__module__, + func_obj.__qualname__, + explicit_key, + ) + + explicit_key = _explicit_cache_key(callback, is_array=is_array) + if explicit_key is not None: + return ( + "callable", + type(callback), + getattr(callback, "__module__", None), + getattr(callback, "__qualname__", None), + explicit_key, + ) + + if ( + inspect.isfunction(callback) + and callback.__closure__ is None + and "" not in callback.__qualname__ + ): + if callback.__name__ == "" or not callback.__qualname__: + return ("function", callback.__module__, callback.__code__) + return ("function", callback.__module__, callback.__qualname__) + + return uncacheable_key() diff --git a/transformer_engine/common/cudnn_frontend.py b/transformer_engine/common/cudnn_frontend.py new file mode 100644 index 00000000000..a1ddb084a38 --- /dev/null +++ b/transformer_engine/common/cudnn_frontend.py @@ -0,0 +1,58 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Framework-independent helpers for constructing cuDNN frontend graphs.""" + +from __future__ import annotations + +import importlib +from typing import Any + + +def import_cudnn_frontend(*, feature: str, requirement: str): + """Import the cuDNN frontend Python package lazily.""" + + try: + return importlib.import_module("cudnn") + except ImportError as exc: + raise ImportError( + f"{feature} requires the cuDNN frontend Python package. Install {requirement}." + ) from exc + + +def make_cudnn_graph( + cudnn, + io_dtype: Any, + *, + name: str | None = None, + handle: Any = None, +): + """Create a cuDNN graph with TE's standard compute and intermediate types.""" + + kwargs = { + "io_data_type": io_dtype, + "intermediate_data_type": cudnn.data_type.FLOAT, + "compute_data_type": cudnn.data_type.FLOAT, + } + if name is not None: + kwargs["name"] = name + if handle is not None: + kwargs["handle"] = handle + return cudnn.pygraph(**kwargs) + + +def build_cudnn_graph(cudnn, graph, *, description: str) -> int: + """Validate and plan a graph, returning a nonzero workspace size.""" + + graph.validate() + graph.build_operation_graph() + try: + graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) + graph.check_support() + except cudnn.cudnnGraphNotSupportedError as exc: + raise RuntimeError( + f"cuDNN {description} graph is not supported: {exc}" + ) from exc + graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) + return max(int(graph.get_workspace_size()), 1) diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index 5f6d437de72..cec1e679e6e 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -14,6 +14,15 @@ import jax.numpy as jnp import numpy as np +from transformer_engine.common.attention.cudnn import ( + AttentionLayout, + FusedAttentionConfig, + check_f16_fused_attention_support, + normalize_attention_mask, + ragged_batch_bucket, + ragged_token_bucket, +) + from .cudnn_graph import ( GraphBinding, SerializedGraph, @@ -21,6 +30,7 @@ dtype_name, finalize_graph, import_cudnn, + make_graph, serialized_graph, ) from .misc import get_all_device_compute_capability, get_cudnn_version @@ -217,21 +227,12 @@ def ragged_graph_batch_size(input_batch: int, max_segments_per_seq: int) -> int: # Older versions use dense stats and require the physical metadata extent. if get_cudnn_version() < (9, 6, 0) or _device_arch() == 120: return batch - if batch <= 32: - return 32 - if batch <= 512: - return 1 << (batch - 1).bit_length() - return ((batch + 511) // 512) * 512 + return ragged_batch_bucket(batch) def _ragged_graph_token_count(tokens: int) -> int: """Preserve the legacy cuDNN graph token-count bucket for ragged attention.""" - tokens = int(tokens) - if tokens <= 1024: - return 1024 - if tokens <= 32768: - return 1 << (tokens - 1).bit_length() - return ((tokens + 32767) // 32768) * 32768 + return ragged_token_bucket(tokens) def _graph_dimensions(info: _LayoutInfo, config): @@ -310,38 +311,39 @@ def _ragged_offset_spec(cudnn): def _mask_options(cudnn, info: _LayoutInfo, config): - is_padding = _is_padding(config) - causal = _is_causal(config) - bottom_right = _is_bottom_right(config) - bottom_right_diagonal = bool(config.bottom_right_diagonal) - if bottom_right and info.q_max_seqlen == info.kv_max_seqlen and not is_padding: - causal = True - bottom_right = False - bottom_right_diagonal = False window_left, window_right = ( config.cp_striped_window_size if config.cp_striped_window_size is not None else config.window_size ) + mask = normalize_attention_mask( + causal=_is_causal(config), + bottom_right=_is_bottom_right(config), + padding=_is_padding(config), + bottom_right_diagonal=bool(config.bottom_right_diagonal), + window_size=(window_left, window_right), + max_seqlen_q=info.q_max_seqlen, + max_seqlen_kv=info.kv_max_seqlen, + ) cudnn_version = get_cudnn_version() options = { "diagonal_alignment": ( cudnn.diagonal_alignment.BOTTOM_RIGHT - if bottom_right_diagonal or bottom_right + if mask.bottom_right_diagonal or mask.bottom_right else cudnn.diagonal_alignment.TOP_LEFT ), } # Before cuDNN 9.6 the preferred right-band API was unavailable, so preserve # the legacy causal flags used by the C++ frontend graph. if cudnn_version < (9, 6, 0): - options["use_causal_mask"] = causal - options["use_causal_mask_bottom_right"] = bottom_right - if cudnn_version >= (9, 2, 0) and window_left != -1: - options["diagonal_band_left_bound"] = int(window_left) + 1 + options["use_causal_mask"] = mask.causal + options["use_causal_mask_bottom_right"] = mask.bottom_right + if cudnn_version >= (9, 2, 0) and mask.window_left != -1: + options["diagonal_band_left_bound"] = mask.window_left + 1 if cudnn_version >= (9, 6, 0): - if window_right != -1: - options["diagonal_band_right_bound"] = int(window_right) - elif causal or bottom_right: + if mask.window_right != -1: + options["diagonal_band_right_bound"] = mask.window_right + elif mask.causal or mask.bottom_right: options["diagonal_band_right_bound"] = 0 return options @@ -383,11 +385,7 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap _graph_dimensions(info, config) ) io_dtype = cudnn_data_type(cudnn, q_aval.dtype) - graph = cudnn.pygraph( - io_data_type=io_dtype, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - ) + graph = make_graph(cudnn, io_dtype) q = _tensor( graph, @@ -667,11 +665,7 @@ def _build_bwd_graph( _graph_dimensions(info, config) ) io_dtype = cudnn_data_type(cudnn, q_aval.dtype) - graph = cudnn.pygraph( - io_data_type=io_dtype, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - ) + graph = make_graph(cudnn, io_dtype) def io_tensor(name, dim, stride, uid, dtype=io_dtype): return _tensor( @@ -925,273 +919,64 @@ def clear_graph_cache(): _graph_cache.clear() -def _encoded_cudnn_version() -> int: - major, minor, patch = get_cudnn_version() - magnitude = 1000 if major < 9 else 10000 - return major * magnitude + minor * 100 + patch - +def _policy_layout(layout) -> AttentionLayout: + qkv_format = layout.get_qkv_format().name.lower() + if layout.is_qkvpacked(): + layout_group = "qkv_packed" + elif layout.is_kvpacked(): + layout_group = "kv_packed" + else: + layout_group = "separate" + return AttentionLayout( + qkv_format=qkv_format, + q_format=qkv_format, + kv_format=qkv_format, + layout_group=layout_group, + is_qkvpacked=layout.is_qkvpacked(), + ) -def is_fused_attn_supported(helper) -> bool: - """JAX-local port of the F16/BF16 cuDNN attention compatibility policy. - - cuDNN frontend ``check_support`` remains authoritative when a concrete graph is - built. This early policy preserves the public fallback behavior for callers that - ask about availability before Q/K/V abstract values exist. - """ - if jnp.dtype(helper.q_dtype) not in ( - jnp.dtype(jnp.float16), - jnp.dtype(jnp.bfloat16), - ): - return False - if jnp.dtype(helper.q_dtype) != jnp.dtype(helper.kv_dtype): - return False - version = _encoded_cudnn_version() - arch = _device_arch() - is_thd = helper.qkv_layout.is_thd() - is_training = bool(helper.is_training) - sq, skv = int(helper.q_max_seqlen), int(helper.kv_max_seqlen) - h, hg = int(helper.q_num_heads), int(helper.kv_num_heads) - dqk, dv = int(helper.head_dim_qk), int(helper.head_dim_v) - dropout = float(helper.dropout_probability) - bias_name = helper.attn_bias_type.name - mask_name = helper.attn_mask_type.name - softmax_name = helper.softmax_type.name - left, right = helper.window_size - deterministic = not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))) - - architecture_ok = ( - (version < 8903 and arch in (80, 90)) - or (version >= 8903 and 80 <= arch < 100) - or (version >= 90700 and arch >= 100) - ) - if version < 8900 or not architecture_ok: - return False - if version < 90000 and (sq % 64 or skv % 64): - return False - if version < 8907 and h != hg: - return False - if dqk % 8 or dv % 8: - return False - - standard_dim = dqk <= 128 and dv <= 128 - hopper_large_dim = ( - dqk <= 256 - and dv <= 256 - and ( - (not is_training and arch == 90 and version >= 90100) - or (is_training and arch == 90 and version >= 90500) - ) - ) - blackwell_fwd_any_dim = ( - not is_training and arch >= 100 and version >= 90900 and sq > 1 - ) - generic_fwd_any_dim = ( - not is_training - and version >= 91002 - and ( - sq > 1 - or (sq == 1 and mask_name not in ("CAUSAL_MASK", "PADDING_CAUSAL_MASK")) - ) - ) - blackwell_mla_bwd = ( - dqk == 192 and dv == 128 and is_training and arch >= 100 and version >= 91100 - ) - blackwell_d256_bwd = ( - dqk == 256 - and dv == 256 - and is_training - and 100 <= arch < 110 - and version >= (92500 if is_thd else 92300) - and bias_name == "NO_BIAS" - and dropout == 0.0 - and softmax_name == "VANILLA_SOFTMAX" - and ( - (left == -1 and right == -1) - or ( - mask_name - in ( - "CAUSAL_MASK", - "PADDING_CAUSAL_MASK", - "CAUSAL_BOTTOM_RIGHT_MASK", - "PADDING_CAUSAL_BOTTOM_RIGHT_MASK", - ) - and right in (-1, 0) - ) - ) - ) - if not ( - standard_dim - or hopper_large_dim - or blackwell_fwd_any_dim - or generic_fwd_any_dim - or blackwell_mla_bwd - or blackwell_d256_bwd - ): - return False - if ( - version >= 91100 - and is_training - and arch == 90 - and dqk >= 128 - and dv >= 128 - and (dqk, dv) != (192, 128) - and dqk != dv - ): - return False +def _policy_mask_name(mask) -> str: + return { + "NO_MASK": "no_mask", + "CAUSAL_MASK": "causal", + "PADDING_MASK": "padding", + "PADDING_CAUSAL_MASK": "padding_causal", + "CAUSAL_BOTTOM_RIGHT_MASK": "causal_bottom_right", + "PADDING_CAUSAL_BOTTOM_RIGHT_MASK": "padding_causal_bottom_right", + }[mask.name] - post_scale_bias_supported = bias_name == "POST_SCALE_BIAS" and ( - (version >= 8906 and arch >= 90) or (version >= 90000 and arch >= 80) - ) - if bias_name != "NO_BIAS" and not post_scale_bias_supported: - return False - - dense_basic_masks = mask_name in ( - "NO_MASK", - "CAUSAL_MASK", - "PADDING_MASK", - "PADDING_CAUSAL_MASK", - ) - thd_basic_masks = mask_name in ("PADDING_MASK", "PADDING_CAUSAL_MASK") - br_mask = mask_name == "CAUSAL_BOTTOM_RIGHT_MASK" - padding_br_mask = mask_name == "PADDING_CAUSAL_BOTTOM_RIGHT_MASK" - mask_ok = False - if version < 8906: - mask_ok = mask_name == "CAUSAL_MASK" and not is_thd - elif not is_thd and dense_basic_masks: - mask_ok = True - if version >= 90100 and is_thd and thd_basic_masks: - mask_ok = True - if ( - version >= 90300 - and not is_thd - and br_mask - and sq % 64 == 0 - and skv % 64 == 0 - and sq <= skv - and bias_name == "NO_BIAS" - and dropout == 0.0 - ): - mask_ok = True - if ( - version >= 90600 - and padding_br_mask - and sq % 64 == 0 - and skv % 64 == 0 - and sq <= skv - and bias_name == "NO_BIAS" - and dropout == 0.0 - ): - mask_ok = True - if version >= 90700: - mask_ok = ( - mask_name in ("NO_MASK", "CAUSAL_MASK") - or ( - mask_name - in ( - "PADDING_MASK", - "PADDING_CAUSAL_MASK", - "PADDING_CAUSAL_BOTTOM_RIGHT_MASK", - ) - and bias_name == "NO_BIAS" - and dropout == 0.0 - ) - or ( - mask_name - in ("CAUSAL_BOTTOM_RIGHT_MASK", "PADDING_CAUSAL_BOTTOM_RIGHT_MASK") - and sq <= skv - ) - ) - if not mask_ok: - return False - if ( - mask_name in ("PADDING_MASK", "PADDING_CAUSAL_MASK") - and bias_name == "POST_SCALE_BIAS" - ): - return False - if is_thd and not ( - arch >= 90 and ((version >= 90100 and h == hg) or version >= 90600) - ): - return False - - full_window = left == -1 and right == -1 - if version < 90200: - window_ok = left == -1 and right in (-1, 0) - elif version < 90600: - window_ok = (full_window and mask_name == "NO_MASK") or ( - left >= -1 - and right == 0 - and mask_name in ("NO_MASK", "CAUSAL_MASK", "CAUSAL_BOTTOM_RIGHT_MASK") - and (mask_name != "CAUSAL_BOTTOM_RIGHT_MASK" or sq == skv) - and sq <= skv - and dropout == 0.0 - and bias_name == "NO_BIAS" - and not is_thd - ) - else: - bottom_right_swa_supported = ( - mask_name - not in ("CAUSAL_BOTTOM_RIGHT_MASK", "PADDING_CAUSAL_BOTTOM_RIGHT_MASK") - or arch < 100 - or sq == skv - or version > 90700 - ) - window_ok = ( - left == -1 - and right in (-1, 0) - or ( - left >= -1 - and right >= -1 - and mask_name - in ( - "NO_MASK", - "CAUSAL_MASK", - "PADDING_MASK", - "PADDING_CAUSAL_MASK", - "CAUSAL_BOTTOM_RIGHT_MASK", - "PADDING_CAUSAL_BOTTOM_RIGHT_MASK", - ) - and sq <= skv - and bias_name == "NO_BIAS" - and dropout == 0.0 - and bottom_right_swa_supported - ) +def is_fused_attn_supported(helper) -> bool: + """Apply the shared F16/BF16 cuDNN attention compatibility policy.""" + + support = check_f16_fused_attention_support( + FusedAttentionConfig( + is_training=bool(helper.is_training), + q_dtype=str(jnp.dtype(helper.q_dtype)), + kv_dtype=str(jnp.dtype(helper.kv_dtype)), + layout=_policy_layout(helper.qkv_layout), + bias_type=helper.attn_bias_type.name.lower(), + mask_type=_policy_mask_name(helper.attn_mask_type), + softmax_type=helper.softmax_type.name.lower().removesuffix("_softmax"), + dropout=float(helper.dropout_probability), + num_attn_heads=int(helper.q_num_heads), + num_gqa_groups=int(helper.kv_num_heads), + max_seqlen_q=int(helper.q_max_seqlen), + max_seqlen_kv=int(helper.kv_max_seqlen), + head_dim_qk=int(helper.head_dim_qk), + head_dim_v=int(helper.head_dim_v), + window_size=tuple(int(value) for value in helper.window_size), + return_max_logit=bool(helper.return_max_logit), + cuda_graph=False, + deterministic=not bool( + int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) + ), + cudnn_version=get_cudnn_version(), + sm_arch=_device_arch(), + allow_alibi=False, + allow_extended_causal_window=True, + modern_mask_rules_override=True, ) - if not window_ok: - return False - - if is_thd: - if helper.qkv_layout.is_qkvpacked(): - max_offset = 3 * h * dqk * sq - elif helper.qkv_layout.is_kvpacked(): - max_offset = max(h * dqk * sq, 2 * hg * dqk * skv) - else: - max_offset = max(h * dqk * sq, hg * dqk * skv, hg * dv * skv) - if max_offset > np.iinfo(np.int32).max and version < 90500: - return False - - if version in (91000, 91001): - return False - if version < 91301 and softmax_name != "VANILLA_SOFTMAX": - return False - if helper.return_max_logit and version < 92100: - return False - if arch >= 100 and is_training: - if deterministic: - if version < 91801 or dropout != 0.0 or bias_name != "NO_BIAS": - return False - elif dropout != 0.0 and bias_name != "NO_BIAS": - return False - if arch == 120 and ( - version < 91801 - or (deterministic and is_training) - or (is_thd and helper.qkv_layout.is_qkvpacked()) - ): - return False - return not ( - version == 91400 - and skv > 1024 - and left != -1 - and mask_name not in ("CAUSAL_MASK", "CAUSAL_BOTTOM_RIGHT_MASK") ) + return support.supported diff --git a/transformer_engine/jax/cpp_extensions/cudnn_graph.py b/transformer_engine/jax/cpp_extensions/cudnn_graph.py index 057150750ec..a879e754382 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_graph.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_graph.py @@ -13,7 +13,6 @@ from __future__ import annotations import hashlib -import importlib from collections.abc import Sequence from dataclasses import dataclass from typing import Any @@ -23,6 +22,12 @@ import numpy as np import transformer_engine_jax +from transformer_engine.common.cudnn_frontend import ( + build_cudnn_graph, + import_cudnn_frontend, + make_cudnn_graph, +) + @dataclass(frozen=True) class GraphBinding: @@ -170,16 +175,20 @@ def check_cudnn_frontend_version_match(cudnn) -> int: def import_cudnn(): """Import and validate the cuDNN frontend Python binding.""" - try: - cudnn = importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "JAX fused_attn requires the cuDNN frontend Python package (`cudnn`)." - ) from exc + cudnn = import_cudnn_frontend( + feature="JAX fused attention", + requirement="nvidia-cudnn-frontend", + ) check_cudnn_frontend_version_match(cudnn) return cudnn +def make_graph(cudnn, io_dtype): + """Create a JAX cuDNN graph with TE's standard compute types.""" + + return make_cudnn_graph(cudnn, io_dtype) + + def graph_hash(serialized_graph: bytes) -> tuple[int, int]: """Return two signed int64 values used as the C++ graph-cache key.""" digest = hashlib.sha256(serialized_graph).digest() @@ -237,18 +246,9 @@ def binding_array(bindings, field): def finalize_graph(cudnn, graph, *, description: str) -> tuple[int, bytes, int]: """Validate, plan and serialize a cuDNN frontend graph.""" - graph.validate() - graph.build_operation_graph() - try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError( - f"cuDNN {description} graph is not supported: {exc}" - ) from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) + workspace_size = build_cudnn_graph(cudnn, graph, description=description) return ( - max(int(graph.get_workspace_size()), 1), + workspace_size, bytes(graph.serialize()), check_cudnn_frontend_version_match(cudnn), ) diff --git a/transformer_engine/jax/cpp_extensions/flex_attention.py b/transformer_engine/jax/cpp_extensions/flex_attention.py index 2ddc9afef14..c018182585e 100644 --- a/transformer_engine/jax/cpp_extensions/flex_attention.py +++ b/transformer_engine/jax/cpp_extensions/flex_attention.py @@ -3,7 +3,6 @@ # See LICENSE for license information. """cuDNN frontend score_mod fused attention helpers.""" -import inspect import os from dataclasses import dataclass from typing import Any, Callable, Dict, Mapping, Optional, Sequence, Tuple @@ -13,11 +12,18 @@ import numpy as np from jax import ffi +from transformer_engine.common.attention.score_mod import ( + UncacheableScoreModKey, + is_uncacheable_score_mod_key, + score_mod_callback_cache_key, +) + from .cudnn_graph import ( GraphBinding, SerializedGraph, finalize_graph, import_cudnn, + make_graph, ) from .cudnn_graph import ( bshd_as_bhsd_dim_stride as _bshd_as_bhsd_dim_stride, @@ -137,103 +143,6 @@ class _ScoreModScalarSpec: stride: Tuple[int, ...] = (1, 1, 1, 1) -class _UncacheableScoreModKey: - """Unique static key for callbacks that must not share compiled score_mod graphs.""" - - def __hash__(self): - return id(self) - - def __eq__(self, other): - return self is other - - -def _score_mod_key_is_uncacheable(key: Any) -> bool: - return isinstance(key, _UncacheableScoreModKey) - - -def _freeze_score_mod_cache_key(value: Any) -> Any: - """Convert a user-provided score_mod graph key into a hashable structure.""" - if _is_array_operand(value): - raise TypeError( - "score_mod_graph_cache_key() must not include tensors. Pass runtime tensors " - "through score_mod_tensors or score_mod_bprop_tensors instead." - ) - if isinstance(value, Mapping): - items = ( - ( - _freeze_score_mod_cache_key(key), - _freeze_score_mod_cache_key(val), - ) - for key, val in value.items() - ) - return tuple(sorted(items, key=repr)) - if isinstance(value, (list, tuple)): - return tuple(_freeze_score_mod_cache_key(item) for item in value) - if isinstance(value, (set, frozenset)): - items = (_freeze_score_mod_cache_key(item) for item in value) - return tuple(sorted(items, key=repr)) - try: - hash(value) - except TypeError as exc: - raise TypeError( - "score_mod_graph_cache_key() must return a hashable value or a nested " - "combination of mapping/list/tuple/set values." - ) from exc - return value - - -def _score_mod_explicit_cache_key(callback_owner: Any) -> Optional[Any]: - """Return a user-provided structural graph key for a score_mod callback.""" - explicit_key = getattr(callback_owner, "score_mod_graph_cache_key", None) - if explicit_key is None: - return None - explicit_key = explicit_key() if callable(explicit_key) else explicit_key - return _freeze_score_mod_cache_key(explicit_key) - - -def _score_mod_callback_cache_key(callback: Optional[Callable]) -> Any: - """Create a stable graph cache key for a score_mod callable. - - Module-level functions are assumed to have stable topology. Stateful bound methods and - callable instances need an explicit score_mod_graph_cache_key(); otherwise their graphs - are left uncached to avoid reusing stale graphs after Python object address reuse. - """ - if callback is None: - return None - self_obj = getattr(callback, "__self__", None) - func_obj = getattr(callback, "__func__", None) - if self_obj is not None and func_obj is not None: - explicit_key = _score_mod_explicit_cache_key(self_obj) - if explicit_key is None: - return _UncacheableScoreModKey() - return ( - "bound_method", - type(self_obj), - func_obj.__module__, - func_obj.__qualname__, - explicit_key, - ) - - explicit_key = _score_mod_explicit_cache_key(callback) - if explicit_key is not None: - return ( - "callable", - type(callback), - getattr(callback, "__module__", None), - getattr(callback, "__qualname__", None), - explicit_key, - ) - - if ( - inspect.isfunction(callback) - and callback.__closure__ is None - and "" not in callback.__qualname__ - ): - return ("function", callback.__module__, callback.__qualname__) - - return _UncacheableScoreModKey() - - @dataclass(frozen=True) class _FusedAttnScoreModConfig: """Static configuration for cuDNN frontend score_mod SDPA graphs.""" @@ -311,6 +220,20 @@ def _is_array_operand(value: Any) -> bool: ) +def _score_mod_callback_cache_key(callback: Optional[Callable]) -> Any: + """Compatibility wrapper around the shared score-modification key policy.""" + + return score_mod_callback_cache_key( + callback, + is_array=_is_array_operand, + uncacheable_key_factory=UncacheableScoreModKey, + ) + + +def _score_mod_key_is_uncacheable(key: Any) -> bool: + return is_uncacheable_score_mod_key(key) + + def _scalar_to_spec(name: str, value: Any) -> _ScoreModScalarSpec: if isinstance(value, bool): dtype = np.bool_ @@ -496,11 +419,7 @@ def _build_score_mod_fwd_graph(q_aval, k_aval, v_aval, score_mod_avals, config): cudnn = _import_cudnn_for_score_mod() io_data_type = _cudnn_data_type(cudnn, q_aval.dtype) - graph = cudnn.pygraph( - io_data_type=io_data_type, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - ) + graph = make_graph(cudnn, io_data_type) q_dim, q_stride = _bshd_as_bhsd_dim_stride(q_aval.shape) k_dim, k_stride = _bshd_as_bhsd_dim_stride(k_aval.shape) @@ -577,11 +496,7 @@ def _build_score_mod_bwd_graph( cudnn = _import_cudnn_for_score_mod() io_data_type = _cudnn_data_type(cudnn, q_aval.dtype) - graph = cudnn.pygraph( - io_data_type=io_data_type, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, - ) + graph = make_graph(cudnn, io_data_type) q_dim, q_stride = _bshd_as_bhsd_dim_stride(q_aval.shape) k_dim, k_stride = _bshd_as_bhsd_dim_stride(k_aval.shape) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py index 268af899104..2c5159a5358 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py @@ -2,17 +2,19 @@ # # See LICENSE for license information. -"""cuDNN SDPA capability selection for the PyTorch frontend. - -This is intentionally kept in Python alongside the Python cuDNN graph builder. -The conditions preserve the compatibility policy of the removed TE-common -attention backend selector. -""" +"""cuDNN SDPA capability selection for the PyTorch frontend.""" from __future__ import annotations import warnings +from transformer_engine.common.attention.cudnn import ( + AttentionLayout, + FusedAttentionConfig, + check_f16_fused_attention_support, + encode_cudnn_version, + requires_64bit_ragged_offset, +) from transformer_engine.pytorch.constants import DType from transformer_engine.pytorch.utils import ( get_cudnn_version, @@ -44,34 +46,23 @@ def tensor_format(component: str) -> str: return qkv_format, q_format, kv_format, layout_group -def _requires_64bit_ragged_offset( - qkv_format: str, - layout_group: str, - num_attn_heads: int, - num_gqa_groups: int, - max_seqlen_q: int, - max_seqlen_kv: int, - head_dim_qk: int, - head_dim_v: int, -) -> bool: - if qkv_format != "thd": - return False - if layout_group in ("3hd", "h3d"): - q = k = v = 3 * num_attn_heads * head_dim_qk * max_seqlen_q - elif layout_group in ("hd_2hd", "hd_h2d"): - q = num_attn_heads * head_dim_qk * max_seqlen_q - k = v = 2 * num_gqa_groups * head_dim_qk * max_seqlen_kv - else: - q = num_attn_heads * head_dim_qk * max_seqlen_q - k = num_gqa_groups * head_dim_qk * max_seqlen_kv - v = num_gqa_groups * head_dim_v * max_seqlen_kv - output = num_attn_heads * head_dim_qk * max_seqlen_q - return max(q, k, v, output) > 2**31 - 1 +def _normalized_layout(qkv_layout: str) -> AttentionLayout: + qkv_format, q_format, kv_format, layout_group = _layout_info(qkv_layout) + return AttentionLayout( + qkv_format=qkv_format, + q_format=q_format, + kv_format=kv_format, + layout_group=layout_group, + is_qkvpacked=layout_group in ("3hd", "h3d"), + ) -def _version_number() -> int: - major, minor, patch = get_cudnn_version() - return major * 10000 + minor * 100 + patch +def _dtype_name(dtype: DType) -> str: + if dtype == DType.kFloat16: + return "float16" + if dtype == DType.kBFloat16: + return "bfloat16" + return dtype.name def get_fused_attn_backend( @@ -97,8 +88,7 @@ def get_fused_attn_backend( ): """Return the Python cuDNN SDPA backend value for an attention configuration.""" - # Import lazily to keep this module independent of graph construction and - # avoid a circular import through dot_product_attention.utils. + # Import lazily to avoid a circular import through dot_product_attention.utils. from .cudnn_attention import FusedAttnBackend q_dtype = DType.cast(q_dtype) @@ -108,12 +98,11 @@ def get_fused_attn_backend( major, minor = get_device_compute_capability() sm_arch = major * 10 + minor - cudnn_version = _version_number() - qkv_format, q_format, kv_format, layout_group = _layout_info(qkv_layout) - is_thd_layout = q_format == "thd" or kv_format == "thd" - requires_i64 = _requires_64bit_ragged_offset( - qkv_format, - layout_group, + cudnn_version_tuple = get_cudnn_version() + cudnn_version = encode_cudnn_version(cudnn_version_tuple) + layout = _normalized_layout(qkv_layout) + requires_i64 = requires_64bit_ragged_offset( + layout, num_attn_heads, num_gqa_groups, max_seqlen_q, @@ -121,8 +110,9 @@ def get_fused_attn_backend( head_dim_qk, head_dim_v, ) - supported_ragged_offset_size = not requires_i64 or cudnn_version >= 90500 + # FP8 and MXFP8 remain PyTorch-specific. The shared policy below covers the + # F16/BF16 subset implemented by both framework frontends. fp8_dtype = q_dtype in (DType.kFloat8E4M3, DType.kFloat8E5M2) fp8_shape_mask = ( ( @@ -167,9 +157,9 @@ def get_fused_attn_backend( ) fp8_format_softmax = ( cudnn_version < 92100 - and qkv_format in ("bshd", "sbhd") + and layout.qkv_format in ("bshd", "sbhd") and softmax_type == "vanilla" - ) or (cudnn_version >= 92100 and qkv_format in ("bshd", "sbhd", "bhsd")) + ) or (cudnn_version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd")) if ( fp8_dtype and sm_arch >= 90 @@ -182,366 +172,33 @@ def get_fused_attn_backend( ): return FusedAttnBackend.FP8 - if q_dtype not in (DType.kFloat16, DType.kBFloat16): - return FusedAttnBackend.No_Backend - - arch_supported = ( - (cudnn_version < 8903 and sm_arch in (80, 90)) - or (cudnn_version >= 8903 and 80 <= sm_arch < 100) - or (cudnn_version >= 90700 and sm_arch >= 100) - ) - seq_supported = cudnn_version >= 90000 or ( - max_seqlen_q % 64 == 0 and max_seqlen_kv % 64 == 0 - ) - heads_supported = cudnn_version >= 8907 or num_attn_heads == num_gqa_groups - - dim_supported = ( - head_dim_qk % 8 == 0 - and head_dim_v % 8 == 0 - and ( - (head_dim_qk <= 128 and head_dim_v <= 128) - or ( - head_dim_qk <= 256 - and head_dim_v <= 256 - and ( - (not is_training and sm_arch == 90 and cudnn_version >= 90100) - or (is_training and sm_arch == 90 and cudnn_version >= 90500) - ) - ) - or ( - not is_training - and sm_arch >= 100 - and cudnn_version >= 90900 - and max_seqlen_q > 1 - and layout_group != "paged_separate" - ) - or ( - not is_training - and cudnn_version >= 91002 - and ( - layout_group == "paged_separate" - or max_seqlen_q > 1 - or ( - max_seqlen_q == 1 - and attn_mask_type not in ("causal", "padding_causal") - ) - ) - ) - or ( - head_dim_qk == 192 - and head_dim_v == 128 - and is_training - and sm_arch >= 100 - and cudnn_version >= 91100 - ) - or ( - head_dim_qk == 256 - and head_dim_v == 256 - and is_training - and 100 <= sm_arch < 110 - and cudnn_version >= (92500 if is_thd_layout else 92300) - and layout_group != "paged_separate" - and bias_type == "no_bias" - and dropout == 0.0 - and softmax_type == "vanilla" - and ( - (window_size_left == -1 and window_size_right == -1) - or ( - attn_mask_type - in ( - "causal", - "padding_causal", - "causal_bottom_right", - "padding_causal_bottom_right", - ) - and window_size_right in (-1, 0) - ) - ) - ) - ) - ) - dim_supported = dim_supported and not ( - cudnn_version >= 91100 - and is_training - and sm_arch == 90 - and head_dim_qk >= 128 - and head_dim_v >= 128 - and (head_dim_qk, head_dim_v) != (192, 128) - and head_dim_qk != head_dim_v - ) - - bias_supported = ( - (cudnn_version < 8906 and bias_type == "no_bias") - or ( - cudnn_version >= 8906 - and ( - bias_type == "no_bias" - or ( - bias_type == "alibi" - and attn_mask_type - not in ( - "no_mask", - "padding", - "padding_causal", - "padding_causal_bottom_right", - ) - and sm_arch >= 90 - ) - or (bias_type == "post_scale_bias" and sm_arch >= 90) - ) - ) - or (cudnn_version >= 90000 and bias_type == "post_scale_bias" and sm_arch >= 80) - ) - - standard_format = qkv_format in ("sbhd", "bshd") - mask_supported = ( - (cudnn_version < 8906 and attn_mask_type == "causal") - or ( - cudnn_version >= 8906 - and standard_format - and attn_mask_type in ("causal", "padding", "padding_causal", "no_mask") - ) - or ( - cudnn_version >= 90100 - and qkv_format == "thd" - and attn_mask_type in ("padding", "padding_causal") - ) - or ( - cudnn_version >= 90300 - and standard_format - and attn_mask_type == "causal_bottom_right" - and max_seqlen_q % 64 == 0 - and max_seqlen_kv % 64 == 0 - and max_seqlen_q <= max_seqlen_kv - and bias_type == "no_bias" - and dropout == 0.0 - ) - or ( - cudnn_version >= 90500 - and layout_group == "paged_separate" - and ( - attn_mask_type in ("padding", "padding_causal") - or ( - attn_mask_type == "padding_causal_bottom_right" - and max_seqlen_q % 64 == 0 - and max_seqlen_kv % 64 == 0 - and max_seqlen_q <= max_seqlen_kv - ) - ) - and bias_type == "no_bias" - and dropout == 0.0 - ) - or ( - cudnn_version >= 90600 - and attn_mask_type == "padding_causal_bottom_right" - and max_seqlen_q % 64 == 0 - and max_seqlen_kv % 64 == 0 - and max_seqlen_q <= max_seqlen_kv - and bias_type == "no_bias" - and dropout == 0.0 - ) - or ( - cudnn_version >= 90700 - and ( - attn_mask_type in ("no_mask", "causal") - or ( - attn_mask_type - in ("padding", "padding_causal", "padding_causal_bottom_right") - and bias_type == "no_bias" - and dropout == 0.0 - ) - or ( - attn_mask_type - in ("causal_bottom_right", "padding_causal_bottom_right") - and max_seqlen_q <= max_seqlen_kv - ) - ) - ) - ) - - bias_mask_supported = not ( - cudnn_version >= 8906 - and attn_mask_type in ("padding", "padding_causal") - and bias_type == "post_scale_bias" - ) - format_supported = ( - qkv_format in ("sbhd", "bshd", "bhsd") - or ( - qkv_format == "thd" - and sm_arch >= 90 - and ( - (cudnn_version >= 90100 and num_attn_heads == num_gqa_groups) - or cudnn_version >= 90600 - ) - ) - or ( - q_format in ("sbhd", "bshd", "bhsd", "thd") - and kv_format in ("sbhd", "bshd", "bhsd", "thd") - and (q_format != "thd" or sm_arch >= 90) - and (kv_format != "thd" or sm_arch >= 90) - and cudnn_version >= 90700 + support = check_f16_fused_attention_support( + FusedAttentionConfig( + is_training=bool(is_training), + q_dtype=_dtype_name(q_dtype), + kv_dtype=_dtype_name(kv_dtype), + layout=layout, + bias_type=bias_type, + mask_type=attn_mask_type, + softmax_type=softmax_type, + dropout=float(dropout), + num_attn_heads=int(num_attn_heads), + num_gqa_groups=int(num_gqa_groups), + max_seqlen_q=int(max_seqlen_q), + max_seqlen_kv=int(max_seqlen_kv), + head_dim_qk=int(head_dim_qk), + head_dim_v=int(head_dim_v), + window_size=(int(window_size_left), int(window_size_right)), + return_max_logit=bool(return_max_logit), + cuda_graph=bool(cuda_graph), + deterministic=bool(deterministic), + cudnn_version=cudnn_version_tuple, + sm_arch=sm_arch, + allow_alibi=True, ) ) - - sliding_window_supported = ( - ( - cudnn_version < 90200 - and window_size_left == -1 - and window_size_right in (-1, 0) - ) - or ( - cudnn_version >= 90200 - and ( - ( - window_size_left == -1 - and window_size_right == -1 - and attn_mask_type == "no_mask" - ) - or ( - window_size_left >= -1 - and window_size_right == 0 - and ( - attn_mask_type in ("no_mask", "causal") - or ( - attn_mask_type == "causal_bottom_right" - and max_seqlen_q == max_seqlen_kv - ) - ) - and max_seqlen_q <= max_seqlen_kv - and dropout == 0.0 - and bias_type == "no_bias" - and standard_format - ) - ) - ) - or ( - cudnn_version >= 90600 - and ( - (window_size_left == -1 and window_size_right in (-1, 0)) - or ( - window_size_left >= -1 - and window_size_right >= -1 - and ( - ( - attn_mask_type == "causal_bottom_right" - and ( - sm_arch < 100 - or ( - sm_arch >= 100 - and ( - ( - max_seqlen_q == max_seqlen_kv - and cudnn_version <= 90700 - ) - or cudnn_version > 90700 - ) - ) - ) - ) - or attn_mask_type in ("no_mask", "padding", "padding_causal") - or ( - attn_mask_type == "padding_causal_bottom_right" - and ( - sm_arch < 100 - or ( - sm_arch >= 100 - and ( - ( - max_seqlen_q == max_seqlen_kv - and cudnn_version <= 90700 - ) - or cudnn_version > 90700 - ) - ) - ) - ) - ) - and max_seqlen_q <= max_seqlen_kv - and bias_type == "no_bias" - and dropout == 0.0 - ) - ) - ) - ) - - softmax_supported = cudnn_version >= 91301 or softmax_type == "vanilla" - max_supported = not return_max_logit or cudnn_version >= 92100 - deterministic_supported = sm_arch < 100 or ( - not is_training - or ( - is_training - and not deterministic - and (dropout == 0.0 or bias_type == "no_bias") - ) - or ( - is_training - and deterministic - and cudnn_version >= 91801 - and dropout == 0.0 - and bias_type == "no_bias" - ) - ) - - supported = all( - ( - arch_supported, - seq_supported, - heads_supported, - dim_supported, - bias_supported, - mask_supported, - bias_mask_supported, - format_supported, - sliding_window_supported, - supported_ragged_offset_size, - cudnn_version not in (91000, 91001), - softmax_supported, - max_supported, - deterministic_supported, - ) - ) - backend = ( - FusedAttnBackend.F16_arbitrary_seqlen - if supported - else FusedAttnBackend.No_Backend - ) - - if cudnn_version < 8900 and backend == FusedAttnBackend.F16_arbitrary_seqlen: - backend = FusedAttnBackend.No_Backend - warnings.warn("FP16/BF16 fused attention requires cuDNN 8.9.0 or newer") - if ( - cudnn_version == 91400 - and max_seqlen_kv > 1024 - and window_size_left != -1 - and attn_mask_type not in ("causal", "causal_bottom_right") - ): - backend = FusedAttnBackend.No_Backend - warnings.warn( - "This non-causal sliding-window configuration requires cuDNN > 9.14.0" - ) - if ( - cudnn_version <= 91500 - and is_training - and standard_format - and max_seqlen_kv % 128 != 0 - and cuda_graph - and attn_mask_type - not in ("padding", "padding_causal", "padding_causal_bottom_right") - ): - backend = FusedAttnBackend.No_Backend - warnings.warn( - "This backward CUDA-graph configuration requires cuDNN 9.15.1 or newer" - ) - if backend == FusedAttnBackend.F16_arbitrary_seqlen and sm_arch == 120: - if cudnn_version < 91801: - backend = FusedAttnBackend.No_Backend - warnings.warn("SM120 fused attention requires cuDNN 9.18.1 or newer") - elif deterministic and is_training: - backend = FusedAttnBackend.No_Backend - warnings.warn( - "Deterministic fused-attention backward is not supported on SM120" - ) - elif qkv_layout in ("t3hd", "th3d"): - backend = FusedAttnBackend.No_Backend - warnings.warn("T3HD/TH3D fused attention is not supported on SM120") - return backend + if support.warning is not None: + warnings.warn(support.warning) + if support.supported: + return FusedAttnBackend.F16_arbitrary_seqlen + return FusedAttnBackend.No_Backend diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py index b276613f64e..7ad60182bb8 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py @@ -6,13 +6,20 @@ from __future__ import annotations -from dataclasses import dataclass, field -import importlib import threading +from dataclasses import dataclass, field from typing import Any, Dict, Hashable, Optional, Tuple import torch +from transformer_engine.common.cudnn_frontend import ( + build_cudnn_graph, + make_cudnn_graph, +) +from transformer_engine.common.cudnn_frontend import ( + import_cudnn_frontend as _import_cudnn_frontend, +) + _thread_state = threading.local() @@ -23,13 +30,10 @@ def import_cudnn_frontend(): Engine must not eagerly initialize cuDNN or fail on CPU-only processes. """ - try: - return importlib.import_module("cudnn") - except ImportError as exc: - raise ImportError( - "cuDNN Frontend Python package not found. Install " - "nvidia-cudnn-frontend>=1.27.0." - ) from exc + return _import_cudnn_frontend( + feature="PyTorch fused attention", + requirement="nvidia-cudnn-frontend>=1.28.0", + ) def _device_key(device: torch.device) -> Tuple[str, Optional[int]]: @@ -88,11 +92,10 @@ def make_graph(io_dtype: Any, device: torch.device, *, name: str): """Create an SDPA graph using FP32 intermediate and compute types.""" cudnn = import_cudnn_frontend() - return cudnn.pygraph( + return make_cudnn_graph( + cudnn, + io_dtype, name=name, - io_data_type=io_dtype, - intermediate_data_type=cudnn.data_type.FLOAT, - compute_data_type=cudnn.data_type.FLOAT, handle=current_stream_handle(device), ) @@ -101,15 +104,7 @@ def finalize_graph(graph) -> int: """Build a cuDNN graph and return its required workspace size.""" cudnn = import_cudnn_frontend() - graph.validate() - graph.build_operation_graph() - try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() - except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError(f"cuDNN attention graph is not supported: {exc}") from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) - return max(int(graph.get_workspace_size()), 1) + return build_cudnn_graph(cudnn, graph, description="attention") @dataclass diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 36df9ad416a..9a29ae478cb 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -11,15 +11,24 @@ from __future__ import annotations -from enum import IntEnum import math -from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple, Union +from enum import IntEnum +from typing import Any, Dict, List, Optional, Sequence, Tuple, Union import torch + # The extension is used only to reserve PyTorch's graph-safe Philox state; # graph construction, backend selection, and execution do not call TE common. import transformer_engine_torch as tex +from transformer_engine.common.attention.cudnn import normalize_attention_mask +from transformer_engine.common.attention.cudnn import ( + ragged_batch_bucket as _max_ragged_batch, +) +from transformer_engine.common.attention.cudnn import ( + ragged_token_bucket as _max_ragged_tokens, +) +from transformer_engine.common.attention.cudnn import round_up as _round_up from transformer_engine.pytorch.constants import ( DType, FP8BwdTensorIdx, @@ -157,30 +166,6 @@ def _format_stride(batch: int, heads: int, seqlen: int, dim: int, tensor_format: raise ValueError(f"Unsupported FP8 tensor format {tensor_format!r}.") -def _round_up(value: int, multiple: int) -> int: - return (value + multiple - 1) // multiple * multiple - - -def _max_ragged_tokens(num_tokens: int) -> int: - """Quantize THD token counts to the buckets used by TE's cuDNN path.""" - - if num_tokens <= 1024: - return 1024 - if num_tokens <= 32768: - return 1 << (num_tokens - 1).bit_length() - return _round_up(num_tokens, 32768) - - -def _max_ragged_batch(batch: int) -> int: - """Quantize THD batch sizes to the buckets used by TE's cuDNN path.""" - - if batch <= 32: - return 32 - if batch <= 512: - return 1 << (batch - 1).bit_length() - return _round_up(batch, 512) - - def _padded_sequence_lengths(cu_seqlens: torch.Tensor, batch: int) -> torch.Tensor: lengths = _sequence_lengths(cu_seqlens) if lengths.numel() == batch: @@ -455,38 +440,36 @@ def _mask_options( max_seqlen_q: int, max_seqlen_kv: int, ) -> Dict[str, Any]: - is_causal = attn_mask_type in ("causal", "padding_causal") - is_bottom_right = attn_mask_type in ( - "causal_bottom_right", - "padding_causal_bottom_right", - ) - is_padding = attn_mask_type in ( - "padding", - "padding_causal", - "padding_causal_bottom_right", + mask = normalize_attention_mask( + causal=attn_mask_type in ("causal", "padding_causal"), + bottom_right=attn_mask_type + in ("causal_bottom_right", "padding_causal_bottom_right"), + padding=attn_mask_type + in ("padding", "padding_causal", "padding_causal_bottom_right"), + bottom_right_diagonal=bottom_right_diagonal, + window_size=window_size, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, ) - if is_bottom_right and max_seqlen_q == max_seqlen_kv and not is_padding: - is_causal = True - is_bottom_right = False - bottom_right_diagonal = False options: Dict[str, Any] = { - "use_causal_mask": is_causal, - "use_causal_mask_bottom_right": is_bottom_right, + "use_causal_mask": mask.causal, + "use_causal_mask_bottom_right": mask.bottom_right, "diagonal_alignment": ( cudnn.diagonal_alignment.BOTTOM_RIGHT - if bottom_right_diagonal + if mask.bottom_right_diagonal else cudnn.diagonal_alignment.TOP_LEFT ), } - left, right = window_size - if left != -1: - options["diagonal_band_left_bound"] = left + 1 + if mask.window_left != -1: + options["diagonal_band_left_bound"] = mask.window_left + 1 # ``use_causal_mask`` already imposes a right bound of zero. The Python # frontend rejects specifying that same bound through both attributes. - if right != -1 and not ((is_causal or is_bottom_right) and right == 0): - options["diagonal_band_right_bound"] = right - options["is_padding"] = is_padding + if mask.window_right != -1 and not ( + (mask.causal or mask.bottom_right) and mask.window_right == 0 + ): + options["diagonal_band_right_bound"] = mask.window_right + options["is_padding"] = mask.padding return options diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index cc19ffd02a9..1e2bc101d6c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -5,11 +5,15 @@ """cuDNN-backed Flex Attention helpers.""" from dataclasses import dataclass -import inspect from typing import Any, Callable, Dict, Optional, Tuple import torch +from transformer_engine.common.attention.score_mod import ( + UNCACHEABLE_SCORE_MOD, + score_mod_callback_cache_key, +) + from ._cudnn_graph import ( current_stream_handle, finalize_graph, @@ -19,7 +23,7 @@ ) _cudnn_score_mod_graph_cache: Dict[Tuple[Any, ...], Any] = {} -_SCORE_MOD_UNCACHEABLE = object() +_SCORE_MOD_UNCACHEABLE = UNCACHEABLE_SCORE_MOD def _import_cudnn_frontend(): @@ -53,91 +57,13 @@ def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): # score_mod graph cache helpers. -def _freeze_score_mod_cache_key(value: Any) -> Any: - """Convert a user-provided score_mod graph key into a hashable structure.""" - if isinstance(value, torch.Tensor): - raise TypeError( - "score_mod_graph_cache_key() must not include tensors. Pass runtime tensors " - "through score_mod_tensors or score_mod_bprop_tensors instead." - ) - if isinstance(value, dict): - items = ( - ( - _freeze_score_mod_cache_key(key), - _freeze_score_mod_cache_key(val), - ) - for key, val in value.items() - ) - return tuple(sorted(items, key=repr)) - if isinstance(value, (list, tuple)): - return tuple(_freeze_score_mod_cache_key(item) for item in value) - if isinstance(value, (set, frozenset)): - items = (_freeze_score_mod_cache_key(item) for item in value) - return tuple(sorted(items, key=repr)) - try: - hash(value) - except TypeError as exc: - raise TypeError( - "score_mod_graph_cache_key() must return a hashable value or a nested " - "combination of dict/list/tuple/set values." - ) from exc - return value - - -def _score_mod_explicit_cache_key(callback_owner: Any) -> Optional[Any]: - """Return a user-provided structural graph key for a score_mod callback.""" - explicit_key = getattr(callback_owner, "score_mod_graph_cache_key", None) - if explicit_key is None: - return None - explicit_key = explicit_key() if callable(explicit_key) else explicit_key - return _freeze_score_mod_cache_key(explicit_key) - - def _score_mod_callback_cache_key(callback: Optional[Callable]) -> Any: - """Create a stable graph cache key for a score_mod callable. - - Module-level named functions are assumed to have stable topology. Anonymous functions - are keyed by code object because lambdas in the same module can share the same - qualname. Stateful bound methods and callable instances need an explicit - score_mod_graph_cache_key(); otherwise their graphs are left uncached to avoid reusing - stale graphs after Python object address reuse. - """ - if callback is None: - return None - self_obj = getattr(callback, "__self__", None) - func_obj = getattr(callback, "__func__", None) - if self_obj is not None and func_obj is not None: - explicit_key = _score_mod_explicit_cache_key(self_obj) - if explicit_key is None: - return _SCORE_MOD_UNCACHEABLE - return ( - "bound_method", - type(self_obj), - func_obj.__module__, - func_obj.__qualname__, - explicit_key, - ) - - explicit_key = _score_mod_explicit_cache_key(callback) - if explicit_key is not None: - return ( - "callable", - type(callback), - getattr(callback, "__module__", None), - getattr(callback, "__qualname__", None), - explicit_key, - ) + """Compatibility wrapper around the shared score-modification key policy.""" - if ( - inspect.isfunction(callback) - and callback.__closure__ is None - and "" not in callback.__qualname__ - ): - if callback.__name__ == "" or not callback.__qualname__: - return ("function", callback.__module__, callback.__code__) - return ("function", callback.__module__, callback.__qualname__) - - return _SCORE_MOD_UNCACHEABLE + return score_mod_callback_cache_key( + callback, + is_array=lambda item: isinstance(item, torch.Tensor), + ) def _score_mod_device_key(device: torch.device) -> Tuple[Any, ...]: @@ -289,8 +215,8 @@ def _cudnn_score_mod_fwd_cache_key( ) -> Optional[Tuple[Any, ...]]: """Pre-build cache key for score_mod fprop execution plans. - cuDNN exposes graph.key(), but only after graph construction has run the user callback. - This key avoids rebuilding the Python graph on cache hits. + cuDNN exposes graph.key(), but only after graph construction has run the user + callback. This key avoids rebuilding the Python graph on cache hits. """ score_mod_key = _score_mod_callback_cache_key(score_mod) if score_mod_key is _SCORE_MOD_UNCACHEABLE: @@ -668,7 +594,8 @@ def backward(ctx, d_out: torch.Tensor): # pylint: disable=missing-function-docstring if not ctx.is_training: raise RuntimeError( - "score_mod backward requires DotProductAttention to be in training mode." + "score_mod backward requires DotProductAttention to be in " + "training mode." ) saved_tensors = ctx.saved_tensors From b7b2cac491a38a3d0d2a2e85afff3f3f69995ace Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 00:46:55 +0000 Subject: [PATCH 16/36] [JAX] Add FP8 cuDNN attention parity Implement graph-backed JAX dot-product attention for DelayedScaling, Float8CurrentScaling, and MXFP8BlockScaling recipes. Quantize Q/K/V and dO internally while preserving FP16/BF16 module inputs, outputs, and gradients, and propagate delayed-scaling state through the custom VJP. Factor FP8 backend selection, layout parsing, mask normalization, graph operation dispatch, tensor strides, MXFP8 padding, and scale swizzling into shared helpers. Migrate the PyTorch cuDNN attention path to these helpers without intentionally changing its supported configurations. Align the JAX frontend with PyTorch by exposing cuDNN ALiBi attention, supporting explicit bottom-right diagonal alignment, and consolidating previously framework-specific compatibility policy flags. Add common policy and graph-dispatch tests, JAX numerical forward/backward coverage for all three FP8 recipes across FP16 and BF16 boundaries, MXFP8 scale-layout coverage, and ALiBi and bottom-right parity tests. Wire the common suite into both L0 framework jobs and the JAX FP8 suite into L0 JAX CI. Signed-off-by: Vladimir Cherepanov --- qa/L0_jax_unittest/test.sh | 4 +- qa/L0_pytorch_unittest/test.sh | 1 + tests/jax/test_fp8_fused_attn.py | 257 ++++ tests/test_common_attention_helpers.py | 165 ++- transformer_engine/common/attention/cudnn.py | 177 ++- transformer_engine/common/attention/fp8.py | 171 +++ transformer_engine/jax/attention.py | 90 +- .../jax/cpp_extensions/__init__.py | 1 + .../jax/cpp_extensions/attention.py | 26 +- .../jax/cpp_extensions/cudnn_attention.py | 43 +- .../jax/cpp_extensions/cudnn_graph.py | 6 + .../jax/cpp_extensions/fp8_attention.py | 1238 +++++++++++++++++ transformer_engine/jax/cpp_extensions/gemm.py | 24 +- .../jax/csrc/extensions/pybind.cpp | 3 +- transformer_engine/jax/flax/module.py | 18 + transformer_engine/jax/flax/transformer.py | 26 +- transformer_engine/jax/quantize/helper.py | 21 +- transformer_engine/jax/quantize/quantizer.py | 24 + .../dot_product_attention/_cudnn_backend.py | 139 +- .../dot_product_attention/cudnn_attention.py | 219 ++- 20 files changed, 2352 insertions(+), 301 deletions(-) create mode 100644 tests/jax/test_fp8_fused_attn.py create mode 100644 transformer_engine/common/attention/fp8.py create mode 100644 transformer_engine/jax/cpp_extensions/fp8_attention.py diff --git a/qa/L0_jax_unittest/test.sh b/qa/L0_jax_unittest/test.sh index ad1157cbadc..897c2c7bd2c 100644 --- a/qa/L0_jax_unittest/test.sh +++ b/qa/L0_jax_unittest/test.sh @@ -28,7 +28,9 @@ pip3 install pytest==8.2.1 pytest-timeout==2.4.0 || error_exit "Failed to instal : ${XML_LOG_DIR:=/logs} mkdir -p "$XML_LOG_DIR" -python3 -m pytest -c $TE_PATH/tests/jax/pytest.ini -v --junitxml=$XML_LOG_DIR/pytest_jax_not_distributed.xml $TE_PATH/tests/jax --ignore=$TE_PATH/tests/jax/test_multi_process_ep.py -k 'not distributed' || test_fail "tests/jax/*not_distributed_*" +python3 -m pytest --import-mode=importlib -v --junitxml=$XML_LOG_DIR/pytest_test_common_attention_helpers.xml $TE_PATH/tests/test_common_attention_helpers.py || test_fail "test_common_attention_helpers.py" +python3 -m pytest -c $TE_PATH/tests/jax/pytest.ini -v --junitxml=$XML_LOG_DIR/pytest_jax_not_distributed.xml $TE_PATH/tests/jax --ignore=$TE_PATH/tests/jax/test_multi_process_ep.py --ignore=$TE_PATH/tests/jax/test_fp8_fused_attn.py -k 'not distributed' || test_fail "tests/jax/*not_distributed_*" +python3 -m pytest -c $TE_PATH/tests/jax/pytest.ini -v --junitxml=$XML_LOG_DIR/pytest_jax_fp8_fused_attn.xml $TE_PATH/tests/jax/test_fp8_fused_attn.py || test_fail "tests/jax/test_fp8_fused_attn.py" python3 -m pytest -c $TE_PATH/tests/jax/pytest.ini -v --junitxml=$XML_LOG_DIR/pytest_jax_fused_attn_score_mod.xml $TE_PATH/tests/jax/test_fused_attn_score_mod.py || test_fail "tests/jax/test_fused_attn_score_mod.py" NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 python3 -m pytest -c $TE_PATH/tests/jax/pytest.ini -v --junitxml=$XML_LOG_DIR/pytest_jax_fused_attn_with_determinism.xml $TE_PATH/tests/jax/test_fused_attn.py -k "TestFusedAttnWithDeterminism" || test_fail "tests/jax/test_fused_attn.py" diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index a78a99d7f95..bb868197ea4 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -29,6 +29,7 @@ export NVTE_FLASH_ATTN_V4=0 pip3 install pytest==8.2.1 || error_exit "Failed to install pytest" +python3 -m pytest --import-mode=importlib --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_common_attention_helpers.xml $TE_PATH/tests/test_common_attention_helpers.py || test_fail "test_common_attention_helpers.py" NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_sanity.xml $TE_PATH/tests/pytorch/test_sanity.py || test_fail "test_sanity.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_recipe.xml $TE_PATH/tests/pytorch/test_recipe.py || test_fail "test_recipe.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_custom_recipe.xml $TE_PATH/tests/pytorch/test_custom_recipe.py || test_fail "test_custom_recipe.py" diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py new file mode 100644 index 00000000000..0817eb034da --- /dev/null +++ b/tests/jax/test_fp8_fused_attn.py @@ -0,0 +1,257 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Coverage for JAX FP8 DPA and attention features shared with PyTorch.""" + +from math import sqrt + +import jax +import jax.numpy as jnp +import numpy as np +import pytest +from transformer_engine_jax import get_cudnn_version, get_device_compute_capability + +from transformer_engine.common import recipe +from transformer_engine.jax import autocast +from transformer_engine.jax.attention import ( + AttnBiasType, + AttnMaskType, + AttnSoftmaxType, + QKVLayout, + SequenceDescriptor, + fused_attn, +) +from transformer_engine.jax.cpp_extensions import FusedAttnHelper +from transformer_engine.jax.cpp_extensions.fp8_attention import _mxfp8_scale_inv +from transformer_engine.jax.flax import DotProductAttention +from transformer_engine.jax.quantize import ( + BlockScaleQuantizer, + QuantizeLayout, + ScalingMode, +) +from transformer_engine.jax.sharding import MeshResource + + +def _require_gpu(min_arch=90, min_cudnn=90700): + try: + if not any(device.platform == "gpu" for device in jax.devices()): + pytest.skip("A CUDA device is required.") + arch = get_device_compute_capability(0) + except RuntimeError as exc: + pytest.skip(f"A usable CUDA device is required: {exc}") + if arch < min_arch: + pytest.skip(f"This test requires SM{min_arch} or newer, found SM{arch}.") + cudnn_version = get_cudnn_version() + if cudnn_version < min_cudnn: + pytest.skip(f"This test requires cuDNN {min_cudnn}, found {cudnn_version}.") + if cudnn_version == 91000: + pytest.skip("cuDNN 9.10.0 has known FP8 SDPA issues.") + return arch + + +def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): + q_seqlen, kv_seqlen = q.shape[1], k.shape[1] + scores = jnp.einsum( + "bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32) + ) / sqrt(q.shape[-1]) + q_pos = jnp.arange(q_seqlen)[:, None] + kv_pos = jnp.arange(kv_seqlen)[None, :] + shift = kv_seqlen - q_seqlen if bottom_right else 0 + if alibi: + heads = q.shape[-2] + power_of_two_heads = 2 ** int(np.floor(np.log2(heads))) + base = 2.0 ** (-8.0 / power_of_two_heads) + slopes = base ** jnp.arange(1, power_of_two_heads + 1, dtype=jnp.float32) + if power_of_two_heads < heads: + extra_base = 2.0 ** (-4.0 / power_of_two_heads) + extra = extra_base ** jnp.arange( + 1, 1 + 2 * (heads - power_of_two_heads), 2, dtype=jnp.float32 + ) + slopes = jnp.concatenate((slopes, extra)) + distance = jnp.abs(q_pos + shift - kv_pos).astype(jnp.float32) + scores -= slopes[None, :, None, None] * distance[None, None, :, :] + allowed = kv_pos <= q_pos + shift + scores = jnp.where(allowed[None, None, :, :], scores, -jnp.inf) + probabilities = jax.nn.softmax(scores, axis=-1) + return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype( + q.dtype + ) + + +def _assert_fp8_close(actual, expected): + actual = np.asarray(actual, dtype=np.float32) + expected = np.asarray(expected, dtype=np.float32) + np.testing.assert_allclose(actual, expected, atol=0.5, rtol=0.05) + assert np.sqrt(np.mean(np.square(actual - expected))) < 0.11 + + +@pytest.mark.parametrize( + "fp8_recipe,min_arch,min_cudnn", + ( + pytest.param( + recipe.DelayedScaling(amax_history_len=1, fp8_dpa=True), + 90, + 90700, + id="delayed", + ), + pytest.param( + recipe.Float8CurrentScaling(fp8_dpa=True), + 90, + 90700, + id="current", + ), + pytest.param( + recipe.MXFP8BlockScaling(fp8_dpa=True), + 100, + 92100, + id="mxfp8", + ), + ), +) +@pytest.mark.parametrize( + "input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16") +) +def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): + """Each supported recipe executes FP8 DPA behind FP16/BF16 module boundaries.""" + + _require_gpu(min_arch, min_cudnn) + batch, seqlen, heads, dim = 2, 128, 8, 128 + q_key, k_key, v_key, do_key = jax.random.split(jax.random.PRNGKey(1234), 4) + shape = (batch, seqlen, heads, dim) + q = jax.random.uniform(q_key, shape, input_dtype, minval=-0.5, maxval=0.5) + k = jax.random.uniform(k_key, shape, input_dtype, minval=-0.5, maxval=0.5) + v = jax.random.uniform(v_key, shape, input_dtype, minval=-0.5, maxval=0.5) + doutput = jax.random.uniform(do_key, shape, input_dtype, minval=-0.5, maxval=0.5) + seqlens = jnp.full((batch,), seqlen, dtype=jnp.int32) + descriptor = SequenceDescriptor.from_seqlens((seqlens, seqlens)) + module = DotProductAttention( + head_dim=dim, + num_attention_heads=heads, + num_gqa_groups=heads, + attn_mask_type="causal", + qkv_layout="bshd_bshd_bshd", + transpose_batch_sequence=False, + ) + + def loss_fn(variables, query, key, value): + output = module.apply( + variables, query, key, value, descriptor, deterministic=True + ) + loss = jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)) + return loss, output + + def reference_loss(query, key, value): + output = _reference_attention(query, key, value) + loss = jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)) + return loss, output + + with autocast(enabled=True, recipe=fp8_recipe, mesh_resource=MeshResource()): + variables = module.init( + jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=True + ) + (_, output), (_, dq, dk, dv) = jax.value_and_grad( + loss_fn, argnums=(0, 1, 2, 3), has_aux=True + )(variables, q, k, v) + (_, reference), (dq_ref, dk_ref, dv_ref) = jax.value_and_grad( + reference_loss, argnums=(0, 1, 2), has_aux=True + )(q, k, v) + + assert output.dtype == q.dtype + _assert_fp8_close(output, reference) + for actual, expected in zip((dq, dk, dv), (dq_ref, dk_ref, dv_ref)): + _assert_fp8_close(actual, expected) + + +def test_mxfp8_attention_scale_layout(): + """Attention pads, permutes, and swizzles compact JAX MXFP8 scales for cuDNN.""" + + quantizer = BlockScaleQuantizer( + q_dtype=jnp.float8_e4m3fn, + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + q_layout=QuantizeLayout.ROWWISE_COLWISE, + data_layout="NN", + ) + tensor = quantizer.quantize( + jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2 + ) + assert tensor.rowwise_tensor.scale_inv.shape == (2, 64, 8, 2) + assert tensor.colwise_tensor.scale_inv.shape == (2, 2, 8, 64) + assert _mxfp8_scale_inv(tensor).shape == (2, 8, 128, 4) + assert _mxfp8_scale_inv(tensor, colwise=True).shape == (2, 8, 4, 128) + + +@pytest.mark.parametrize("feature", ("alibi", "bottom_right")) +def test_fused_attention_parity_features(feature): + """JAX executes ALiBi and explicit bottom-right diagonal attention.""" + + _require_gpu(90, 90700) + batch, heads, dim = 2, 8, 64 + q_seqlen, kv_seqlen = (128, 128) if feature == "alibi" else (64, 128) + q_key, k_key, v_key, do_key = jax.random.split(jax.random.PRNGKey(4321), 4) + q = jax.random.normal(q_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 + k = jax.random.normal(k_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 + v = jax.random.normal(v_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 + doutput = ( + jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 + ) + q_lengths = jnp.full((batch,), q_seqlen, dtype=jnp.int32) + kv_lengths = jnp.full((batch,), kv_seqlen, dtype=jnp.int32) + descriptor = SequenceDescriptor.from_seqlens((q_lengths, kv_lengths)) + bias_type = AttnBiasType.ALIBI if feature == "alibi" else AttnBiasType.NO_BIAS + helper = FusedAttnHelper( + True, + q.dtype, + k.dtype, + QKVLayout.BSHD_BSHD_BSHD, + bias_type, + AttnMaskType.CAUSAL_MASK, + AttnSoftmaxType.VANILLA_SOFTMAX, + 0.0, + heads, + heads, + q_seqlen, + kv_seqlen, + dim, + dim, + (-1, -1), + ) + if not helper.is_fused_attn_kernel_available(): + pytest.skip("No fused-attention kernel supports this configuration.") + + def te_loss(query, key, value): + output = fused_attn( + (query, key, value), + None, + descriptor, + None, + bias_type, + AttnMaskType.CAUSAL_MASK, + QKVLayout.BSHD_BSHD_BSHD, + AttnSoftmaxType.VANILLA_SOFTMAX, + 1.0 / sqrt(dim), + 0.0, + True, + bottom_right_diagonal=feature == "bottom_right", + ) + return jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)), output + + def reference_loss(query, key, value): + output = _reference_attention( + query, + key, + value, + bottom_right=feature == "bottom_right", + alibi=feature == "alibi", + ) + return jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)), output + + (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)( + q, k, v + ) + (_, reference), reference_grads = jax.value_and_grad( + reference_loss, argnums=(0, 1, 2), has_aux=True + )(q, k, v) + np.testing.assert_allclose(output, reference, atol=0.02, rtol=0.02) + for actual, expected in zip(grads, reference_grads): + np.testing.assert_allclose(actual, expected, atol=0.03, rtol=0.03) diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py index 0f9c3d3e41c..d44aeb74caa 100644 --- a/tests/test_common_attention_helpers.py +++ b/tests/test_common_attention_helpers.py @@ -13,11 +13,21 @@ AttentionLayout, FusedAttentionConfig, check_f16_fused_attention_support, + check_fp8_fused_attention_support, + cudnn_mask_options, encode_cudnn_version, normalize_attention_mask, + parse_attention_layout, ragged_batch_bucket, ragged_token_bucket, ) +from transformer_engine.common.attention.fp8 import ( + FP8AttentionGraphConfig, + attention_format_stride, + build_fp8_backward_operation, + build_fp8_forward_operation, + mxfp8_padded_sizes, +) from transformer_engine.common.attention.score_mod import ( UNCACHEABLE_SCORE_MOD, freeze_score_mod_cache_key, @@ -95,20 +105,16 @@ def test_shared_f16_policy_basic_support_and_rejection(): assert "architecture" in unsupported.reason -def test_shared_f16_policy_explicit_frontend_capabilities(): +def test_shared_f16_policy_framework_parity_features(): alibi = _attention_config(bias_type="alibi", mask_type="causal") - assert not check_f16_fused_attention_support(alibi).supported - assert check_f16_fused_attention_support(replace(alibi, allow_alibi=True)).supported + assert check_f16_fused_attention_support(alibi).supported extended_causal = _attention_config( mask_type="causal", window_size=(128, 64), sm_arch=100, ) - assert not check_f16_fused_attention_support(extended_causal).supported - assert check_f16_fused_attention_support( - replace(extended_causal, allow_extended_causal_window=True) - ).supported + assert check_f16_fused_attention_support(extended_causal).supported modern_padding = _attention_config( mask_type="padding", @@ -116,9 +122,150 @@ def test_shared_f16_policy_explicit_frontend_capabilities(): cudnn_version=(9, 7, 0), ) assert check_f16_fused_attention_support(modern_padding).supported - assert not check_f16_fused_attention_support( - replace(modern_padding, modern_mask_rules_override=True) + + +def test_shared_fp8_policy(): + fp8 = _attention_config(q_dtype="float8_e4m3", kv_dtype="float8_e4m3", sm_arch=100) + assert check_fp8_fused_attention_support(fp8).supported + assert check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 10, 1)) + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 10, 0)) + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, bias_type="alibi") + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, return_max_logit=True) ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, head_dim_qk=200) + ).supported + + +@pytest.mark.parametrize( + "layout, expected", + [ + ("bs3hd", ("bshd", "bshd", "3hd")), + ("bshd_bs2hd", ("bshd", "bshd", "hd_2hd")), + ("bhsd_bhsd_bhsd", ("bhsd", "bhsd", "sd_sd_sd")), + ("paged_kv_bshd_bshd_bshd", ("bshd", "bshd", "paged_separate")), + ], +) +def test_parse_attention_layout(layout, expected): + parsed = parse_attention_layout(layout) + assert (parsed.q_format, parsed.kv_format, parsed.layout_group) == expected + + +def test_shared_mask_options_use_modern_band_api(): + options = cudnn_mask_options( + causal=True, + bottom_right=False, + padding=False, + bottom_right_diagonal=True, + window_size=(32, -1), + max_seqlen_q=64, + max_seqlen_kv=128, + cudnn_version=(9, 6, 0), + ) + assert options == { + "diagonal_alignment": "bottom_right", + "is_padding": False, + "diagonal_band_left_bound": 33, + "diagonal_band_right_bound": 0, + } + + +def test_shared_fp8_shape_helpers(): + assert attention_format_stride(2, 8, 128, 64, "bshd") == (65536, 64, 512, 1) + assert attention_format_stride(2, 8, 128, 64, "bhsd") == (65536, 8192, 64, 1) + assert mxfp8_padded_sizes(129, 33, 160, 96) == { + "s_q_padded": 256, + "s_kv_padded": 128, + "s_q_scale_padded": 8, + "s_kv_scale_padded": 4, + "d_qk_padded": 256, + "d_v_padded": 128, + "d_qk_scale_padded": 8, + "d_v_scale_padded": 4, + } + + +class _FakeFP8Graph: + def __init__(self): + self.call = None + + def sdpa_fp8(self, *args, **kwargs): + self.call = ("fp8_fwd", args, kwargs) + return "o", "stats", "amax_s", "amax_o" + + def sdpa_mxfp8(self, *args, **kwargs): + self.call = ("mx_fwd", args, kwargs) + return "o", "stats", "amax_o" + + def sdpa_fp8_backward(self, *args, **kwargs): + self.call = ("fp8_bwd", args, kwargs) + return "dq", "dk", "dv", "aq", "ak", "av", "ap" + + def sdpa_mxfp8_backward(self, *args, **kwargs): + self.call = ("mx_bwd", args, kwargs) + return "dq", "dk", "dv", "aq", "ak", "av" + + +def test_shared_fp8_graph_operation_dispatch(): + graph = _FakeFP8Graph() + forward = build_fp8_forward_operation( + graph, + { + "q": 1, + "k": 2, + "v": 3, + "descale_q": 4, + "descale_k": 5, + "descale_v": 6, + "descale_s": 7, + "scale_s": 8, + "scale_o": 9, + }, + {"attn_scale": 0.125}, + FP8AttentionGraphConfig("delayed", "forward"), + ) + assert forward == { + "output": "o", + "stats": "stats", + "amax_s": "amax_s", + "amax_o": "amax_o", + } + assert graph.call[0] == "fp8_fwd" + + backward = build_fp8_backward_operation( + graph, + { + **{name: name for name in ("q", "k", "v", "o", "do", "stats")}, + **{ + name: name + for name in ( + "descale_q", + "descale_k", + "descale_v", + "descale_o", + "descale_do", + "descale_s", + "descale_dp", + "scale_s", + "scale_dq", + "scale_dk", + "scale_dv", + "scale_dp", + ) + }, + }, + {"attn_scale": 0.125}, + FP8AttentionGraphConfig("current", "backward"), + ) + assert backward["amax_dp"] == "ap" + assert graph.call[0] == "fp8_bwd" def _not_an_array(_value): diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index 7194d50f9de..cb2cce09775 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -50,9 +50,6 @@ class FusedAttentionConfig: deterministic: bool cudnn_version: tuple[int, int, int] sm_arch: int - allow_alibi: bool = False - allow_extended_causal_window: bool = False - modern_mask_rules_override: bool = False @dataclass(frozen=True) @@ -76,6 +73,38 @@ class AttentionMask: window_right: int +def parse_attention_layout(qkv_layout: str) -> AttentionLayout: + """Normalize a TE QKV layout string for framework-independent policy checks.""" + + paged = qkv_layout.startswith("paged_kv_") + layout = qkv_layout.removeprefix("paged_kv_") + components = layout.split("_") + + def tensor_format(component: str) -> str: + return "".join(char for char in component if char.isalpha()) + + q_format = tensor_format(components[0]) + kv_format = tensor_format(components[-1]) if len(components) > 1 else q_format + qkv_format = q_format if q_format == kv_format else f"{q_format}_2{kv_format}" + if paged: + layout_group = "paged_separate" + elif len(components) == 1 and "3" in components[0]: + layout_group = "h3d" if "h3d" in components[0] else "3hd" + elif len(components) == 2: + layout_group = "hd_h2d" if "h2d" in components[1] else "hd_2hd" + elif q_format == "bhsd": + layout_group = "sd_sd_sd" + else: + layout_group = "separate" + return AttentionLayout( + qkv_format=qkv_format, + q_format=q_format, + kv_format=kv_format, + layout_group=layout_group, + is_qkvpacked=layout_group in ("3hd", "h3d"), + ) + + def encode_cudnn_version(version: tuple[int, int, int]) -> int: """Encode a cuDNN backend version using its native integer convention.""" @@ -138,6 +167,48 @@ def normalize_attention_mask( ) +def cudnn_mask_options( + *, + causal: bool, + bottom_right: bool, + padding: bool, + bottom_right_diagonal: bool, + window_size: tuple[int, int], + max_seqlen_q: int, + max_seqlen_kv: int, + cudnn_version: tuple[int, int, int], +) -> dict[str, bool | int | str]: + """Return canonical cuDNN SDPA mask options using framework-neutral values.""" + + mask = normalize_attention_mask( + causal=causal, + bottom_right=bottom_right, + padding=padding, + bottom_right_diagonal=bottom_right_diagonal, + window_size=window_size, + max_seqlen_q=max_seqlen_q, + max_seqlen_kv=max_seqlen_kv, + ) + version = encode_cudnn_version(cudnn_version) + options: dict[str, bool | int | str] = { + "diagonal_alignment": ( + "bottom_right" if mask.bottom_right_diagonal else "top_left" + ), + "is_padding": mask.padding, + } + if version < 90600: + options["use_causal_mask"] = mask.causal + options["use_causal_mask_bottom_right"] = mask.bottom_right + if version >= 90200 and mask.window_left != -1: + options["diagonal_band_left_bound"] = mask.window_left + 1 + if version >= 90600: + if mask.window_right != -1: + options["diagonal_band_right_bound"] = mask.window_right + elif mask.causal or mask.bottom_right: + options["diagonal_band_right_bound"] = 0 + return options + + def requires_64bit_ragged_offset( layout: AttentionLayout, num_attn_heads: int, @@ -290,8 +361,7 @@ def check_f16_fused_attention_support( ) alibi_supported = ( - config.allow_alibi - and bias == "alibi" + bias == "alibi" and version >= 8906 and arch >= 90 and mask @@ -373,11 +443,7 @@ def check_f16_fused_attention_support( and sq <= skv ) ) - mask_ok = ( - modern_mask_ok - if config.modern_mask_rules_override - else mask_ok or modern_mask_ok - ) + mask_ok = mask_ok or modern_mask_ok if not mask_ok: return _unsupported("attention mask is not supported") if mask in ("padding", "padding_causal") and bias == "post_scale_bias": @@ -437,7 +503,7 @@ def check_f16_fused_attention_support( "causal_bottom_right", "padding_causal_bottom_right", ) - or (config.allow_extended_causal_window and mask == "causal") + or mask == "causal" ) ) and sq <= skv @@ -460,8 +526,8 @@ def check_f16_fused_attention_support( ) if requires_i64 and version < 90500: return _unsupported("ragged offsets require int64 support") - if version in (91000, 91001): - return _unsupported("cuDNN 9.10.0 and 9.10.1 have known SDPA issues") + if version == 91000: + return _unsupported("cuDNN 9.10.0 has known SDPA issues") if version < 91301 and softmax != "vanilla": return _unsupported("this softmax type requires cuDNN 9.13.1 or newer") if config.return_max_logit and version < 92100: @@ -513,3 +579,88 @@ def check_f16_fused_attention_support( ) return FusedAttentionSupport(True) + + +def check_fp8_fused_attention_support( + config: FusedAttentionConfig, +) -> FusedAttentionSupport: + """Check the shared cuDNN FP8/MXFP8 fused-attention compatibility policy.""" + + if config.q_dtype != config.kv_dtype: + return _unsupported("Q and KV must have the same data type") + if config.q_dtype not in ("float8_e4m3", "float8_e5m2"): + return _unsupported("only FP8 E4M3 and E5M2 are supported") + + version = encode_cudnn_version(config.cudnn_version) + arch = int(config.sm_arch) + layout = config.layout + sq = int(config.max_seqlen_q) + skv = int(config.max_seqlen_kv) + dqk = int(config.head_dim_qk) + dv = int(config.head_dim_v) + mask = config.mask_type + + if arch < 90: + return _unsupported("FP8 attention requires SM90 or newer") + if config.bias_type != "no_bias": + return _unsupported("FP8 attention does not support attention bias") + if config.return_max_logit: + return _unsupported("FP8 attention does not support returning max logits") + if version == 91000: + return _unsupported("cuDNN 9.10.0 has known SDPA issues") + if requires_64bit_ragged_offset( + layout, + config.num_attn_heads, + config.num_gqa_groups, + sq, + skv, + dqk, + dv, + ): + return _unsupported("FP8 attention does not support 64-bit ragged offsets") + + shape_mask_ok = ( + ( + version >= 90201 + and arch < 100 + and sq % 128 == 0 + and skv % 128 == 0 + and dqk == 128 + and dv == 128 + and mask in ("causal", "no_mask") + ) + or ( + version >= 90700 + and ( + (arch < 100 and not config.is_training and dqk <= 256 and dv <= 256) + or (arch < 100 and config.is_training and dqk == 128 and dv == 128) + or (arch >= 100 and dqk <= 128 and dv <= 128) + ) + and dqk % 16 == 0 + and dv % 16 == 0 + and mask in ("no_mask", "causal", "padding", "padding_causal") + ) + or ( + version >= 92100 + and arch >= 100 + and dqk <= 192 + and dv <= 128 + and dqk % 16 == 0 + and dv % 16 == 0 + and mask in ("no_mask", "causal", "causal_bottom_right") + ) + ) + if not shape_mask_ok: + return _unsupported("FP8 attention shape or mask is not supported") + + format_softmax_ok = ( + version < 92100 + and layout.qkv_format in ("bshd", "sbhd") + and config.softmax_type == "vanilla" + ) or ( + version >= 92100 + and layout.qkv_format in ("bshd", "sbhd", "bhsd") + ) + if not format_softmax_ok: + return _unsupported("FP8 attention layout or softmax type is not supported") + return FusedAttentionSupport(True) diff --git a/transformer_engine/common/attention/fp8.py b/transformer_engine/common/attention/fp8.py new file mode 100644 index 00000000000..5675a68d82c --- /dev/null +++ b/transformer_engine/common/attention/fp8.py @@ -0,0 +1,171 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Framework-neutral helpers for cuDNN FP8 attention graph construction.""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +from .cudnn import round_up + + +@dataclass(frozen=True) +class FP8AttentionGraphConfig: + """Static choices that select a cuDNN FP8 attention graph family.""" + + mode: str + name: str + + def __post_init__(self): + if self.mode not in ("delayed", "current", "mxfp8"): + raise ValueError(f"Unknown FP8 attention scaling mode {self.mode!r}.") + + @property + def is_mxfp8(self) -> bool: + """Return whether the graph uses microscaling FP8 nodes.""" + + return self.mode == "mxfp8" + + +def attention_format_stride( + batch: int, heads: int, seqlen: int, dim: int, tensor_format: str +) -> tuple[int, int, int, int]: + """Describe a contiguous TE attention tensor as logical BHSD.""" + + if tensor_format in ("bshd", "thd"): + return (seqlen * heads * dim, dim, heads * dim, 1) + if tensor_format == "sbhd": + return (heads * dim, dim, batch * heads * dim, 1) + if tensor_format == "bhsd": + return (heads * seqlen * dim, seqlen * dim, dim, 1) + raise ValueError(f"Unsupported FP8 tensor format {tensor_format!r}.") + + +def mxfp8_padded_sizes(s_q: int, s_kv: int, d_qk: int, d_v: int) -> dict[str, int]: + """Return the padded data and E8M0-scale dimensions required by cuDNN MXFP8.""" + + return { + "s_q_padded": round_up(s_q, 128), + "s_kv_padded": round_up(s_kv, 128), + "s_q_scale_padded": round_up((s_q + 31) // 32, 4), + "s_kv_scale_padded": round_up((s_kv + 31) // 32, 4), + "d_qk_padded": round_up(d_qk, 128), + "d_v_padded": round_up(d_v, 128), + "d_qk_scale_padded": round_up((d_qk + 31) // 32, 4), + "d_v_scale_padded": round_up((d_v + 31) // 32, 4), + } + + +def build_fp8_forward_operation( + graph: Any, + tensors: Mapping[str, Any], + options: Mapping[str, Any], + config: FP8AttentionGraphConfig, +) -> dict[str, Any]: + """Add the selected FP8 SDPA forward operation to a cuDNN frontend graph.""" + + kwargs = dict(options) + if config.is_mxfp8: + output, stats, amax_o = graph.sdpa_mxfp8( + tensors["q"], + tensors["k"], + tensors["v"], + tensors["descale_q"], + tensors["descale_k"], + tensors["descale_v"], + name=config.name, + **kwargs, + ) + return {"output": output, "stats": stats, "amax_o": amax_o} + + output, stats, amax_s, amax_o = graph.sdpa_fp8( + tensors["q"], + tensors["k"], + tensors["v"], + tensors["descale_q"], + tensors["descale_k"], + tensors["descale_v"], + tensors["descale_s"], + tensors["scale_s"], + tensors["scale_o"], + name=config.name, + **kwargs, + ) + return { + "output": output, + "stats": stats, + "amax_s": amax_s, + "amax_o": amax_o, + } + + +def build_fp8_backward_operation( + graph: Any, + tensors: Mapping[str, Any], + options: Mapping[str, Any], + config: FP8AttentionGraphConfig, +) -> dict[str, Any]: + """Add the selected FP8 SDPA backward operation to a cuDNN frontend graph.""" + + kwargs = dict(options) + if config.is_mxfp8: + outputs = graph.sdpa_mxfp8_backward( + tensors["q"], + tensors["q_t"], + tensors["k"], + tensors["k_t"], + tensors["v"], + tensors["o"], + tensors["do_f16"], + tensors["do"], + tensors["do_t"], + tensors["stats"], + tensors["descale_q"], + tensors["descale_q_t"], + tensors["descale_k"], + tensors["descale_k_t"], + tensors["descale_v"], + tensors["descale_do"], + tensors["descale_do_t"], + name=config.name, + **kwargs, + ) + dq, dk, dv, *amax = outputs + return {"dq": dq, "dk": dk, "dv": dv, "amax": tuple(amax)} + + outputs = graph.sdpa_fp8_backward( + tensors["q"], + tensors["k"], + tensors["v"], + tensors["o"], + tensors["do"], + tensors["stats"], + tensors["descale_q"], + tensors["descale_k"], + tensors["descale_v"], + tensors["descale_o"], + tensors["descale_do"], + tensors["descale_s"], + tensors["descale_dp"], + tensors["scale_s"], + tensors["scale_dq"], + tensors["scale_dk"], + tensors["scale_dv"], + tensors["scale_dp"], + name=config.name, + **kwargs, + ) + dq, dk, dv, amax_dq, amax_dk, amax_dv, amax_dp = outputs + return { + "dq": dq, + "dk": dk, + "dv": dv, + "amax_dq": amax_dq, + "amax_dk": amax_dk, + "amax_dv": amax_dv, + "amax_dp": amax_dp, + } diff --git a/transformer_engine/jax/attention.py b/transformer_engine/jax/attention.py index 6d7c823e122..40fac63393b 100644 --- a/transformer_engine/jax/attention.py +++ b/transformer_engine/jax/attention.py @@ -20,6 +20,7 @@ from transformer_engine_jax import NVTE_Softmax_Type from . import cpp_extensions as tex +from .quantize import AttentionQuantizerSet class AttnBiasType(Enum): @@ -32,6 +33,8 @@ class AttnBiasType(Enum): NO_BIAS = NVTE_Bias_Type.NVTE_NO_BIAS PRE_SCALE_BIAS = NVTE_Bias_Type.NVTE_PRE_SCALE_BIAS POST_SCALE_BIAS = NVTE_Bias_Type.NVTE_POST_SCALE_BIAS + # Keep source-tree imports working before the local extension is rebuilt. + ALIBI = getattr(NVTE_Bias_Type, "NVTE_ALIBI", 3) class AttnMaskType(Enum): @@ -1056,6 +1059,7 @@ def _legacy_fused_attn( context_parallel_axis: str = "", softmax_offset: Optional[jnp.ndarray] = None, return_max_logit: bool = False, + bottom_right_diagonal: Optional[bool] = None, ): """ Perform non-THD (non-packed) cuDNN fused attention. @@ -1150,6 +1154,11 @@ def _legacy_fused_attn( context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, return_max_logit=return_max_logit, + bottom_right_diagonal=( + attn_mask_type.is_bottom_right() + if bottom_right_diagonal is None + else bottom_right_diagonal + ), ) return output @@ -1238,7 +1247,7 @@ def fused_attn_thd( @partial( jax.custom_vjp, - nondiff_argnums=(5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19), + nondiff_argnums=(5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20), ) def _fused_attn( qkv: Tuple[jnp.ndarray, ...], @@ -1261,6 +1270,7 @@ def _fused_attn( context_checkpoint_name: str = "context", stripe_size: int | None = None, return_max_logit: bool = False, + bottom_right_diagonal: bool = False, ): output, _ = _fused_attn_fwd_rule( qkv, @@ -1283,6 +1293,7 @@ def _fused_attn( context_checkpoint_name=context_checkpoint_name, stripe_size=stripe_size, return_max_logit=return_max_logit, + bottom_right_diagonal=bottom_right_diagonal, ) return output @@ -1308,6 +1319,7 @@ def _fused_attn_fwd_rule( context_checkpoint_name, stripe_size, return_max_logit, + bottom_right_diagonal, ): output, softmax_aux, rng_state, max_logit = tex.fused_attn_fwd( qkv, @@ -1329,6 +1341,7 @@ def _fused_attn_fwd_rule( context_parallel_axis=context_parallel_axis, stripe_size=stripe_size, return_max_logit=return_max_logit, + bottom_right_diagonal=bottom_right_diagonal, ) output = checkpoint_name(output, context_checkpoint_name) softmax_aux = checkpoint_name(softmax_aux, context_checkpoint_name) @@ -1362,6 +1375,7 @@ def _fused_attn_bwd_rule( context_checkpoint_name, stripe_size, return_max_logit, + bottom_right_diagonal, ctx, dz, ): @@ -1399,8 +1413,9 @@ def _fused_attn_bwd_rule( context_parallel_causal_load_balanced=context_parallel_causal_load_balanced, context_parallel_axis=context_parallel_axis, stripe_size=stripe_size, + bottom_right_diagonal=bottom_right_diagonal, ) - if attn_bias_type == AttnBiasType.NO_BIAS: + if attn_bias_type in (AttnBiasType.NO_BIAS, AttnBiasType.ALIBI): grad_bias = None if softmax_type != AttnSoftmaxType.LEARNABLE_SOFTMAX: grad_softmax_offset = None @@ -1416,6 +1431,28 @@ def _fused_attn_bwd_rule( _fused_attn.defvjp(_fused_attn_fwd_rule, _fused_attn_bwd_rule) +@partial(jax.custom_vjp, nondiff_argnums=(4,)) +def _fused_attn_fp8(qkv, sequence_descriptor, seed, quantizer_set, config): + output, _ = tex.fused_attn_fp8_fwd( + qkv, sequence_descriptor, seed, quantizer_set, config + ) + return output + + +def _fused_attn_fp8_fwd_rule(qkv, sequence_descriptor, seed, quantizer_set, config): + return tex.fused_attn_fp8_fwd( + qkv, sequence_descriptor, seed, quantizer_set, config + ) + + +def _fused_attn_fp8_bwd_rule(config, ctx, doutput): + grad_qkv, quantizer_set = tex.fused_attn_fp8_bwd(ctx, doutput, config) + return grad_qkv, None, None, quantizer_set + + +_fused_attn_fp8.defvjp(_fused_attn_fp8_fwd_rule, _fused_attn_fp8_bwd_rule) + + @partial(jax.custom_vjp, nondiff_argnums=(3, 4)) def _fused_attn_score_mod( qkv: Tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray], @@ -1494,6 +1531,8 @@ def fused_attn( score_mod_tensors: Optional[Mapping[str, Any]] = None, score_mod_bprop_tensors: Optional[Mapping[str, Any]] = None, return_max_logit: bool = False, + bottom_right_diagonal: Optional[bool] = None, + quantizer_set: Optional[AttentionQuantizerSet] = None, ): """ Perform cuDNN fused attention. @@ -1552,6 +1591,11 @@ def fused_attn( Python/NumPy scalars made available to `score_mod_bprop`. return_max_logit (bool): If True, also return per-head maximum attention logits with shape ``[h]``. + bottom_right_diagonal (Optional[bool]): Explicitly select bottom-right diagonal + alignment independently of the mask type. By default it follows bottom-right masks. + quantizer_set (Optional[AttentionQuantizerSet]): Quantizers for FP8 DPA. When + provided, Q/K/V and dO are quantized internally while public inputs and outputs + remain FP16/BF16. Returns: jnp.ndarray: Attention output when ``return_max_logit`` is False. @@ -1600,6 +1644,8 @@ def fused_attn( if score_mod_only_args: raise ValueError(f"{', '.join(score_mod_only_args)} require score_mod to be provided.") else: + if quantizer_set is not None: + raise NotImplementedError("JAX FP8 attention does not support score_mod.") if return_max_logit: raise ValueError("return_max_logit is not supported with score_mod fused_attn.") tex.validate_fused_attn_score_mod( @@ -1637,6 +1683,10 @@ def fused_attn( ) if sequence_descriptor is None or isinstance(sequence_descriptor, jnp.ndarray): + if quantizer_set is not None: + raise NotImplementedError( + "JAX FP8 attention requires a SequenceDescriptor instead of a legacy mask." + ) warnings.warn( "Pass mask to fused_attn is deprecated, please use SequenceDescriptor instead. " + "See help(transformer_engine.jax.attention.SequenceDescriptor) for details.", @@ -1662,6 +1712,7 @@ def fused_attn( context_parallel_axis=context_parallel_axis, softmax_offset=softmax_offset, return_max_logit=return_max_logit, + bottom_right_diagonal=bottom_right_diagonal, ) if max_segments_per_seq > 1 and not qkv_layout.is_thd(): warnings.warn( @@ -1672,6 +1723,36 @@ def fused_attn( UserWarning, stacklevel=2, ) + if quantizer_set is not None: + if not isinstance(quantizer_set, AttentionQuantizerSet): + raise TypeError("quantizer_set must be an AttentionQuantizerSet.") + if qkv_layout.is_thd() or max_segments_per_seq != 1: + raise NotImplementedError("JAX FP8 attention does not support THD sequence packing.") + if context_parallel_axis or context_parallel_strategy != CPStrategy.DEFAULT: + raise NotImplementedError("JAX FP8 attention does not support context parallelism.") + if bias is not None or attn_bias_type != AttnBiasType.NO_BIAS: + raise NotImplementedError("JAX FP8 attention does not support attention bias.") + if return_max_logit: + raise NotImplementedError("JAX FP8 attention does not support return_max_logit.") + if softmax_offset is not None: + raise NotImplementedError("JAX FP8 attention does not support a softmax offset.") + diagonal = ( + attn_mask_type.is_bottom_right() + if bottom_right_diagonal is None + else bottom_right_diagonal + ) + config = tex.FP8AttentionConfig( + attn_bias_type=attn_bias_type, + attn_mask_type=attn_mask_type, + softmax_type=softmax_type, + qkv_layout=qkv_layout, + scaling_factor=scaling_factor, + dropout_probability=dropout_probability, + is_training=is_training, + window_size=(-1, -1) if window_size is None else window_size, + bottom_right_diagonal=diagonal, + ) + return _fused_attn_fp8(qkv, sequence_descriptor, seed, quantizer_set, config) output = _fused_attn( qkv, bias, @@ -1693,5 +1774,10 @@ def fused_attn( context_checkpoint_name=context_checkpoint_name, stripe_size=stripe_size, return_max_logit=return_max_logit, + bottom_right_diagonal=( + attn_mask_type.is_bottom_right() + if bottom_right_diagonal is None + else bottom_right_diagonal + ), ) return output diff --git a/transformer_engine/jax/cpp_extensions/__init__.py b/transformer_engine/jax/cpp_extensions/__init__.py index c9647afb826..68eaa7d6776 100644 --- a/transformer_engine/jax/cpp_extensions/__init__.py +++ b/transformer_engine/jax/cpp_extensions/__init__.py @@ -5,6 +5,7 @@ from .activation import * from .amax import * from .attention import * +from .fp8_attention import * from .flex_attention import * from .normalization import * from .quantization import * diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index e2cb5eef463..3b697f569b1 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -3538,6 +3538,7 @@ def fused_attn_fwd( context_parallel_axis: str = "", stripe_size: int | None = None, return_max_logit: bool = False, + bottom_right_diagonal: bool | None = None, ) -> jnp.ndarray: """ Perform the forward pass of with cuDNN fused attention implementations. @@ -3578,6 +3579,8 @@ def fused_attn_fwd( context_parallel_axis (str): The name of the context parallel axis. stripe_size (int | None): Indicates the striping height to be used for ReorderStrategy.Striped Load Balancing return_max_logit (bool): Whether to return the per-head maximum attention logit. + bottom_right_diagonal (bool | None): Explicit diagonal alignment. When unset, it + follows whether ``attn_mask_type`` is a bottom-right mask. Returns: (jnp.ndarray): The output tensor from the fused attention. """ @@ -3601,10 +3604,10 @@ def fused_attn_fwd( else: raise ValueError(f"Unknown {qkv_layout=}") - if attn_bias_type == AttnBiasType.NO_BIAS: + if attn_bias_type in (AttnBiasType.NO_BIAS, AttnBiasType.ALIBI): assert ( bias is None - ), f"bias must be None when attn_bias_type is NO_BIAS, but got bias={bias}" + ), f"bias must be None when attn_bias_type is {attn_bias_type}, but got bias={bias}" bias = jnp.zeros(0, dtype=qkv[0].dtype) if softmax_offset is None: @@ -3648,7 +3651,11 @@ def fused_attn_fwd( is_training=is_training, max_segments_per_seq=max_segments_per_seq, window_size=(-1, -1) if window_size is None else window_size, - bottom_right_diagonal=attn_mask_type.is_bottom_right(), + bottom_right_diagonal=( + attn_mask_type.is_bottom_right() + if bottom_right_diagonal is None + else bottom_right_diagonal + ), context_parallel_load_balanced=context_parallel_causal_load_balanced, cp_axis=_maybe_context_parallel_axis(context_parallel_axis), cp_striped_window_size=None, @@ -3705,6 +3712,7 @@ def fused_attn_bwd( context_parallel_causal_load_balanced: bool = False, context_parallel_axis: str = "", stripe_size: int | None = None, + bottom_right_diagonal: bool | None = None, ): """ Perform the backward pass of the cuDNN fused attention implementations. @@ -3745,6 +3753,8 @@ def fused_attn_bwd( Indicates the sequences are ordered for causal mask load balancing when running context parallelism. context_parallel_axis (str): The name of the context parallel axis. stripe_size (int | None): Indicates the striping height to be used for ReorderStrategy.Striped Load Balancing + bottom_right_diagonal (bool | None): Explicit diagonal alignment. When unset, it + follows whether ``attn_mask_type`` is a bottom-right mask. Returns: Tuple[jnp.ndarray, ...], jnp.ndarray: - The first tuple contains the gradients with respect to the input `qkv` tensors in the @@ -3770,10 +3780,10 @@ def fused_attn_bwd( else: raise ValueError(f"Unknown {qkv_layout=}") - if attn_bias_type == AttnBiasType.NO_BIAS: + if attn_bias_type in (AttnBiasType.NO_BIAS, AttnBiasType.ALIBI): assert ( bias is None - ), f"bias must be None when attn_bias_type is NO_BIAS, but got bias with type={type(bias)}" + ), f"bias must be None when attn_bias_type is {attn_bias_type}, but got bias with type={type(bias)}" bias = jnp.zeros(0, dtype=qkv[0].dtype) if softmax_offset is None: @@ -3824,7 +3834,11 @@ def fused_attn_bwd( is_training=is_training, max_segments_per_seq=max_segments_per_seq, window_size=(-1, -1) if window_size is None else window_size, - bottom_right_diagonal=attn_mask_type.is_bottom_right(), + bottom_right_diagonal=( + attn_mask_type.is_bottom_right() + if bottom_right_diagonal is None + else bottom_right_diagonal + ), context_parallel_load_balanced=context_parallel_causal_load_balanced, cp_axis=_maybe_context_parallel_axis(context_parallel_axis), cp_striped_window_size=None, diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index cec1e679e6e..db974d1def6 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -18,7 +18,7 @@ AttentionLayout, FusedAttentionConfig, check_f16_fused_attention_support, - normalize_attention_mask, + cudnn_mask_options, ragged_batch_bucket, ragged_token_bucket, ) @@ -311,12 +311,13 @@ def _ragged_offset_spec(cudnn): def _mask_options(cudnn, info: _LayoutInfo, config): + cp_striped_window_size = getattr(config, "cp_striped_window_size", None) window_left, window_right = ( - config.cp_striped_window_size - if config.cp_striped_window_size is not None + cp_striped_window_size + if cp_striped_window_size is not None else config.window_size ) - mask = normalize_attention_mask( + options = cudnn_mask_options( causal=_is_causal(config), bottom_right=_is_bottom_right(config), padding=_is_padding(config), @@ -324,27 +325,14 @@ def _mask_options(cudnn, info: _LayoutInfo, config): window_size=(window_left, window_right), max_seqlen_q=info.q_max_seqlen, max_seqlen_kv=info.kv_max_seqlen, + cudnn_version=get_cudnn_version(), + ) + options.pop("is_padding") + options["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if options["diagonal_alignment"] == "bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT ) - cudnn_version = get_cudnn_version() - options = { - "diagonal_alignment": ( - cudnn.diagonal_alignment.BOTTOM_RIGHT - if mask.bottom_right_diagonal or mask.bottom_right - else cudnn.diagonal_alignment.TOP_LEFT - ), - } - # Before cuDNN 9.6 the preferred right-band API was unavailable, so preserve - # the legacy causal flags used by the C++ frontend graph. - if cudnn_version < (9, 6, 0): - options["use_causal_mask"] = mask.causal - options["use_causal_mask_bottom_right"] = mask.bottom_right - if cudnn_version >= (9, 2, 0) and mask.window_left != -1: - options["diagonal_band_left_bound"] = mask.window_left + 1 - if cudnn_version >= (9, 6, 0): - if mask.window_right != -1: - options["diagonal_band_right_bound"] = mask.window_right - elif mask.causal or mask.bottom_right: - options["diagonal_band_right_bound"] = 0 return options @@ -434,6 +422,8 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap "attn_scale": scale, **_mask_options(cudnn, info, config), } + if getattr(config.attn_bias_type, "name", "") == "ALIBI": + kwargs["use_alibi_mask"] = True if _is_bias(config): *bias_batch_shape, bias_heads, bias_sq, bias_skv = bias_aval.shape @@ -743,6 +733,8 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): "attn_scale": scale, **_mask_options(cudnn, info, config), } + if getattr(config.attn_bias_type, "name", "") == "ALIBI": + kwargs["use_alibi_mask"] = True if get_cudnn_version() >= (9, 0, 0): kwargs["use_deterministic_algorithm"] = not bool( int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) @@ -974,9 +966,6 @@ def is_fused_attn_supported(helper) -> bool: ), cudnn_version=get_cudnn_version(), sm_arch=_device_arch(), - allow_alibi=False, - allow_extended_causal_window=True, - modern_mask_rules_override=True, ) ) return support.supported diff --git a/transformer_engine/jax/cpp_extensions/cudnn_graph.py b/transformer_engine/jax/cpp_extensions/cudnn_graph.py index a879e754382..8ba743e9079 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_graph.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_graph.py @@ -110,6 +110,12 @@ def cudnn_data_type(cudnn, dtype): return cudnn.data_type.HALF if dtype == jnp.bfloat16: return cudnn.data_type.BFLOAT16 + if dtype == jnp.float8_e4m3fn: + return cudnn.data_type.FP8_E4M3 + if dtype == jnp.float8_e5m2: + return cudnn.data_type.FP8_E5M2 + if hasattr(jnp, "float8_e8m0fnu") and dtype == jnp.float8_e8m0fnu: + return cudnn.data_type.FP8_E8M0 if dtype == jnp.float32: return cudnn.data_type.FLOAT if dtype == jnp.float64: diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py new file mode 100644 index 00000000000..bfc76c7a8fb --- /dev/null +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -0,0 +1,1238 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""JAX execution adapter for common cuDNN FP8 attention graph construction.""" + +from __future__ import annotations + +import copy +import operator +import os +from dataclasses import dataclass +from functools import reduce +from typing import Any + +import jax +import jax.numpy as jnp +from jax import ffi + +from transformer_engine.common.attention.cudnn import ( + FusedAttentionConfig, + check_fp8_fused_attention_support, +) +from transformer_engine.common.attention.fp8 import ( + FP8AttentionGraphConfig, + attention_format_stride, + build_fp8_backward_operation, + build_fp8_forward_operation, + mxfp8_padded_sizes, +) + +from ..quantize import ScalingMode, TensorUsage, swizzle_mxfp8_scale +from .attention import _FusedAttnRNGStateChecker +from .cudnn_attention import ( + _UID_DK, + _UID_DO, + _UID_DQ, + _UID_DROPOUT_OFFSET, + _UID_DROPOUT_SEED, + _UID_DV, + _UID_K, + _UID_O, + _UID_Q, + _UID_SEQ_KV, + _UID_SEQ_Q, + _UID_STATS, + _UID_V, + _device_arch, + _is_dropout, + _is_padding, + _layout_info, + _mask_options, + _matrix_stride, + _policy_layout, + _policy_mask_name, + _qkv_bindings, +) +from .cudnn_graph import ( + GraphBinding, + SerializedGraph, + cudnn_data_type, + dtype_name, + finalize_graph, + import_cudnn, + make_graph, + serialized_graph, +) +from .misc import get_cudnn_version +from .quantization import quantize + +__all__ = ["FP8AttentionConfig", "fused_attn_fp8_bwd", "fused_attn_fp8_fwd"] + + +@dataclass(frozen=True) +class FP8AttentionConfig: + """Static configuration for JAX FP8 dot-product attention.""" + + attn_bias_type: Any + attn_mask_type: Any + softmax_type: Any + qkv_layout: Any + scaling_factor: float + dropout_probability: float + is_training: bool + window_size: tuple[int, int] + bottom_right_diagonal: bool = False + + +_UID_DESCALE_Q = 101 +_UID_DESCALE_K = 102 +_UID_DESCALE_V = 103 +_UID_DESCALE_S = 104 +_UID_SCALE_S = 105 +_UID_SCALE_O = 106 +_UID_AMAX_S = 107 +_UID_AMAX_O = 108 +_UID_DESCALE_O = 109 +_UID_DESCALE_DO = 110 +_UID_DESCALE_DP = 111 +_UID_SCALE_DQ = 112 +_UID_SCALE_DK = 113 +_UID_SCALE_DV = 114 +_UID_SCALE_DP = 115 +_UID_AMAX_DQ = 116 +_UID_AMAX_DK = 117 +_UID_AMAX_DV = 118 +_UID_AMAX_DP = 119 +_UID_Q_T = 120 +_UID_K_T = 121 +_UID_DO_T = 122 +_UID_DESCALE_Q_T = 123 +_UID_DESCALE_K_T = 124 +_UID_DESCALE_DO_T = 125 +_UID_DO_F16 = 126 + + +@dataclass(frozen=True) +class FP8AttentionGraphInfo: + """Serialized graph and result metadata for an FP8 attention call.""" + + graph: SerializedGraph + output_shape: tuple[int, ...] + q_shape: tuple[int, ...] + k_shape: tuple[int, ...] + v_shape: tuple[int, ...] + + +_graph_cache: dict[tuple[Any, ...], FP8AttentionGraphInfo] = {} + + +def _cache_key(direction, mode, avals, config, output_dtype): + return ( + direction, + mode, + config, + dtype_name(output_dtype), + tuple((tuple(aval.shape), dtype_name(aval.dtype)) for aval in avals), + get_cudnn_version(), + _device_arch(), + ) + + +def _tensor(graph, *, name, dim, stride, dtype, uid): + return graph.tensor( + name=name, + dim=tuple(int(value) for value in dim), + stride=tuple(int(value) for value in stride), + data_type=dtype, + uid=uid, + ) + + +def _scalar(graph, cudnn, name, uid): + return _tensor( + graph, + name=name, + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.FLOAT, + uid=uid, + ) + + +def _mx_scale( + graph, + cudnn, + *, + name, + uid, + batch, + heads, + seqlen, + dim, + tensor_format="bshd", +): + return _tensor( + graph, + name=name, + dim=(batch, heads, seqlen, dim), + stride=attention_format_stride(batch, heads, seqlen, dim, tensor_format), + dtype=cudnn.data_type.FP8_E8M0, + uid=uid, + ).set_reordering_type(cudnn.tensor_reordering.F8_128x4) + + +def _logical_shapes(info): + q = (*info.batch_shape, info.q_max_seqlen, info.q_heads, info.qk_dim) + k = (*info.batch_shape, info.kv_max_seqlen, info.kv_heads, info.qk_dim) + v = (*info.batch_shape, info.kv_max_seqlen, info.kv_heads, info.v_dim) + o = (*info.batch_shape, info.q_max_seqlen, info.q_heads, info.v_dim) + return q, k, v, o + + +def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): + info = _layout_info(q_aval, k_aval, v_aval, config.qkv_layout) + io_dtype = cudnn_data_type(cudnn, q_aval.dtype) + q = _tensor( + graph, + name="Q", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.qk_dim), + stride=_matrix_stride( + info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen + ), + dtype=io_dtype, + uid=_UID_Q, + ) + k = _tensor( + graph, + name="K", + dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.qk_dim), + stride=_matrix_stride( + info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen + ), + dtype=io_dtype, + uid=_UID_K, + ) + v = _tensor( + graph, + name="V", + dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.v_dim), + stride=_matrix_stride( + info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen + ), + dtype=io_dtype, + uid=_UID_V, + ) + return info, io_dtype, q, k, v + + +def _fp8_options(cudnn, info, config): + options = _mask_options(cudnn, info, config) + is_padding = options.pop("is_padding", _is_padding(config)) + if "diagonal_band_left_bound" in options: + options["left_bound"] = options.pop("diagonal_band_left_bound") + if "diagonal_band_right_bound" in options: + options["right_bound"] = options.pop("diagonal_band_right_bound") + options.update( + attn_scale=float(config.scaling_factor), + use_padding_mask=is_padding, + ) + return options, is_padding + + +def build_fp8_fwd_graph( + q_aval, + k_aval, + v_aval, + q_scale_aval, + k_scale_aval, + v_scale_aval, + config, + mode: str, + output_dtype, +) -> FP8AttentionGraphInfo: + """Build or retrieve a dense JAX FP8 attention forward graph.""" + + avals = (q_aval, k_aval, v_aval, q_scale_aval, k_scale_aval, v_scale_aval) + key = _cache_key("fwd", mode, avals, config, output_dtype) + if key in _graph_cache: + return _graph_cache[key] + + cudnn = import_cudnn() + graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) + info, _, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) + q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) + input_bindings = list( + _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) + ) + tensors = {"q": q, "k": k, "v": v} + options, is_padding = _fp8_options(cudnn, info, config) + options["generate_stats"] = True + + if mode == "mxfp8": + if is_padding: + raise ValueError("JAX MXFP8 attention does not support padding masks.") + options.pop("use_padding_mask", None) + # The current cuDNN MXFP8 forward binding uses diagonal_band_* names. + if "left_bound" in options: + options["diagonal_band_left_bound"] = options.pop("left_bound") + if "right_bound" in options: + options["diagonal_band_right_bound"] = options.pop("right_bound") + padded = mxfp8_padded_sizes( + info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim + ) + scale_specs = ( + ( + "descale_q", + _UID_DESCALE_Q, + info.q_heads, + "s_q_padded", + "d_qk_scale_padded", + 3, + ), + ( + "descale_k", + _UID_DESCALE_K, + info.kv_heads, + "s_kv_padded", + "d_qk_scale_padded", + 4, + ), + ( + "descale_v", + _UID_DESCALE_V, + info.kv_heads, + "s_kv_scale_padded", + "d_v_padded", + 6, + ), + ) + for name, uid, heads, s_key, d_key, buffer_index in scale_specs: + tensors[name] = _mx_scale( + graph, + cudnn, + name=name, + uid=uid, + batch=info.input_batch, + heads=heads, + seqlen=padded[s_key], + dim=padded[d_key], + ) + input_bindings.append(GraphBinding(uid, buffer_index)) + else: + for name, uid, index in ( + ("descale_q", _UID_DESCALE_Q, 3), + ("descale_k", _UID_DESCALE_K, 4), + ("descale_v", _UID_DESCALE_V, 6), + ("descale_s", _UID_DESCALE_S, 7), + ("scale_s", _UID_SCALE_S, 8), + ("scale_o", _UID_SCALE_O, 9), + ): + tensors[name] = _scalar(graph, cudnn, name, uid) + input_bindings.append(GraphBinding(uid, index)) + + if is_padding: + seq_q = _tensor( + graph, + name="seq_len_q", + dim=(info.input_batch, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=_UID_SEQ_Q, + ) + seq_kv = _tensor( + graph, + name="seq_len_kv", + dim=(info.input_batch, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=_UID_SEQ_KV, + ) + options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) + input_bindings.extend( + (GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11)) + ) + + output_bindings = [GraphBinding(_UID_O, 0), GraphBinding(_UID_STATS, 1)] + if _is_dropout(config): + seed = _tensor( + graph, + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT64, + uid=_UID_DROPOUT_SEED, + ) + offset = _tensor( + graph, + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT64, + uid=_UID_DROPOUT_OFFSET, + ) + options["dropout"] = (float(config.dropout_probability), seed, offset) + output_bindings.extend( + ( + GraphBinding(_UID_DROPOUT_SEED, 3), + GraphBinding(_UID_DROPOUT_OFFSET, 3, 8), + ) + ) + + op = build_fp8_forward_operation( + graph, + tensors, + options, + FP8AttentionGraphConfig(mode, "te_jax_fp8_sdpa_forward"), + ) + output = op["output"] + output.set_output(True).set_uid(_UID_O).set_data_type( + cudnn_data_type(cudnn, output_dtype) + ).set_dim( + (info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim) + ).set_stride( + attention_format_stride( + info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" + ) + ) + stats = op["stats"] + stats.set_output(True).set_uid(_UID_STATS).set_data_type( + cudnn.data_type.FLOAT + ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, 1)).set_stride( + (info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1) + ) + if mode != "mxfp8": + for name, uid, offset in ( + ("amax_s", _UID_AMAX_S, 0), + ("amax_o", _UID_AMAX_O, 4), + ): + op[name].set_output(True).set_uid(uid).set_data_type( + cudnn.data_type.FLOAT + ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) + output_bindings.append(GraphBinding(uid, 2, offset)) + else: + op["amax_o"].set_output(False) + + workspace, data, version = finalize_graph( + cudnn, graph, description=f"JAX {mode} FP8 attention forward" + ) + result = serialized_graph( + serialized_graph_data=data, + cudnn_frontend_version=version, + workspace_size=workspace, + input_bindings=input_bindings, + output_bindings=output_bindings, + ) + info_result = FP8AttentionGraphInfo(result, o_shape, q_shape, k_shape, v_shape) + _graph_cache[key] = info_result + return info_result + + +def _mx_bwd_scales(graph, cudnn, info, padded, input_bindings): + specs = ( + ( + "descale_q", + _UID_DESCALE_Q, + info.q_heads, + "s_q_padded", + "d_qk_scale_padded", + 11, + ), + ( + "descale_q_t", + _UID_DESCALE_Q_T, + info.q_heads, + "s_q_scale_padded", + "d_qk_padded", + 25, + ), + ( + "descale_k", + _UID_DESCALE_K, + info.kv_heads, + "s_kv_padded", + "d_qk_scale_padded", + 12, + ), + ( + "descale_k_t", + _UID_DESCALE_K_T, + info.kv_heads, + "s_kv_scale_padded", + "d_qk_padded", + 26, + ), + ( + "descale_v", + _UID_DESCALE_V, + info.kv_heads, + "s_kv_padded", + "d_v_scale_padded", + 13, + ), + ( + "descale_do", + _UID_DESCALE_DO, + info.q_heads, + "s_q_padded", + "d_v_scale_padded", + 15, + ), + ( + "descale_do_t", + _UID_DESCALE_DO_T, + info.q_heads, + "s_q_scale_padded", + "d_v_padded", + 24, + ), + ) + tensors = {} + for name, uid, heads, s_key, d_key, index in specs: + tensors[name] = _mx_scale( + graph, + cudnn, + name=name, + uid=uid, + batch=info.input_batch, + heads=heads, + seqlen=padded[s_key], + dim=padded[d_key], + ) + input_bindings.append(GraphBinding(uid, index)) + return tensors + + +def build_fp8_bwd_graph( + q_aval, + k_aval, + v_aval, + stats_aval, + output_aval, + doutput_aval, + config, + mode: str, + grad_dtype, +) -> FP8AttentionGraphInfo: + """Build or retrieve a dense JAX FP8 attention backward graph.""" + + avals = (q_aval, k_aval, v_aval, stats_aval, output_aval, doutput_aval) + key = _cache_key("bwd", mode, avals, config, grad_dtype) + if key in _graph_cache: + return _graph_cache[key] + + cudnn = import_cudnn() + graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) + info, io_dtype, q, k, v = _graph_io_tensors( + graph, cudnn, q_aval, k_aval, v_aval, config + ) + q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) + input_bindings = list( + _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) + ) + o = _tensor( + graph, + name="O", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim), + stride=attention_format_stride( + info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" + ), + dtype=cudnn_data_type(cudnn, output_aval.dtype), + uid=_UID_O, + ) + do = _tensor( + graph, + name="dO", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim), + stride=attention_format_stride( + info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" + ), + dtype=cudnn_data_type(cudnn, doutput_aval.dtype), + uid=_UID_DO, + ) + stats = _tensor( + graph, + name="Stats", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, 1), + stride=(info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1), + dtype=cudnn.data_type.FLOAT, + uid=_UID_STATS, + ) + input_bindings.extend( + (GraphBinding(_UID_STATS, 5), GraphBinding(_UID_O, 7), GraphBinding(_UID_DO, 8)) + ) + tensors = {"q": q, "k": k, "v": v, "o": o, "do": do, "stats": stats} + options, is_padding = _fp8_options(cudnn, info, config) + options["use_deterministic_algorithm"] = not bool( + int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) + ) + if is_padding: + seq_q = _tensor( + graph, + name="seq_len_q", + dim=(info.input_batch, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=_UID_SEQ_Q, + ) + seq_kv = _tensor( + graph, + name="seq_len_kv", + dim=(info.input_batch, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT32, + uid=_UID_SEQ_KV, + ) + options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) + input_bindings.extend( + (GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10)) + ) + + if _is_dropout(config): + seed = _tensor( + graph, + name="dropout_seed", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT64, + uid=_UID_DROPOUT_SEED, + ) + offset = _tensor( + graph, + name="dropout_offset", + dim=(1, 1, 1, 1), + stride=(1, 1, 1, 1), + dtype=cudnn.data_type.INT64, + uid=_UID_DROPOUT_OFFSET, + ) + options["dropout"] = (float(config.dropout_probability), seed, offset) + input_bindings.extend( + ( + GraphBinding(_UID_DROPOUT_SEED, 6), + GraphBinding(_UID_DROPOUT_OFFSET, 6, 8), + ) + ) + + if mode == "mxfp8": + if is_padding: + raise ValueError("JAX MXFP8 attention does not support padding masks.") + q_t = _tensor( + graph, + name="Q_T", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.qk_dim), + stride=attention_format_stride( + info.input_batch, info.q_heads, info.q_max_seqlen, info.qk_dim, "bshd" + ), + dtype=io_dtype, + uid=_UID_Q_T, + ) + k_t = _tensor( + graph, + name="K_T", + dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.qk_dim), + stride=attention_format_stride( + info.input_batch, info.kv_heads, info.kv_max_seqlen, info.qk_dim, "bshd" + ), + dtype=io_dtype, + uid=_UID_K_T, + ) + do_t = _tensor( + graph, + name="dO_T", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim), + stride=attention_format_stride( + info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" + ), + dtype=cudnn_data_type(cudnn, doutput_aval.dtype), + uid=_UID_DO_T, + ) + do_f16 = _tensor( + graph, + name="dO_f16", + dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim), + stride=attention_format_stride( + info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" + ), + dtype=cudnn_data_type(cudnn, output_aval.dtype), + uid=_UID_DO_F16, + ) + tensors.update(q_t=q_t, k_t=k_t, do_t=do_t, do_f16=do_f16) + input_bindings.extend( + ( + GraphBinding(_UID_Q_T, 3), + GraphBinding(_UID_K_T, 4), + GraphBinding(_UID_DO_T, 23), + GraphBinding(_UID_DO_F16, 27), + ) + ) + padded = mxfp8_padded_sizes( + info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim + ) + tensors.update(_mx_bwd_scales(graph, cudnn, info, padded, input_bindings)) + else: + for name, uid, index in ( + ("descale_q", _UID_DESCALE_Q, 11), + ("descale_k", _UID_DESCALE_K, 12), + ("descale_v", _UID_DESCALE_V, 13), + ("descale_o", _UID_DESCALE_O, 14), + ("descale_do", _UID_DESCALE_DO, 15), + ("descale_s", _UID_DESCALE_S, 16), + ("descale_dp", _UID_DESCALE_DP, 17), + ("scale_s", _UID_SCALE_S, 18), + ("scale_dq", _UID_SCALE_DQ, 19), + ("scale_dk", _UID_SCALE_DK, 20), + ("scale_dv", _UID_SCALE_DV, 21), + ("scale_dp", _UID_SCALE_DP, 22), + ): + tensors[name] = _scalar(graph, cudnn, name, uid) + input_bindings.append(GraphBinding(uid, index)) + + op = build_fp8_backward_operation( + graph, + tensors, + options, + FP8AttentionGraphConfig(mode, "te_jax_fp8_sdpa_backward"), + ) + output_bindings = [] + for name, uid, index, shape in ( + ("dq", _UID_DQ, 0, q_shape), + ("dk", _UID_DK, 1, k_shape), + ("dv", _UID_DV, 2, v_shape), + ): + batch = reduce(operator.mul, shape[:-3], 1) + seqlen, heads, dim = shape[-3:] + op[name].set_output(True).set_uid(uid).set_data_type( + cudnn_data_type(cudnn, grad_dtype) + ).set_dim((batch, heads, seqlen, dim)).set_stride( + attention_format_stride(batch, heads, seqlen, dim, "bshd") + ) + output_bindings.append(GraphBinding(uid, index)) + + if mode != "mxfp8": + for name, uid, offset in ( + ("amax_dq", _UID_AMAX_DQ, 0), + ("amax_dk", _UID_AMAX_DK, 4), + ("amax_dv", _UID_AMAX_DV, 8), + ("amax_dp", _UID_AMAX_DP, 12), + ): + op[name].set_output(True).set_uid(uid).set_data_type( + cudnn.data_type.FLOAT + ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) + output_bindings.append(GraphBinding(uid, 3, offset)) + else: + for amax in op["amax"]: + amax.set_output(False) + + workspace, data, version = finalize_graph( + cudnn, graph, description=f"JAX {mode} FP8 attention backward" + ) + result = serialized_graph( + serialized_graph_data=data, + cudnn_frontend_version=version, + workspace_size=workspace, + input_bindings=input_bindings, + output_bindings=output_bindings, + ) + graph_info = FP8AttentionGraphInfo(result, o_shape, q_shape, k_shape, v_shape) + _graph_cache[key] = graph_info + return graph_info + + +def execute_fp8_fwd( + q, + k, + v, + q_scale_inv, + k_scale_inv, + v_scale_inv, + s_scale_inv, + s_scale, + o_scale, + q_seqlen, + kv_seqlen, + seed, + *, + config, + mode, + output_dtype, +): + """Execute a serialized dense FP8 attention forward graph.""" + + graph_info = build_fp8_fwd_graph( + jax.ShapeDtypeStruct(q.shape, q.dtype), + jax.ShapeDtypeStruct(k.shape, k.dtype), + jax.ShapeDtypeStruct(v.shape, v.dtype), + jax.ShapeDtypeStruct(q_scale_inv.shape, q_scale_inv.dtype), + jax.ShapeDtypeStruct(k_scale_inv.shape, k_scale_inv.dtype), + jax.ShapeDtypeStruct(v_scale_inv.shape, v_scale_inv.dtype), + config, + mode, + output_dtype, + ) + result_specs = ( + jax.ShapeDtypeStruct(graph_info.output_shape, output_dtype), + jax.ShapeDtypeStruct( + ( + *graph_info.output_shape[:-3], + graph_info.output_shape[-2], + graph_info.output_shape[-3], + 1, + ), + jnp.float32, + ), + jax.ShapeDtypeStruct((2,), jnp.float32), + jax.ShapeDtypeStruct((seed.shape[0], 4), jnp.uint32), + jax.ShapeDtypeStruct((graph_info.graph.workspace_size,), jnp.uint8), + ) + return ffi.ffi_call("te_fused_attn_forward_ffi", result_specs)( + q, + k, + v, + q_scale_inv, + k_scale_inv, + seed, + v_scale_inv, + s_scale_inv, + s_scale, + o_scale, + q_seqlen, + kv_seqlen, + is_ragged=False, + rng_offset_increment=16, + **graph_info.graph.ffi_attrs(), + ) + + +def execute_fp8_bwd( + q, + k, + v, + q_t, + k_t, + stats, + rng_state, + output, + doutput, + q_seqlen, + kv_seqlen, + q_scale_inv, + k_scale_inv, + v_scale_inv, + o_scale_inv, + do_scale_inv, + s_scale_inv, + dp_scale_inv, + s_scale, + dq_scale, + dk_scale, + dv_scale, + dp_scale, + do_t, + do_scale_inv_t, + q_scale_inv_t, + k_scale_inv_t, + doutput_f16, + *, + config, + mode, + grad_dtype, +): + """Execute a serialized dense FP8 attention backward graph.""" + + graph_info = build_fp8_bwd_graph( + jax.ShapeDtypeStruct(q.shape, q.dtype), + jax.ShapeDtypeStruct(k.shape, k.dtype), + jax.ShapeDtypeStruct(v.shape, v.dtype), + jax.ShapeDtypeStruct(stats.shape, stats.dtype), + jax.ShapeDtypeStruct(output.shape, output.dtype), + jax.ShapeDtypeStruct(doutput.shape, doutput.dtype), + config, + mode, + grad_dtype, + ) + result_specs = ( + jax.ShapeDtypeStruct(graph_info.q_shape, grad_dtype), + jax.ShapeDtypeStruct(graph_info.k_shape, grad_dtype), + jax.ShapeDtypeStruct(graph_info.v_shape, grad_dtype), + jax.ShapeDtypeStruct((4,), jnp.float32), + jax.ShapeDtypeStruct((0,), output.dtype), + jax.ShapeDtypeStruct((graph_info.graph.workspace_size,), jnp.uint8), + ) + return ffi.ffi_call("te_fused_attn_backward_ffi", result_specs)( + q, + k, + v, + q_t, + k_t, + stats, + rng_state, + output, + doutput, + q_seqlen, + kv_seqlen, + q_scale_inv, + k_scale_inv, + v_scale_inv, + o_scale_inv, + do_scale_inv, + s_scale_inv, + dp_scale_inv, + s_scale, + dq_scale, + dk_scale, + dv_scale, + dp_scale, + do_t, + do_scale_inv_t, + q_scale_inv_t, + k_scale_inv_t, + doutput_f16, + is_ragged=False, + **graph_info.graph.ffi_attrs(), + ) + + +def _scaling_mode(quantizer) -> str: + mode = quantizer.scaling_mode + if mode == ScalingMode.DELAYED_TENSOR_SCALING: + return "delayed" + if mode == ScalingMode.CURRENT_TENSOR_SCALING: + return "current" + if mode == ScalingMode.MXFP8_1D_SCALING: + return "mxfp8" + raise ValueError(f"FP8 attention does not support scaling mode {mode}.") + + +def _rowwise(tensor): + return tensor.get_tensor(TensorUsage.LHS) + + +def _colwise(tensor): + return tensor.get_tensor(TensorUsage.RHS) + + +def _quantize_many(values, quantizer, *, both=False): + tensors = [] + amaxes = [] + for value in values: + local_quantizer = copy.copy(quantizer) + tensor = quantize(value, quantizer=local_quantizer, flatten_axis=-2) + tensors.append(tensor) + rowwise_amax = _rowwise(tensor).amax + if rowwise_amax is not None: + amaxes.append(rowwise_amax) + if quantizer.scaling_mode == ScalingMode.DELAYED_TENSOR_SCALING and amaxes: + quantizer.update(jnp.max(jnp.stack(amaxes))) + if both: + # Access both layouts here so an invalid recipe/layout fails before graph construction. + for tensor in tensors: + _rowwise(tensor) + _colwise(tensor) + return tuple(tensors) + + +def _quantized_operands(qkv, layout, quantizer, *, both=False): + quantized = _quantize_many(qkv, quantizer, both=both) + empty = jnp.zeros((0,), dtype=quantizer.q_dtype) + if layout.is_qkvpacked(): + tensor = quantized[0] + return (tensor, tensor, tensor), (_rowwise(tensor).data, empty, empty) + if layout.is_kvpacked(): + q_tensor, kv_tensor = quantized + return ( + q_tensor, + kv_tensor, + kv_tensor, + ), (_rowwise(q_tensor).data, _rowwise(kv_tensor).data, empty) + if layout.is_separate(): + return quantized, tuple(_rowwise(tensor).data for tensor in quantized) + raise ValueError(f"FP8 attention does not support layout {layout}.") + + +def _scale_inv(tensor, *, colwise=False): + return (_colwise(tensor) if colwise else _rowwise(tensor)).scale_inv + + +def _mxfp8_scale_inv(tensor, *, colwise=False): + """Prepare a compact BSHD MXFP8 scale tensor for cuDNN's F8_128x4 layout.""" + + component = _colwise(tensor) if colwise else _rowwise(tensor) + *batch_shape, seqlen, heads, dim = component.data.shape + batch = reduce(operator.mul, batch_shape, 1) + if colwise: + scale = component.scale_inv.reshape(batch, seqlen // 32, heads, dim) + target_seqlen = ((seqlen + 127) // 128) * 4 + target_dim = ((dim + 127) // 128) * 128 + else: + scale = component.scale_inv.reshape(batch, seqlen, heads, dim // 32) + target_seqlen = ((seqlen + 127) // 128) * 128 + target_dim = ((dim + 127) // 128) * 4 + scale = jnp.pad( + scale, + ( + (0, 0), + (0, target_seqlen - scale.shape[1]), + (0, 0), + (0, target_dim - scale.shape[3]), + ), + mode="constant", + constant_values=2**-127, + ) + scale = jnp.transpose(scale, (0, 2, 1, 3)) + return swizzle_mxfp8_scale(scale, -1, colwise) + + +def _tensor_scale(quantizer): + if quantizer.scaling_mode == ScalingMode.DELAYED_TENSOR_SCALING: + return quantizer.scale + return jnp.ones((1,), dtype=jnp.float32) + + +def _tensor_scale_inv(quantizer): + return jnp.reciprocal(_tensor_scale(quantizer)) + + +def _graph_scale_inv(tensor, mode, *, colwise=False): + if mode == "mxfp8": + return _mxfp8_scale_inv(tensor, colwise=colwise) + return _scale_inv(tensor, colwise=colwise) + + +def _sequence_lengths(sequence_descriptor, config): + (q_seqlen, kv_seqlen), _ = sequence_descriptor.get_seqlens_and_offsets( + config.attn_mask_type, + config.qkv_layout, + config.window_size, + 1, + ) + return q_seqlen.flatten(), kv_seqlen.flatten() + + +def _validate_fp8_support(qkv, quantizers, config, mode): + if config.qkv_layout.is_qkvpacked(): + q = k = v = qkv[0] + elif config.qkv_layout.is_kvpacked(): + q, k = qkv + v = k + else: + q, k, v = qkv + info = _layout_info(q, k, v, config.qkv_layout) + q_dtype = jnp.dtype(quantizers.qkv.q_dtype) + dtype_name_ = { + jnp.dtype(jnp.float8_e4m3fn): "float8_e4m3", + jnp.dtype(jnp.float8_e5m2): "float8_e5m2", + }.get(q_dtype, str(q_dtype)) + support = check_fp8_fused_attention_support( + FusedAttentionConfig( + is_training=bool(config.is_training), + q_dtype=dtype_name_, + kv_dtype=dtype_name_, + layout=_policy_layout(config.qkv_layout), + bias_type=config.attn_bias_type.name.lower(), + mask_type=_policy_mask_name(config.attn_mask_type), + softmax_type=config.softmax_type.name.lower().removesuffix("_softmax"), + dropout=float(config.dropout_probability), + num_attn_heads=info.q_heads, + num_gqa_groups=info.kv_heads, + max_seqlen_q=info.q_max_seqlen, + max_seqlen_kv=info.kv_max_seqlen, + head_dim_qk=info.qk_dim, + head_dim_v=info.v_dim, + window_size=tuple(int(value) for value in config.window_size), + return_max_logit=False, + cuda_graph=False, + deterministic=not bool( + int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) + ), + cudnn_version=get_cudnn_version(), + sm_arch=_device_arch(), + ) + ) + if not support.supported: + raise ValueError( + f"Unsupported JAX FP8 attention configuration: {support.reason}." + ) + if mode == "mxfp8" and (get_cudnn_version() < (9, 21, 0) or _device_arch() < 100): + raise ValueError("MXFP8 attention requires cuDNN 9.21 and SM100 or newer.") + + +def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): + """Quantize high-precision inputs and execute dense FP8 attention forward.""" + + mode = _scaling_mode(quantizers.qkv) + if any( + _scaling_mode(q) != mode + for q in ( + quantizers.s, + quantizers.o, + quantizers.do, + quantizers.dp, + quantizers.dqkv, + ) + ): + raise ValueError("All FP8 attention quantizers must use the same scaling mode.") + if config.qkv_layout.is_thd(): + raise NotImplementedError("FP8 attention does not support THD layouts in JAX.") + if mode == "mxfp8" and not config.qkv_layout.is_separate(): + raise NotImplementedError( + "JAX MXFP8 attention currently requires separate BSHD Q/K/V." + ) + if getattr(config.attn_bias_type, "name", "") != "NO_BIAS": + raise NotImplementedError("FP8 attention does not support attention bias.") + if getattr(config.softmax_type, "name", "") != "VANILLA_SOFTMAX": + raise NotImplementedError( + "JAX FP8 attention currently supports vanilla softmax only." + ) + _validate_fp8_support(qkv, quantizers, config, mode) + if mode == "mxfp8" and _is_padding(config): + raise NotImplementedError("JAX MXFP8 attention does not support padding masks.") + + quantized, data = _quantized_operands( + qkv, config.qkv_layout, quantizers.qkv, both=mode == "mxfp8" + ) + q_tensor, k_tensor, v_tensor = quantized + q_data, k_data, v_data = data + # MXFP8 forward consumes V in the columnwise orientation. + if mode == "mxfp8": + v_data = _colwise(v_tensor).data + q_seqlen, kv_seqlen = _sequence_lengths(sequence_descriptor, config) + seed = _FusedAttnRNGStateChecker().check_seed( + seed, config.dropout_probability, config.is_training + ) + output_dtype = quantizers.o.q_dtype if mode == "delayed" else qkv[0].dtype + s_scale = _tensor_scale(quantizers.s) + s_scale_inv = jnp.reciprocal(s_scale) + output_scale_inv = _tensor_scale_inv(quantizers.o) + raw_output, stats, amax, rng_state, _ = execute_fp8_fwd( + q_data, + k_data, + v_data, + _graph_scale_inv(q_tensor, mode), + _graph_scale_inv(k_tensor, mode), + _graph_scale_inv(v_tensor, mode, colwise=mode == "mxfp8"), + s_scale_inv, + s_scale, + _tensor_scale(quantizers.o), + q_seqlen, + kv_seqlen, + seed, + config=config, + mode=mode, + output_dtype=output_dtype, + ) + if mode == "delayed": + quantizers.s.update(amax[0:1]) + quantizers.o.update(amax[1:2]) + output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype( + qkv[0].dtype + ) + else: + output = raw_output + return output, ( + quantized, + raw_output, + stats, + rng_state, + q_seqlen, + kv_seqlen, + quantizers, + s_scale, + s_scale_inv, + output_scale_inv, + ) + + +def _split_gradient_outputs(dq, dk, dv, layout): + if layout.is_qkvpacked(): + return (jnp.stack((dq, dk, dv), axis=-3),) + if layout.is_kvpacked(): + return dq, jnp.stack((dk, dv), axis=-3) + return dq, dk, dv + + +def fused_attn_fp8_bwd(ctx, doutput, config): + """Execute FP8 attention backward and update delayed-scaling quantizers.""" + + ( + quantized, + raw_output, + stats, + rng_state, + q_seqlen, + kv_seqlen, + quantizers, + s_scale, + s_scale_inv, + output_scale_inv, + ) = ctx + q_tensor, k_tensor, v_tensor = quantized + mode = _scaling_mode(quantizers.qkv) + input_dtype = _rowwise(q_tensor).dq_dtype + (do_tensor,) = _quantize_many((doutput,), quantizers.do, both=mode == "mxfp8") + q_data, k_data, v_data = ( + _rowwise(q_tensor).data, + _rowwise(k_tensor).data, + _rowwise(v_tensor).data, + ) + empty_data = jnp.zeros((0,), dtype=q_data.dtype) + empty_scale = jnp.ones((1,), dtype=jnp.float32) + if config.qkv_layout.is_qkvpacked(): + k_data = v_data = empty_data + elif config.qkv_layout.is_kvpacked(): + v_data = empty_data + if mode == "mxfp8": + q_t, k_t = _colwise(q_tensor).data, _colwise(k_tensor).data + do_t = _colwise(do_tensor).data + q_scale_t, k_scale_t = ( + _graph_scale_inv(q_tensor, mode, colwise=True), + _graph_scale_inv(k_tensor, mode, colwise=True), + ) + do_scale_t = _graph_scale_inv(do_tensor, mode, colwise=True) + else: + q_t = k_t = do_t = empty_data + q_scale_t = k_scale_t = do_scale_t = empty_scale + + grad_dtype = quantizers.dqkv.q_dtype if mode == "delayed" else input_dtype + grad_scale_inv = _tensor_scale_inv(quantizers.dqkv) + dq, dk, dv, amax, _, _ = execute_fp8_bwd( + q_data, + k_data, + v_data, + q_t, + k_t, + stats, + rng_state, + raw_output, + _rowwise(do_tensor).data, + q_seqlen, + kv_seqlen, + _graph_scale_inv(q_tensor, mode), + _graph_scale_inv(k_tensor, mode), + _graph_scale_inv(v_tensor, mode), + output_scale_inv, + _graph_scale_inv(do_tensor, mode), + s_scale_inv, + _tensor_scale_inv(quantizers.dp), + s_scale, + _tensor_scale(quantizers.dqkv), + _tensor_scale(quantizers.dqkv), + _tensor_scale(quantizers.dqkv), + _tensor_scale(quantizers.dp), + do_t, + do_scale_t, + q_scale_t, + k_scale_t, + doutput, + config=config, + mode=mode, + grad_dtype=grad_dtype, + ) + if mode == "delayed": + quantizers.dqkv.update(jnp.max(amax[:3])) + quantizers.dp.update(amax[3:4]) + dq, dk, dv = ( + (tensor.astype(input_dtype) * grad_scale_inv).astype(input_dtype) + for tensor in (dq, dk, dv) + ) + return _split_gradient_outputs(dq, dk, dv, config.qkv_layout), quantizers diff --git a/transformer_engine/jax/cpp_extensions/gemm.py b/transformer_engine/jax/cpp_extensions/gemm.py index 438eaa16ea7..77ad82f52a0 100644 --- a/transformer_engine/jax/cpp_extensions/gemm.py +++ b/transformer_engine/jax/cpp_extensions/gemm.py @@ -47,6 +47,7 @@ noop_quantizer_set, is_fp8_gemm_with_all_layouts_supported, apply_padding_to_scale_inv, + swizzle_mxfp8_scale, QuantizeLayout, ) from .misc import get_padded_spec, is_all_reduce_in_float32, get_min_device_compute_capability @@ -398,21 +399,6 @@ def create(forward_collective_op: CollectiveOp): noop_collective_op_set = CollectiveOpSet.create(forward_collective_op=CollectiveOp.NONE) -@partial(jax.jit, static_argnums=(1, 2)) -def swizzled_scale(scale_inv, flatten_axis, is_colwise): - "Swizzle scale_inv via JAX transpose ops" - original_shape = scale_inv.shape - shape_2d = (math.prod(original_shape[:flatten_axis]), math.prod(original_shape[flatten_axis:])) - if is_colwise: - scale_inv = jnp.transpose(scale_inv.reshape(shape_2d)) - cols, rows = shape_2d - else: - rows, cols = shape_2d - reshape = scale_inv.reshape(rows // 128, 4, 32, cols // 4, 4) - swizzled = jnp.transpose(reshape, (0, 3, 2, 1, 4)) - return swizzled.reshape(original_shape) - - def get_lhs_axis_boundary(lhs_cdims, is_transposed): """Get the axis boundary for the LHS operand.""" return max(lhs_cdims) + 1 if is_transposed else min(lhs_cdims) @@ -759,8 +745,12 @@ def impl( # Only perform JAX-based swizzle for MXFP8, NVFP4 swizzle will go though nvte kernel if scaling_mode.is_mxfp8_scaling: - lhs_scale_inv = swizzled_scale(lhs_scale_inv, lhs_flatten_axis, lhs_transposed) - rhs_scale_inv = swizzled_scale(rhs_scale_inv, rhs_flatten_axis, not rhs_transposed) + lhs_scale_inv = swizzle_mxfp8_scale( + lhs_scale_inv, lhs_flatten_axis, lhs_transposed + ) + rhs_scale_inv = swizzle_mxfp8_scale( + rhs_scale_inv, rhs_flatten_axis, not rhs_transposed + ) # Determine if we need to reorder the tensor so that the input/output are in the correct layout for the collective operation need_reorder = not transpose_batch_sequence and not is_outer and not collective_op.is_none diff --git a/transformer_engine/jax/csrc/extensions/pybind.cpp b/transformer_engine/jax/csrc/extensions/pybind.cpp index effa4382aa4..6cf57543896 100644 --- a/transformer_engine/jax/csrc/extensions/pybind.cpp +++ b/transformer_engine/jax/csrc/extensions/pybind.cpp @@ -180,7 +180,8 @@ PYBIND11_MODULE(transformer_engine_jax, m) { pybind11::enum_(m, "NVTE_Bias_Type", pybind11::module_local()) .value("NVTE_NO_BIAS", NVTE_Bias_Type::NVTE_NO_BIAS) .value("NVTE_PRE_SCALE_BIAS", NVTE_Bias_Type::NVTE_PRE_SCALE_BIAS) - .value("NVTE_POST_SCALE_BIAS", NVTE_Bias_Type::NVTE_POST_SCALE_BIAS); + .value("NVTE_POST_SCALE_BIAS", NVTE_Bias_Type::NVTE_POST_SCALE_BIAS) + .value("NVTE_ALIBI", NVTE_Bias_Type::NVTE_ALIBI); pybind11::enum_(m, "NVTE_Mask_Type", pybind11::module_local()) .value("NVTE_NO_MASK", NVTE_Mask_Type::NVTE_NO_MASK) diff --git a/transformer_engine/jax/flax/module.py b/transformer_engine/jax/flax/module.py index 17c9a242f04..5f6bb268d15 100644 --- a/transformer_engine/jax/flax/module.py +++ b/transformer_engine/jax/flax/module.py @@ -37,6 +37,7 @@ jax_scaled_upper_triang_masked_softmax, ) from ..quantize import ( + AttentionQuantizerSet, QuantizerFactory, get_global_quantize_recipe, QuantizeMetaSet, @@ -417,6 +418,23 @@ def generate_quantizer_set( ) return quantizer_set + def generate_attention_quantizer_set(self, fp8_recipe=None): + """Generate independent quantizers for all FP8 DPA tensor roles.""" + + if fp8_recipe is None: + fp8_recipe = get_global_quantize_recipe() + first = self.generate_quantizer_set(postfix="_attention_qkv_s_do", fp8_recipe=fp8_recipe) + second = self.generate_quantizer_set(postfix="_attention_o_dp", fp8_recipe=fp8_recipe) + third = self.generate_quantizer_set(postfix="_attention_dqkv", fp8_recipe=fp8_recipe) + return AttentionQuantizerSet( + qkv=first.x, + s=first.kernel, + o=second.x, + do=first.dgrad, + dp=second.dgrad, + dqkv=third.dgrad, + ) + class DenseGeneral(TransformerEngineBase): r""" diff --git a/transformer_engine/jax/flax/transformer.py b/transformer_engine/jax/flax/transformer.py index 4b497826cc7..27ba24d5a8d 100644 --- a/transformer_engine/jax/flax/transformer.py +++ b/transformer_engine/jax/flax/transformer.py @@ -21,7 +21,7 @@ from jax import lax, vmap from jax.ad_checkpoint import checkpoint_name -from .module import DenseGeneral, LayerNormDenseGeneral, LayerNormMLP +from .module import DenseGeneral, LayerNormDenseGeneral, LayerNormMLP, TransformerEngineBase from .module import LayerNorm, Softmax from ..attention import ( AttnBiasType, @@ -33,6 +33,7 @@ from ..attention import is_fused_attn_kernel_available, make_swa_mask, canonicalize_attn_mask_type from ..attention import fused_attn from ..attention import CPStrategy +from ..quantize import get_global_quantize_recipe from ..softmax import SoftmaxFusionType from ..sharding import num_of_devices from ..sharding import get_sharding_map_logic_axis_to_mesh_axis @@ -291,7 +292,7 @@ def convert_to_softmax_fusion_type(attn_mask_type, mask): return jnp.einsum("bhqk,bkhd->bqhd", attn_weights, value) -class _FusedDotProductAttention(nn.Module): # pylint: disable=too-few-public-methods +class _FusedDotProductAttention(TransformerEngineBase): # pylint: disable=too-few-public-methods attention_dropout: float = 0.0 attn_mask_type: AttnMaskType = AttnMaskType.CAUSAL_MASK attn_bias_type: Optional[AttnBiasType] = None @@ -309,6 +310,7 @@ class _FusedDotProductAttention(nn.Module): # pylint: disable=too-few-public-me score_mod_bprop: Optional[Callable] = None score_mod_requested: bool = False return_max_logit: bool = False + bottom_right_diagonal: Optional[bool] = None @nn.compact def __call__( @@ -365,7 +367,18 @@ def __call__( "score_mod_tensors": score_mod_tensors, "score_mod_bprop_tensors": score_mod_bprop_tensors, "return_max_logit": self.return_max_logit, + "bottom_right_diagonal": self.bottom_right_diagonal, } + fp8_recipe = get_global_quantize_recipe() + if fp8_recipe is not None and getattr(fp8_recipe, "fp8_dpa", False): + if getattr(fp8_recipe, "fp8_mha", False): + raise NotImplementedError( + "JAX FP8 attention currently supports FP16/BF16 DPA boundaries only; " + "fp8_mha is not implemented." + ) + fused_attn_kwargs["quantizer_set"] = self.generate_attention_quantizer_set( + fp8_recipe + ) if self.qkv_layout.is_qkvpacked(): """qkvpacked format, treat @@ -629,6 +642,9 @@ class DotProductAttention(nn.Module): # pylint: disable=too-few-public-methods return_max_logit: bool, default = False If True, return ``(output, max_logit)`` where ``max_logit`` contains the per-head maximum attention logits with shape ``[h]``. This path requires fused attention. + bottom_right_diagonal: Optional[bool], default = None + Explicit diagonal alignment for fused attention. When unset, bottom-right mask types + use bottom-right alignment and other masks use top-left alignment. Optimization parameters ----------------------- @@ -658,6 +674,7 @@ class DotProductAttention(nn.Module): # pylint: disable=too-few-public-methods score_mod: Optional[Callable] = None score_mod_bprop: Optional[Callable] = None return_max_logit: bool = False + bottom_right_diagonal: Optional[bool] = None def __post_init__(self): # TODO(KshitijLakhani): Remove warning in TransformerEngine v2.12 @@ -758,10 +775,10 @@ def __call__( or score_mod_bprop_tensors is not None ) - if attn_bias_type == AttnBiasType.NO_BIAS: + if attn_bias_type in (AttnBiasType.NO_BIAS, AttnBiasType.ALIBI): assert ( bias is None - ), f"bias must be None when attn_bias_type is NO_BIAS, but got bias={bias}" + ), f"bias must be None when attn_bias_type is {attn_bias_type}, but got bias={bias}" else: assert ( bias is not None @@ -936,6 +953,7 @@ def __call__( score_mod_bprop=self.score_mod_bprop, score_mod_requested=score_mod_requested, return_max_logit=self.return_max_logit, + bottom_right_diagonal=self.bottom_right_diagonal, )( query, key, diff --git a/transformer_engine/jax/quantize/helper.py b/transformer_engine/jax/quantize/helper.py index 3a93af4a68d..0540e1edc2f 100644 --- a/transformer_engine/jax/quantize/helper.py +++ b/transformer_engine/jax/quantize/helper.py @@ -14,7 +14,7 @@ from enum import Enum import hashlib from typing import Optional, Tuple, Dict, Union, Sequence, Type, List -from functools import reduce +from functools import partial, reduce import operator import warnings @@ -58,6 +58,7 @@ "update_collections", "apply_padding_to_scale_inv", "remove_padding_from_scale_inv", + "swizzle_mxfp8_scale", "NVTE_FP8_COLLECTION_NAME", "TensorSource", ] @@ -69,6 +70,24 @@ NVTE_FP8_COLLECTION_NAME = "fp8_metas" +@partial(jax.jit, static_argnums=(1, 2)) +def swizzle_mxfp8_scale(scale_inv, flatten_axis, is_colwise): + """Convert a padded MXFP8 scale tensor to the F8_128x4 physical layout.""" + + original_shape = scale_inv.shape + shape_2d = ( + reduce(operator.mul, original_shape[:flatten_axis], 1), + reduce(operator.mul, original_shape[flatten_axis:], 1), + ) + if is_colwise: + scale_inv = jnp.transpose(scale_inv.reshape(shape_2d)) + cols, rows = shape_2d + else: + rows, cols = shape_2d + reshaped = scale_inv.reshape(rows // 128, 4, 32, cols // 4, 4) + return jnp.transpose(reshaped, (0, 3, 2, 1, 4)).reshape(original_shape) + + def _check_delayed_scaling_fp8_support(gpu_arch) -> Tuple[bool, str]: """Check if delayed scaling FP8 is supported on the given GPU architecture. diff --git a/transformer_engine/jax/quantize/quantizer.py b/transformer_engine/jax/quantize/quantizer.py index db56db935dc..85506475539 100644 --- a/transformer_engine/jax/quantize/quantizer.py +++ b/transformer_engine/jax/quantize/quantizer.py @@ -39,6 +39,7 @@ __all__ = [ "Quantizer", "QuantizerSet", + "AttentionQuantizerSet", "CurrentScaleQuantizer", "DelayedScaleQuantizer", "BlockScaleQuantizer", @@ -876,6 +877,29 @@ def tree_unflatten(cls, aux_data, children): return cls(*aux_data, *children) +@register_pytree_node_class +@dataclass +class AttentionQuantizerSet: + """Quantizers for the six independent FP8 dot-product-attention roles.""" + + qkv: Quantizer + s: Quantizer + o: Quantizer + do: Quantizer + dp: Quantizer + dqkv: Quantizer + + def tree_flatten(self): + """Flatten all quantizers so delayed-scaling state participates in autodiff.""" + + return (self.qkv, self.s, self.o, self.do, self.dp, self.dqkv), () + + @classmethod + def tree_unflatten(cls, aux_data, children): + del aux_data + return cls(*children) + + @register_pytree_node_class @dataclass class GroupedQuantizer(Quantizer): diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py index 2c5159a5358..8d8fffcc91a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py @@ -9,11 +9,10 @@ import warnings from transformer_engine.common.attention.cudnn import ( - AttentionLayout, FusedAttentionConfig, check_f16_fused_attention_support, - encode_cudnn_version, - requires_64bit_ragged_offset, + check_fp8_fused_attention_support, + parse_attention_layout, ) from transformer_engine.pytorch.constants import DType from transformer_engine.pytorch.utils import ( @@ -22,46 +21,15 @@ ) -def _layout_info(qkv_layout: str): - paged = qkv_layout.startswith("paged_kv_") - layout = qkv_layout.removeprefix("paged_kv_") - components = layout.split("_") - - def tensor_format(component: str) -> str: - return "".join(char for char in component if char.isalpha()) - - q_format = tensor_format(components[0]) - kv_format = tensor_format(components[-1]) if len(components) > 1 else q_format - qkv_format = q_format if q_format == kv_format else f"{q_format}_2{kv_format}" - if paged: - layout_group = "paged_separate" - elif len(components) == 1 and "3" in components[0]: - layout_group = "h3d" if "h3d" in components[0] else "3hd" - elif len(components) == 2: - layout_group = "hd_h2d" if "h2d" in components[1] else "hd_2hd" - elif q_format == "bhsd": - layout_group = "sd_sd_sd" - else: - layout_group = "separate" - return qkv_format, q_format, kv_format, layout_group - - -def _normalized_layout(qkv_layout: str) -> AttentionLayout: - qkv_format, q_format, kv_format, layout_group = _layout_info(qkv_layout) - return AttentionLayout( - qkv_format=qkv_format, - q_format=q_format, - kv_format=kv_format, - layout_group=layout_group, - is_qkvpacked=layout_group in ("3hd", "h3d"), - ) - - def _dtype_name(dtype: DType) -> str: if dtype == DType.kFloat16: return "float16" if dtype == DType.kBFloat16: return "bfloat16" + if dtype == DType.kFloat8E4M3: + return "float8_e4m3" + if dtype == DType.kFloat8E5M2: + return "float8_e5m2" return dtype.name @@ -99,78 +67,36 @@ def get_fused_attn_backend( major, minor = get_device_compute_capability() sm_arch = major * 10 + minor cudnn_version_tuple = get_cudnn_version() - cudnn_version = encode_cudnn_version(cudnn_version_tuple) - layout = _normalized_layout(qkv_layout) - requires_i64 = requires_64bit_ragged_offset( - layout, - num_attn_heads, - num_gqa_groups, - max_seqlen_q, - max_seqlen_kv, - head_dim_qk, - head_dim_v, - ) + layout = parse_attention_layout(qkv_layout) - # FP8 and MXFP8 remain PyTorch-specific. The shared policy below covers the - # F16/BF16 subset implemented by both framework frontends. fp8_dtype = q_dtype in (DType.kFloat8E4M3, DType.kFloat8E5M2) - fp8_shape_mask = ( - ( - cudnn_version >= 90201 - and sm_arch < 100 - and max_seqlen_q % 128 == 0 - and max_seqlen_kv % 128 == 0 - and head_dim_qk == 128 - and head_dim_v == 128 - and attn_mask_type in ("causal", "no_mask") - ) - or ( - cudnn_version >= 90700 - and ( - ( - sm_arch < 100 - and not is_training - and head_dim_qk <= 256 - and head_dim_v <= 256 - ) - or ( - sm_arch < 100 - and is_training - and head_dim_qk == 128 - and head_dim_v == 128 - ) - or (sm_arch >= 100 and head_dim_qk <= 128 and head_dim_v <= 128) + if fp8_dtype: + support = check_fp8_fused_attention_support( + FusedAttentionConfig( + is_training=bool(is_training), + q_dtype=_dtype_name(q_dtype), + kv_dtype=_dtype_name(kv_dtype), + layout=layout, + bias_type=bias_type, + mask_type=attn_mask_type, + softmax_type=softmax_type, + dropout=float(dropout), + num_attn_heads=int(num_attn_heads), + num_gqa_groups=int(num_gqa_groups), + max_seqlen_q=int(max_seqlen_q), + max_seqlen_kv=int(max_seqlen_kv), + head_dim_qk=int(head_dim_qk), + head_dim_v=int(head_dim_v), + window_size=(int(window_size_left), int(window_size_right)), + return_max_logit=bool(return_max_logit), + cuda_graph=bool(cuda_graph), + deterministic=bool(deterministic), + cudnn_version=cudnn_version_tuple, + sm_arch=sm_arch, ) - and head_dim_qk % 16 == 0 - and head_dim_v % 16 == 0 - and attn_mask_type in ("no_mask", "causal", "padding", "padding_causal") ) - or ( - cudnn_version >= 92100 - and sm_arch >= 100 - and head_dim_qk <= 192 - and head_dim_v <= 128 - and head_dim_qk % 16 == 0 - and head_dim_v % 16 == 0 - and attn_mask_type in ("no_mask", "causal", "causal_bottom_right") - ) - ) - fp8_format_softmax = ( - cudnn_version < 92100 - and layout.qkv_format in ("bshd", "sbhd") - and softmax_type == "vanilla" - ) or (cudnn_version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd")) - if ( - fp8_dtype - and sm_arch >= 90 - and bias_type == "no_bias" - and fp8_shape_mask - and fp8_format_softmax - and not requires_i64 - and cudnn_version != 91000 - and not return_max_logit - ): - return FusedAttnBackend.FP8 + if support.supported: + return FusedAttnBackend.FP8 support = check_f16_fused_attention_support( FusedAttentionConfig( @@ -194,7 +120,6 @@ def get_fused_attn_backend( deterministic=bool(deterministic), cudnn_version=cudnn_version_tuple, sm_arch=sm_arch, - allow_alibi=True, ) ) if support.warning is not None: diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 9a29ae478cb..d9ee2d6428a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -21,14 +21,20 @@ # graph construction, backend selection, and execution do not call TE common. import transformer_engine_torch as tex -from transformer_engine.common.attention.cudnn import normalize_attention_mask +from transformer_engine.common.attention.cudnn import cudnn_mask_options from transformer_engine.common.attention.cudnn import ( ragged_batch_bucket as _max_ragged_batch, ) from transformer_engine.common.attention.cudnn import ( ragged_token_bucket as _max_ragged_tokens, ) -from transformer_engine.common.attention.cudnn import round_up as _round_up +from transformer_engine.common.attention.fp8 import ( + FP8AttentionGraphConfig, + attention_format_stride as _format_stride, + build_fp8_backward_operation, + build_fp8_forward_operation, + mxfp8_padded_sizes as _mxfp8_padded_sizes, +) from transformer_engine.pytorch.constants import ( DType, FP8BwdTensorIdx, @@ -36,6 +42,7 @@ TE_DType_To_Torch, ) from transformer_engine.pytorch.quantized_tensor import QuantizedTensorStorage +from transformer_engine.pytorch.utils import get_cudnn_version from transformer_engine.pytorch.tensor.float8_tensor import ( Float8CurrentScalingQuantizer, Float8Quantizer, @@ -156,16 +163,6 @@ def _constant_graph_tensor(graph, cudnn, name: str): return _scalar_graph_tensor(graph, cudnn, name) -def _format_stride(batch: int, heads: int, seqlen: int, dim: int, tensor_format: str): - if tensor_format in ("bshd", "thd"): - return (seqlen * heads * dim, dim, heads * dim, 1) - if tensor_format == "sbhd": - return (heads * dim, dim, batch * heads * dim, 1) - if tensor_format == "bhsd": - return (heads * seqlen * dim, seqlen * dim, dim, 1) - raise ValueError(f"Unsupported FP8 tensor format {tensor_format!r}.") - - def _padded_sequence_lengths(cu_seqlens: torch.Tensor, batch: int) -> torch.Tensor: lengths = _sequence_lengths(cu_seqlens) if lengths.numel() == batch: @@ -190,19 +187,6 @@ def _element_ragged_offsets( return offsets * multiplier -def _mxfp8_padded_sizes(s_q: int, s_kv: int, d_qk: int, d_v: int) -> Dict[str, int]: - return { - "s_q_padded": _round_up(s_q, 128), - "s_kv_padded": _round_up(s_kv, 128), - "s_q_scale_padded": _round_up((s_q + 31) // 32, 4), - "s_kv_scale_padded": _round_up((s_kv + 31) // 32, 4), - "d_qk_padded": _round_up(d_qk, 128), - "d_v_padded": _round_up(d_v, 128), - "d_qk_scale_padded": _round_up((d_qk + 31) // 32, 4), - "d_v_scale_padded": _round_up((d_v + 31) // 32, 4), - } - - def _make_mxfp8_scale_tensor( graph, cudnn, @@ -440,7 +424,7 @@ def _mask_options( max_seqlen_q: int, max_seqlen_kv: int, ) -> Dict[str, Any]: - mask = normalize_attention_mask( + options = cudnn_mask_options( causal=attn_mask_type in ("causal", "padding_causal"), bottom_right=attn_mask_type in ("causal_bottom_right", "padding_causal_bottom_right"), @@ -450,26 +434,13 @@ def _mask_options( window_size=window_size, max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, + cudnn_version=get_cudnn_version(), + ) + options["diagonal_alignment"] = ( + cudnn.diagonal_alignment.BOTTOM_RIGHT + if options["diagonal_alignment"] == "bottom_right" + else cudnn.diagonal_alignment.TOP_LEFT ) - - options: Dict[str, Any] = { - "use_causal_mask": mask.causal, - "use_causal_mask_bottom_right": mask.bottom_right, - "diagonal_alignment": ( - cudnn.diagonal_alignment.BOTTOM_RIGHT - if mask.bottom_right_diagonal - else cudnn.diagonal_alignment.TOP_LEFT - ), - } - if mask.window_left != -1: - options["diagonal_band_left_bound"] = mask.window_left + 1 - # ``use_causal_mask`` already imposes a right bound of zero. The Python - # frontend rejects specifying that same bound through both attributes. - if mask.window_right != -1 and not ( - (mask.causal or mask.bottom_right) and mask.window_right == 0 - ): - options["diagonal_band_right_bound"] = mask.window_right - options["is_padding"] = mask.padding return options @@ -1285,16 +1256,20 @@ def _build_fp8_fwd_graph( tensor_format=scale_format_kv, ) tensors.update(descale_q=descale_q, descale_k=descale_k, descale_v=descale_v) - output_t, stats_t, amax_o_t = graph.sdpa_mxfp8( - q_t, - k_t, - v_t, - descale_q, - descale_k, - descale_v, - name="te_sdpa_mxfp8", - **options, - ) + op = build_fp8_forward_operation( + graph, + { + "q": q_t, + "k": k_t, + "v": v_t, + "descale_q": descale_q, + "descale_k": descale_k, + "descale_v": descale_v, + }, + options, + FP8AttentionGraphConfig("mxfp8", "te_sdpa_mxfp8"), + ) + output_t, stats_t, amax_o_t = op["output"], op["stats"], op["amax_o"] amax_o_t.set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( (1, 1, 1, 1) ).set_stride((1, 1, 1, 1)) @@ -1321,19 +1296,27 @@ def _build_fp8_fwd_graph( else: scale_o = _constant_graph_tensor(graph, cudnn, "Current_Scale_O") tensors["constant_scale_o"] = scale_o - output_t, stats_t, amax_s_t, amax_o_t = graph.sdpa_fp8( - q_t, - k_t, - v_t, - descale_q, - descale_k, - descale_v, - descale_s, - scale_s, - scale_o, - name="te_sdpa_fp8", - **options, - ) + op = build_fp8_forward_operation( + graph, + { + "q": q_t, + "k": k_t, + "v": v_t, + "descale_q": descale_q, + "descale_k": descale_k, + "descale_v": descale_v, + "descale_s": descale_s, + "scale_s": scale_s, + "scale_o": scale_o, + }, + options, + FP8AttentionGraphConfig( + "delayed" if isinstance(o_quantizer, Float8Quantizer) else "current", + "te_sdpa_fp8", + ), + ) + output_t, stats_t = op["output"], op["stats"] + amax_s_t, amax_o_t = op["amax_s"], op["amax_o"] amax_s_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( (1, 1, 1, 1) ).set_stride((1, 1, 1, 1)) @@ -2194,28 +2177,31 @@ def mx_scale(name, h, s, d, fmt): descale_do_t = mx_scale( "descale_do_t", heads, "s_q_scale_padded", "d_v_padded", scale_format_do ) - outputs = graph.sdpa_mxfp8_backward( - q_t, - q_col_t, - k_t, - k_col_t, - v_t, - o_t, - do_f16_t, - do_t, - do_col_t, - stats_t, - descale_q, - descale_q_t, - descale_k, - descale_k_t, - descale_v, - descale_do, - descale_do_t, - name="te_sdpa_mxfp8_backward", - **options, - ) - dq_t, dk_t, dv_t, *amax_outputs = outputs + op = build_fp8_backward_operation( + graph, + { + "q": q_t, + "q_t": q_col_t, + "k": k_t, + "k_t": k_col_t, + "v": v_t, + "o": o_t, + "do_f16": do_f16_t, + "do": do_t, + "do_t": do_col_t, + "stats": stats_t, + "descale_q": descale_q, + "descale_q_t": descale_q_t, + "descale_k": descale_k, + "descale_k_t": descale_k_t, + "descale_v": descale_v, + "descale_do": descale_do, + "descale_do_t": descale_do_t, + }, + options, + FP8AttentionGraphConfig("mxfp8", "te_sdpa_mxfp8_backward"), + ) + dq_t, dk_t, dv_t, amax_outputs = op["dq"], op["dk"], op["dv"], op["amax"] for amax_t in amax_outputs: amax_t.set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( (1, 1, 1, 1) @@ -2283,29 +2269,36 @@ def mx_scale(name, h, s, d, fmt): if "scale_dv" in tensors else tensors["constant_scale_dv"] ) - outputs = graph.sdpa_fp8_backward( - q_t, - k_t, - v_t, - o_t, - do_t, - stats_t, - tensors["descale_q"], - tensors["descale_k"], - tensors["descale_v"], - descale_o_arg, - tensors["descale_do"], - descale_s_arg, - descale_dp_arg, - scale_s_arg, - scale_dq_arg, - scale_dk_arg, - scale_dv_arg, - scale_dp_arg, - name="te_sdpa_fp8_backward", - **options, - ) - dq_t, dk_t, dv_t, amax_dq_t, amax_dk_t, amax_dv_t, amax_dp_t = outputs + op = build_fp8_backward_operation( + graph, + { + "q": q_t, + "k": k_t, + "v": v_t, + "o": o_t, + "do": do_t, + "stats": stats_t, + "descale_q": tensors["descale_q"], + "descale_k": tensors["descale_k"], + "descale_v": tensors["descale_v"], + "descale_o": descale_o_arg, + "descale_do": tensors["descale_do"], + "descale_s": descale_s_arg, + "descale_dp": descale_dp_arg, + "scale_s": scale_s_arg, + "scale_dq": scale_dq_arg, + "scale_dk": scale_dk_arg, + "scale_dv": scale_dv_arg, + "scale_dp": scale_dp_arg, + }, + options, + FP8AttentionGraphConfig( + "delayed" if delayed else "current", "te_sdpa_fp8_backward" + ), + ) + dq_t, dk_t, dv_t = op["dq"], op["dk"], op["dv"] + amax_dq_t, amax_dk_t = op["amax_dq"], op["amax_dk"] + amax_dv_t, amax_dp_t = op["amax_dv"], op["amax_dp"] for name, amax_t in ( ("amax_dq", amax_dq_t), ("amax_dk", amax_dk_t), From d6f38aabe916ebc645e592e5597aba77d9c1fa15 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 01:17:12 +0000 Subject: [PATCH 17/36] [JAX] Fix attention regression failures Update the score-mod fused-attention test double for the newly propagated bottom-right diagonal argument so plumbing tests continue to model the public call signature. Provide dtype, dimensions, and strides for unused MXFP8 amax tensors in forward and backward cuDNN graphs. This satisfies frontend tensor validation while keeping those tensors excluded from graph outputs. Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fused_attn_score_mod.py | 2 ++ transformer_engine/jax/cpp_extensions/fp8_attention.py | 8 ++++++-- 2 files changed, 8 insertions(+), 2 deletions(-) diff --git a/tests/jax/test_fused_attn_score_mod.py b/tests/jax/test_fused_attn_score_mod.py index 86596826007..181a6344499 100644 --- a/tests/jax/test_fused_attn_score_mod.py +++ b/tests/jax/test_fused_attn_score_mod.py @@ -427,6 +427,7 @@ def fake_fused_attn( score_mod_tensors=None, score_mod_bprop_tensors=None, return_max_logit=False, + bottom_right_diagonal=None, ): captured.update( qkv=qkv, @@ -452,6 +453,7 @@ def fake_fused_attn( score_mod_bprop=score_mod_bprop, score_mod_tensors=score_mod_tensors, score_mod_bprop_tensors=score_mod_bprop_tensors, + bottom_right_diagonal=bottom_right_diagonal, ) return qkv[0] diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index bfc76c7a8fb..2becff11275 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -412,7 +412,9 @@ def build_fp8_fwd_graph( ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 2, offset)) else: - op["amax_o"].set_output(False) + op["amax_o"].set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) workspace, data, version = finalize_graph( cudnn, graph, description=f"JAX {mode} FP8 attention forward" @@ -722,7 +724,9 @@ def build_fp8_bwd_graph( output_bindings.append(GraphBinding(uid, 3, offset)) else: for amax in op["amax"]: - amax.set_output(False) + amax.set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) workspace, data, version = finalize_graph( cudnn, graph, description=f"JAX {mode} FP8 attention backward" From 13ae9af6fcef8d0fc0433726fe71d5c3f50d03aa Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 02:08:14 +0000 Subject: [PATCH 18/36] [JAX] Fix MXFP8 attention scale strides JAX transposes compact MXFP8 scale-inverse buffers to contiguous BHSD before applying the cuDNN F8_128x4 swizzle. Describe those buffers with matching BHSD strides instead of the original BSHD data layout, preventing invalid GPU memory accesses during MXFP8 attention. Add a focused regression test for the graph scale dimensions, strides, and reordering contract. The complete JAX FP8 attention test file passes on an SM100 CUDA device. Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fp8_fused_attn.py | 42 ++++++++++++++++++- .../jax/cpp_extensions/fp8_attention.py | 5 ++- 2 files changed, 44 insertions(+), 3 deletions(-) diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py index 0817eb034da..6156a9813ba 100644 --- a/tests/jax/test_fp8_fused_attn.py +++ b/tests/jax/test_fp8_fused_attn.py @@ -23,7 +23,10 @@ fused_attn, ) from transformer_engine.jax.cpp_extensions import FusedAttnHelper -from transformer_engine.jax.cpp_extensions.fp8_attention import _mxfp8_scale_inv +from transformer_engine.jax.cpp_extensions.fp8_attention import ( + _mx_scale, + _mxfp8_scale_inv, +) from transformer_engine.jax.flax import DotProductAttention from transformer_engine.jax.quantize import ( BlockScaleQuantizer, @@ -181,6 +184,43 @@ def test_mxfp8_attention_scale_layout(): assert _mxfp8_scale_inv(tensor, colwise=True).shape == (2, 8, 4, 128) +def test_mxfp8_attention_scale_graph_stride(): + """The cuDNN scale descriptor matches the contiguous BHSD swizzle buffer.""" + + class FakeTensor: + def set_reordering_type(self, reordering): + self.reordering = reordering + return self + + class FakeGraph: + def tensor(self, **kwargs): + self.kwargs = kwargs + return FakeTensor() + + class FakeCudnn: + class data_type: + FP8_E8M0 = "fp8_e8m0" + + class tensor_reordering: + F8_128x4 = "f8_128x4" + + graph = FakeGraph() + tensor = _mx_scale( + graph, + FakeCudnn, + name="descale_q", + uid=101, + batch=2, + heads=8, + seqlen=128, + dim=4, + ) + + assert graph.kwargs["dim"] == (2, 8, 128, 4) + assert graph.kwargs["stride"] == (4096, 512, 4, 1) + assert tensor.reordering == FakeCudnn.tensor_reordering.F8_128x4 + + @pytest.mark.parametrize("feature", ("alibi", "bottom_right")) def test_fused_attention_parity_features(feature): """JAX executes ALiBi and explicit bottom-right diagonal attention.""" diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index 2becff11275..b8e55a85319 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -171,13 +171,14 @@ def _mx_scale( heads, seqlen, dim, - tensor_format="bshd", ): + # _mxfp8_scale_inv transposes every compact scale buffer to contiguous BHSD + # before applying cuDNN's F8_128x4 physical reordering. return _tensor( graph, name=name, dim=(batch, heads, seqlen, dim), - stride=attention_format_stride(batch, heads, seqlen, dim, tensor_format), + stride=attention_format_stride(batch, heads, seqlen, dim, "bhsd"), dtype=cudnn.data_type.FP8_E8M0, uid=uid, ).set_reordering_type(cudnn.tensor_reordering.F8_128x4) From 028c5831c66f5741ef02f3a1910bec3a75678415 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 03:51:01 +0000 Subject: [PATCH 19/36] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/jax/test_fp8_fused_attn.py | 34 +- tests/jax/test_fused_attn.py | 5 +- tests/test_common_attention_helpers.py | 24 +- transformer_engine/common/attention/cudnn.py | 23 +- .../common/attention/score_mod.py | 8 +- transformer_engine/common/cudnn_frontend.py | 4 +- .../common/fused_attn/fused_attn.cpp | 4 +- transformer_engine/common/util/utils.cu | 5 +- transformer_engine/jax/attention.py | 8 +- .../jax/cpp_extensions/attention.py | 19 +- .../jax/cpp_extensions/cudnn_attention.py | 94 ++--- .../jax/cpp_extensions/cudnn_graph.py | 8 +- .../jax/cpp_extensions/flex_attention.py | 6 +- .../jax/cpp_extensions/fp8_attention.py | 84 ++--- transformer_engine/jax/cpp_extensions/gemm.py | 8 +- .../jax/csrc/extensions/attention.cpp | 104 +++--- .../jax/csrc/extensions/attention_kernels.cu | 2 +- transformer_engine/jax/flax/transformer.py | 4 +- .../dot_product_attention/cudnn_attention.py | 346 +++++------------- .../dot_product_attention/flex_attention.py | 32 +- 20 files changed, 256 insertions(+), 566 deletions(-) diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py index 6156a9813ba..d8844659a33 100644 --- a/tests/jax/test_fp8_fused_attn.py +++ b/tests/jax/test_fp8_fused_attn.py @@ -55,9 +55,9 @@ def _require_gpu(min_arch=90, min_cudnn=90700): def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): q_seqlen, kv_seqlen = q.shape[1], k.shape[1] - scores = jnp.einsum( - "bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32) - ) / sqrt(q.shape[-1]) + scores = jnp.einsum("bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32)) / sqrt( + q.shape[-1] + ) q_pos = jnp.arange(q_seqlen)[:, None] kv_pos = jnp.arange(kv_seqlen)[None, :] shift = kv_seqlen - q_seqlen if bottom_right else 0 @@ -77,9 +77,7 @@ def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): allowed = kv_pos <= q_pos + shift scores = jnp.where(allowed[None, None, :, :], scores, -jnp.inf) probabilities = jax.nn.softmax(scores, axis=-1) - return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype( - q.dtype - ) + return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype(q.dtype) def _assert_fp8_close(actual, expected): @@ -112,9 +110,7 @@ def _assert_fp8_close(actual, expected): ), ), ) -@pytest.mark.parametrize( - "input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16") -) +@pytest.mark.parametrize("input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16")) def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): """Each supported recipe executes FP8 DPA behind FP16/BF16 module boundaries.""" @@ -138,9 +134,7 @@ def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): ) def loss_fn(variables, query, key, value): - output = module.apply( - variables, query, key, value, descriptor, deterministic=True - ) + output = module.apply(variables, query, key, value, descriptor, deterministic=True) loss = jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)) return loss, output @@ -150,9 +144,7 @@ def reference_loss(query, key, value): return loss, output with autocast(enabled=True, recipe=fp8_recipe, mesh_resource=MeshResource()): - variables = module.init( - jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=True - ) + variables = module.init(jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=True) (_, output), (_, dq, dk, dv) = jax.value_and_grad( loss_fn, argnums=(0, 1, 2, 3), has_aux=True )(variables, q, k, v) @@ -175,9 +167,7 @@ def test_mxfp8_attention_scale_layout(): q_layout=QuantizeLayout.ROWWISE_COLWISE, data_layout="NN", ) - tensor = quantizer.quantize( - jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2 - ) + tensor = quantizer.quantize(jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2) assert tensor.rowwise_tensor.scale_inv.shape == (2, 64, 8, 2) assert tensor.colwise_tensor.scale_inv.shape == (2, 2, 8, 64) assert _mxfp8_scale_inv(tensor).shape == (2, 8, 128, 4) @@ -232,9 +222,7 @@ def test_fused_attention_parity_features(feature): q = jax.random.normal(q_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 k = jax.random.normal(k_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 v = jax.random.normal(v_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 - doutput = ( - jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 - ) + doutput = jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 q_lengths = jnp.full((batch,), q_seqlen, dtype=jnp.int32) kv_lengths = jnp.full((batch,), kv_seqlen, dtype=jnp.int32) descriptor = SequenceDescriptor.from_seqlens((q_lengths, kv_lengths)) @@ -286,9 +274,7 @@ def reference_loss(query, key, value): ) return jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)), output - (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)( - q, k, v - ) + (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)(q, k, v) (_, reference), reference_grads = jax.value_and_grad( reference_loss, argnums=(0, 1, 2), has_aux=True )(q, k, v) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 493ddecc4b6..a251ddbdb07 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -390,9 +390,7 @@ def score_mod(_graph, score, _tensors): ((9, 6, 0), (32, 2048, 2048, True)), ], ) -def test_thd_graph_bucketing_requires_cudnn_9_6( - monkeypatch, cudnn_version, expected_dimensions -): +def test_thd_graph_bucketing_requires_cudnn_9_6(monkeypatch, cudnn_version, expected_dimensions): """Pre-9.6 THD graphs retain dense dimensions and metadata extents.""" monkeypatch.setattr(cudnn_attention, "get_cudnn_version", lambda: cudnn_version) monkeypatch.setattr(cudnn_attention, "_device_arch", lambda: 90) @@ -406,6 +404,7 @@ def test_thd_graph_bucketing_requires_cudnn_9_6( qk_dim=128, v_dim=128, ) + class Config: qkv_layout = QKVLayout.THD_THD_THD max_segments_per_seq = 4 diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py index d44aeb74caa..db0912e60ba 100644 --- a/tests/test_common_attention_helpers.py +++ b/tests/test_common_attention_helpers.py @@ -127,21 +127,11 @@ def test_shared_f16_policy_framework_parity_features(): def test_shared_fp8_policy(): fp8 = _attention_config(q_dtype="float8_e4m3", kv_dtype="float8_e4m3", sm_arch=100) assert check_fp8_fused_attention_support(fp8).supported - assert check_fp8_fused_attention_support( - replace(fp8, cudnn_version=(9, 10, 1)) - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, cudnn_version=(9, 10, 0)) - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, bias_type="alibi") - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, return_max_logit=True) - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, head_dim_qk=200) - ).supported + assert check_fp8_fused_attention_support(replace(fp8, cudnn_version=(9, 10, 1))).supported + assert not check_fp8_fused_attention_support(replace(fp8, cudnn_version=(9, 10, 0))).supported + assert not check_fp8_fused_attention_support(replace(fp8, bias_type="alibi")).supported + assert not check_fp8_fused_attention_support(replace(fp8, return_max_logit=True)).supported + assert not check_fp8_fused_attention_support(replace(fp8, head_dim_qk=200)).supported @pytest.mark.parametrize( @@ -384,6 +374,4 @@ def test_shared_cudnn_graph_creation_and_finalization(): def test_shared_cudnn_graph_support_error_has_context(): with pytest.raises(RuntimeError, match="cuDNN test graph is not supported"): - build_cudnn_graph( - _FakeCudnn(), _FakeGraph(unsupported=True), description="test" - ) + build_cudnn_graph(_FakeCudnn(), _FakeGraph(unsupported=True), description="test") diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index cb2cce09775..9d0b19e74dd 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -191,9 +191,7 @@ def cudnn_mask_options( ) version = encode_cudnn_version(cudnn_version) options: dict[str, bool | int | str] = { - "diagonal_alignment": ( - "bottom_right" if mask.bottom_right_diagonal else "top_left" - ), + "diagonal_alignment": ("bottom_right" if mask.bottom_right_diagonal else "top_left"), "is_padding": mask.padding, } if version < 90600: @@ -356,9 +354,7 @@ def check_f16_fused_attention_support( and (dqk, dv) != (192, 128) and dqk != dv ): - return _unsupported( - "this Hopper backward head-dimension combination is unsupported" - ) + return _unsupported("this Hopper backward head-dimension combination is unsupported") alibi_supported = ( bias == "alibi" @@ -438,10 +434,7 @@ def check_f16_fused_attention_support( and bias == "no_bias" and dropout == 0.0 ) - or ( - mask in ("causal_bottom_right", "padding_causal_bottom_right") - and sq <= skv - ) + or (mask in ("causal_bottom_right", "padding_causal_bottom_right") and sq <= skv) ) mask_ok = mask_ok or modern_mask_ok if not mask_ok: @@ -473,10 +466,7 @@ def check_f16_fused_attention_support( or ( left >= -1 and right == 0 - and ( - mask in ("no_mask", "causal") - or (mask == "causal_bottom_right" and sq == skv) - ) + and (mask in ("no_mask", "causal") or (mask == "causal_bottom_right" and sq == skv)) and sq <= skv and dropout == 0.0 and bias == "no_bias" @@ -657,10 +647,7 @@ def check_fp8_fused_attention_support( version < 92100 and layout.qkv_format in ("bshd", "sbhd") and config.softmax_type == "vanilla" - ) or ( - version >= 92100 - and layout.qkv_format in ("bshd", "sbhd", "bhsd") - ) + ) or (version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd")) if not format_softmax_ok: return _unsupported("FP8 attention layout or softmax type is not supported") return FusedAttentionSupport(True) diff --git a/transformer_engine/common/attention/score_mod.py b/transformer_engine/common/attention/score_mod.py index 5e9e770dcca..dc9c16f71bc 100644 --- a/transformer_engine/common/attention/score_mod.py +++ b/transformer_engine/common/attention/score_mod.py @@ -48,9 +48,7 @@ def freeze_score_mod_cache_key(value: Any, *, is_array: Callable[[Any], bool]) - ) return tuple(sorted(items, key=repr)) if isinstance(value, (list, tuple)): - return tuple( - freeze_score_mod_cache_key(item, is_array=is_array) for item in value - ) + return tuple(freeze_score_mod_cache_key(item, is_array=is_array) for item in value) if isinstance(value, (set, frozenset)): items = (freeze_score_mod_cache_key(item, is_array=is_array) for item in value) return tuple(sorted(items, key=repr)) @@ -64,9 +62,7 @@ def freeze_score_mod_cache_key(value: Any, *, is_array: Callable[[Any], bool]) - return value -def _explicit_cache_key( - callback_owner: Any, *, is_array: Callable[[Any], bool] -) -> Any | None: +def _explicit_cache_key(callback_owner: Any, *, is_array: Callable[[Any], bool]) -> Any | None: explicit_key = getattr(callback_owner, "score_mod_graph_cache_key", None) if explicit_key is None: return None diff --git a/transformer_engine/common/cudnn_frontend.py b/transformer_engine/common/cudnn_frontend.py index a1ddb084a38..c6a5ac5c5fc 100644 --- a/transformer_engine/common/cudnn_frontend.py +++ b/transformer_engine/common/cudnn_frontend.py @@ -51,8 +51,6 @@ def build_cudnn_graph(cudnn, graph, *, description: str) -> int: graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) graph.check_support() except cudnn.cudnnGraphNotSupportedError as exc: - raise RuntimeError( - f"cuDNN {description} graph is not supported: {exc}" - ) from exc + raise RuntimeError(f"cuDNN {description} graph is not supported: {exc}") from exc graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) return max(int(graph.get_workspace_size()), 1) diff --git a/transformer_engine/common/fused_attn/fused_attn.cpp b/transformer_engine/common/fused_attn/fused_attn.cpp index c3601e8af6c..c6caf156486 100644 --- a/transformer_engine/common/fused_attn/fused_attn.cpp +++ b/transformer_engine/common/fused_attn/fused_attn.cpp @@ -4,10 +4,10 @@ * See LICENSE for license information. ************************************************************************/ -#include - #include "transformer_engine/fused_attn.h" +#include + #include "../common.h" namespace transformer_engine { diff --git a/transformer_engine/common/util/utils.cu b/transformer_engine/common/util/utils.cu index ae3d1da8ed4..f659759163d 100644 --- a/transformer_engine/common/util/utils.cu +++ b/transformer_engine/common/util/utils.cu @@ -17,9 +17,8 @@ namespace transformer_engine { namespace extract_seed_and_offset { namespace { -__global__ void kernel(int64_t *rng_state_ptr, bool captured, int64_t *seed_ptr, - uint64_t seed_val, int64_t *offset_ptr, uint64_t offset_val, - uint32_t offset_intragraph) { +__global__ void kernel(int64_t *rng_state_ptr, bool captured, int64_t *seed_ptr, uint64_t seed_val, + int64_t *offset_ptr, uint64_t offset_val, uint32_t offset_intragraph) { if (captured) { rng_state_ptr[0] = *seed_ptr; rng_state_ptr[1] = static_cast(*offset_ptr + static_cast(offset_intragraph)); diff --git a/transformer_engine/jax/attention.py b/transformer_engine/jax/attention.py index 40fac63393b..2a0e8c2fd36 100644 --- a/transformer_engine/jax/attention.py +++ b/transformer_engine/jax/attention.py @@ -1433,16 +1433,12 @@ def _fused_attn_bwd_rule( @partial(jax.custom_vjp, nondiff_argnums=(4,)) def _fused_attn_fp8(qkv, sequence_descriptor, seed, quantizer_set, config): - output, _ = tex.fused_attn_fp8_fwd( - qkv, sequence_descriptor, seed, quantizer_set, config - ) + output, _ = tex.fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizer_set, config) return output def _fused_attn_fp8_fwd_rule(qkv, sequence_descriptor, seed, quantizer_set, config): - return tex.fused_attn_fp8_fwd( - qkv, sequence_descriptor, seed, quantizer_set, config - ) + return tex.fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizer_set, config) def _fused_attn_fp8_bwd_rule(config, ctx, doutput): diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index 3b697f569b1..905ef15cd06 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -277,9 +277,7 @@ def _multiply_offsets_as_uint64_words(offsets, multiplier): carry = (product_lo >> 16) + (product_mid_lo & mask) + (product_mid_hi & mask) low_word = (product_lo & mask) | ((carry & mask) << 16) - high_word = ( - product_hi + (product_mid_lo >> 16) + (product_mid_hi >> 16) + (carry >> 16) - ) + high_word = product_hi + (product_mid_lo >> 16) + (product_mid_hi >> 16) + (carry >> 16) return jnp.stack((low_word, high_word), axis=-1) @@ -425,9 +423,7 @@ def abstract( rng_state_shape = (seed_aval.shape[0], checker.rng_state_size) rng_state_aval = seed_aval.update(shape=rng_state_shape, dtype=checker.rng_state_dtype) - wkspace_aval = q_aval.update( - shape=(graph_info.graph.workspace_size,), dtype=jnp.uint8 - ) + wkspace_aval = q_aval.update(shape=(graph_info.graph.workspace_size,), dtype=jnp.uint8) assert ( softmax_offset_aval.dtype == jnp.float32 @@ -965,9 +961,7 @@ def abstract( dk_aval = k_aval.update(shape=k_aval.shape, dtype=k_dtype) dv_aval = v_aval.update(shape=v_aval.shape, dtype=v_dtype) dbias_aval = bias_aval.update(shape=bias_aval.shape, dtype=bias_dtype) - wkspace_aval = q_aval.update( - shape=(graph_info.graph.workspace_size,), dtype=jnp.uint8 - ) + wkspace_aval = q_aval.update(shape=(graph_info.graph.workspace_size,), dtype=jnp.uint8) # Validate incoming softmax_offset shape and dtype assert ( @@ -3781,9 +3775,10 @@ def fused_attn_bwd( raise ValueError(f"Unknown {qkv_layout=}") if attn_bias_type in (AttnBiasType.NO_BIAS, AttnBiasType.ALIBI): - assert ( - bias is None - ), f"bias must be None when attn_bias_type is {attn_bias_type}, but got bias with type={type(bias)}" + assert bias is None, ( + f"bias must be None when attn_bias_type is {attn_bias_type}, but got bias with" + f" type={type(bias)}" + ) bias = jnp.zeros(0, dtype=qkv[0].dtype) if softmax_offset is None: diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index db974d1def6..1678bfe84b1 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -92,9 +92,7 @@ def _layout_info(q_aval, k_aval, v_aval, layout) -> _LayoutInfo: if layout.is_qkvpacked(): *batch_shape, q_seqlen, packed, q_heads, qk_dim = q_aval.shape if packed != 3: - raise ValueError( - f"QKV-packed fused attention expects dimension 3, got {q_aval.shape}." - ) + raise ValueError(f"QKV-packed fused attention expects dimension 3, got {q_aval.shape}.") kv_seqlen = q_seqlen kv_heads = q_heads v_dim = qk_dim @@ -102,27 +100,17 @@ def _layout_info(q_aval, k_aval, v_aval, layout) -> _LayoutInfo: *batch_shape, q_seqlen, q_heads, qk_dim = q_aval.shape *kv_batch_shape, kv_seqlen, packed, kv_heads, v_dim = k_aval.shape if tuple(batch_shape) != tuple(kv_batch_shape) or packed != 2: - raise ValueError( - f"Invalid KV-packed fused-attention shapes: {q_aval}, {k_aval}." - ) + raise ValueError(f"Invalid KV-packed fused-attention shapes: {q_aval}, {k_aval}.") if qk_dim != v_dim: - raise ValueError( - "KV-packed fused attention requires equal QK and V head dimensions." - ) + raise ValueError("KV-packed fused attention requires equal QK and V head dimensions.") elif layout.is_separate(): *batch_shape, q_seqlen, q_heads, qk_dim = q_aval.shape *k_batch_shape, kv_seqlen, kv_heads, k_dim = k_aval.shape *v_batch_shape, v_seqlen, v_heads, v_dim = v_aval.shape - if tuple(batch_shape) != tuple(k_batch_shape) or tuple(batch_shape) != tuple( - v_batch_shape - ): - raise ValueError( - "Separate Q, K and V tensors must have matching batch shapes." - ) + if tuple(batch_shape) != tuple(k_batch_shape) or tuple(batch_shape) != tuple(v_batch_shape): + raise ValueError("Separate Q, K and V tensors must have matching batch shapes.") if qk_dim != k_dim or kv_seqlen != v_seqlen or kv_heads != v_heads: - raise ValueError( - "Separate fused-attention K and V shapes are inconsistent." - ) + raise ValueError("Separate fused-attention K and V shapes are inconsistent.") else: raise ValueError(f"Unsupported JAX fused-attention layout: {layout}.") return _LayoutInfo( @@ -137,9 +125,7 @@ def _layout_info(q_aval, k_aval, v_aval, layout) -> _LayoutInfo: ) -def _matrix_stride( - info: _LayoutInfo, layout, matrix: str, graph_sq: int, graph_skv: int -): +def _matrix_stride(info: _LayoutInfo, layout, matrix: str, graph_sq: int, graph_skv: int): """Port generateMatrixStrides for JAX's BSHD/THD layout subset.""" if matrix in ("q", "o"): heads = info.q_heads @@ -242,17 +228,13 @@ def _graph_dimensions(info: _LayoutInfo, config): arch = _device_arch() use_ragged_stats = is_ragged and cudnn_version >= (9, 6, 0) and arch != 120 if is_ragged: - graph_batch = ragged_graph_batch_size( - info.input_batch, config.max_segments_per_seq - ) + graph_batch = ragged_graph_batch_size(info.input_batch, config.max_segments_per_seq) if cudnn_version < (9, 6, 0) or arch == 120: graph_sq = info.q_max_seqlen graph_skv = info.kv_max_seqlen else: graph_sq = _ragged_graph_token_count(info.input_batch * info.q_max_seqlen) - graph_skv = _ragged_graph_token_count( - info.input_batch * info.kv_max_seqlen - ) + graph_skv = _ragged_graph_token_count(info.input_batch * info.kv_max_seqlen) else: graph_batch = info.input_batch graph_sq = info.q_max_seqlen @@ -313,9 +295,7 @@ def _ragged_offset_spec(cudnn): def _mask_options(cudnn, info: _LayoutInfo, config): cp_striped_window_size = getattr(config, "cp_striped_window_size", None) window_left, window_right = ( - cp_striped_window_size - if cp_striped_window_size is not None - else config.window_size + cp_striped_window_size if cp_striped_window_size is not None else config.window_size ) options = cudnn_mask_options( causal=_is_causal(config), @@ -369,8 +349,8 @@ def build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGraph def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGraphInfo: cudnn = import_cudnn() info = _layout_info(q_aval, k_aval, v_aval, config.qkv_layout) - graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = ( - _graph_dimensions(info, config) + graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = _graph_dimensions( + info, config ) io_dtype = cudnn_data_type(cudnn, q_aval.dtype) graph = make_graph(cudnn, io_dtype) @@ -403,13 +383,9 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap uid=_UID_V, ) - input_bindings = list( - _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) - ) + input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) output_bindings = [GraphBinding(_UID_O, 0), GraphBinding(_UID_STATS, 1)] - scale = _scalar_tensor( - graph, cudnn, "attn_scale", _UID_ATTN_SCALE, cudnn.data_type.FLOAT - ) + scale = _scalar_tensor(graph, cudnn, "attn_scale", _UID_ATTN_SCALE, cudnn.data_type.FLOAT) scalar_uids = [_UID_ATTN_SCALE] scalar_values = [np.asarray(config.scaling_factor, dtype=np.float32).tobytes()] @@ -473,9 +449,7 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap uid=_UID_SEQ_KV, ) kwargs.update(use_padding_mask=True, seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend( - (GraphBinding(_UID_SEQ_Q, 6), GraphBinding(_UID_SEQ_KV, 7)) - ) + input_bindings.extend((GraphBinding(_UID_SEQ_Q, 6), GraphBinding(_UID_SEQ_KV, 7))) offset_q = offset_k = offset_v = offset_o = offset_stats = None if config.qkv_layout.is_thd(): @@ -588,9 +562,7 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap else: stats.set_stride((info.q_heads * graph_sq, graph_sq, 1, 1)) - workspace, data, version = finalize_graph( - cudnn, graph, description="fused-attention forward" - ) + workspace, data, version = finalize_graph(cudnn, graph, description="fused-attention forward") result = serialized_graph( serialized_graph_data=data, cudnn_frontend_version=version, @@ -651,16 +623,14 @@ def _build_bwd_graph( ) -> AttentionGraphInfo: cudnn = import_cudnn() info = _layout_info(q_aval, k_aval, v_aval, config.qkv_layout) - graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = ( - _graph_dimensions(info, config) + graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = _graph_dimensions( + info, config ) io_dtype = cudnn_data_type(cudnn, q_aval.dtype) graph = make_graph(cudnn, io_dtype) def io_tensor(name, dim, stride, uid, dtype=io_dtype): - return _tensor( - graph, cudnn, name=name, dim=dim, stride=stride, dtype=dtype, uid=uid - ) + return _tensor(graph, cudnn, name=name, dim=dim, stride=stride, dtype=dtype, uid=uid) q = io_tensor( "q", @@ -713,12 +683,8 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): GraphBinding(_UID_DO, 8), ) ) - output_bindings = list( - _qkv_bindings(info, config.qkv_layout, itemsize, outputs=True) - ) - scale = _scalar_tensor( - graph, cudnn, "attn_scale", _UID_ATTN_SCALE, cudnn.data_type.FLOAT - ) + output_bindings = list(_qkv_bindings(info, config.qkv_layout, itemsize, outputs=True)) + scale = _scalar_tensor(graph, cudnn, "attn_scale", _UID_ATTN_SCALE, cudnn.data_type.FLOAT) scalar_uids = [_UID_ATTN_SCALE] scalar_values = [np.asarray(config.scaling_factor, dtype=np.float32).tobytes()] @@ -741,11 +707,7 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): ) if ragged_stats: kwargs["max_total_seq_len_q"] = graph_sq - if ( - config.qkv_layout.is_thd() - and get_cudnn_version() >= (9, 6, 0) - and _device_arch() != 120 - ): + if config.qkv_layout.is_thd() and get_cudnn_version() >= (9, 6, 0) and _device_arch() != 120: kwargs["max_total_seq_len_kv"] = graph_skv if _is_bias(config): @@ -806,9 +768,7 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): cudnn.data_type.INT32, ) kwargs.update(use_padding_mask=True, seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend( - (GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10)) - ) + input_bindings.extend((GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10))) if config.qkv_layout.is_thd(): offset_dtype, offset_itemsize = _ragged_offset_spec(cudnn) @@ -891,9 +851,7 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): dk.set_ragged_offset(offset_k) dv.set_ragged_offset(offset_v) - workspace, data, version = finalize_graph( - cudnn, graph, description="fused-attention backward" - ) + workspace, data, version = finalize_graph(cudnn, graph, description="fused-attention backward") result = serialized_graph( serialized_graph_data=data, cudnn_frontend_version=version, @@ -961,9 +919,7 @@ def is_fused_attn_supported(helper) -> bool: window_size=tuple(int(value) for value in helper.window_size), return_max_logit=bool(helper.return_max_logit), cuda_graph=False, - deterministic=not bool( - int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) - ), + deterministic=not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))), cudnn_version=get_cudnn_version(), sm_arch=_device_arch(), ) diff --git a/transformer_engine/jax/cpp_extensions/cudnn_graph.py b/transformer_engine/jax/cpp_extensions/cudnn_graph.py index 8ba743e9079..b5e3a6d4d62 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_graph.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_graph.py @@ -155,9 +155,7 @@ def encode_cudnn_frontend_version(version: str) -> int: public_version = version.split("+", 1)[0].split("-", 1)[0] parts = public_version.split(".") if len(parts) < 3: - raise RuntimeError( - f"Could not parse cuDNN frontend Python version: {version!r}." - ) + raise RuntimeError(f"Could not parse cuDNN frontend Python version: {version!r}.") major, minor, patch = (int(part) for part in parts[:3]) return major * 10000 + minor * 100 + patch @@ -229,9 +227,7 @@ def serialized_graph( scalar_sizes, packed_scalar_values = pack_scalar_values(scalar_values) def binding_array(bindings, field): - return np.asarray( - [getattr(binding, field) for binding in bindings], dtype=np.int64 - ) + return np.asarray([getattr(binding, field) for binding in bindings], dtype=np.int64) return SerializedGraph( serialized_graph=serialized_graph_data, diff --git a/transformer_engine/jax/cpp_extensions/flex_attention.py b/transformer_engine/jax/cpp_extensions/flex_attention.py index c018182585e..de366c8b279 100644 --- a/transformer_engine/jax/cpp_extensions/flex_attention.py +++ b/transformer_engine/jax/cpp_extensions/flex_attention.py @@ -369,12 +369,10 @@ def _serialized_score_mod_graph( cudnn_frontend_version=int(cudnn_frontend_version), workspace_size=int(workspace_size), input_bindings=[ - GraphBinding(uid=int(uid), buffer_index=index) - for index, uid in enumerate(input_uids) + GraphBinding(uid=int(uid), buffer_index=index) for index, uid in enumerate(input_uids) ], output_bindings=[ - GraphBinding(uid=int(uid), buffer_index=index) - for index, uid in enumerate(output_uids) + GraphBinding(uid=int(uid), buffer_index=index) for index, uid in enumerate(output_uids) ], scalar_uids=np.asarray(scalar_uids, dtype=np.int64), scalar_values=scalar_values, diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index b8e55a85319..669c7aafb29 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -199,9 +199,7 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="Q", dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.qk_dim), - stride=_matrix_stride( - info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen - ), + stride=_matrix_stride(info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen), dtype=io_dtype, uid=_UID_Q, ) @@ -209,9 +207,7 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="K", dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.qk_dim), - stride=_matrix_stride( - info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen - ), + stride=_matrix_stride(info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen), dtype=io_dtype, uid=_UID_K, ) @@ -219,9 +215,7 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="V", dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.v_dim), - stride=_matrix_stride( - info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen - ), + stride=_matrix_stride(info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen), dtype=io_dtype, uid=_UID_V, ) @@ -264,9 +258,7 @@ def build_fp8_fwd_graph( graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) info, _, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) - input_bindings = list( - _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) - ) + input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) tensors = {"q": q, "k": k, "v": v} options, is_padding = _fp8_options(cudnn, info, config) options["generate_stats"] = True @@ -280,9 +272,7 @@ def build_fp8_fwd_graph( options["diagonal_band_left_bound"] = options.pop("left_bound") if "right_bound" in options: options["diagonal_band_right_bound"] = options.pop("right_bound") - padded = mxfp8_padded_sizes( - info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim - ) + padded = mxfp8_padded_sizes(info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim) scale_specs = ( ( "descale_q", @@ -351,9 +341,7 @@ def build_fp8_fwd_graph( uid=_UID_SEQ_KV, ) options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend( - (GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11)) - ) + input_bindings.extend((GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11))) output_bindings = [GraphBinding(_UID_O, 0), GraphBinding(_UID_STATS, 1)] if _is_dropout(config): @@ -390,27 +378,23 @@ def build_fp8_fwd_graph( output = op["output"] output.set_output(True).set_uid(_UID_O).set_data_type( cudnn_data_type(cudnn, output_dtype) - ).set_dim( - (info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim) - ).set_stride( + ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim)).set_stride( attention_format_stride( info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" ) ) stats = op["stats"] - stats.set_output(True).set_uid(_UID_STATS).set_data_type( - cudnn.data_type.FLOAT - ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, 1)).set_stride( - (info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1) - ) + stats.set_output(True).set_uid(_UID_STATS).set_data_type(cudnn.data_type.FLOAT).set_dim( + (info.input_batch, info.q_heads, info.q_max_seqlen, 1) + ).set_stride((info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1)) if mode != "mxfp8": for name, uid, offset in ( ("amax_s", _UID_AMAX_S, 0), ("amax_o", _UID_AMAX_O, 4), ): - op[name].set_output(True).set_uid(uid).set_data_type( - cudnn.data_type.FLOAT - ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) + op[name].set_output(True).set_uid(uid).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 2, offset)) else: op["amax_o"].set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( @@ -527,13 +511,9 @@ def build_fp8_bwd_graph( cudnn = import_cudnn() graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) - info, io_dtype, q, k, v = _graph_io_tensors( - graph, cudnn, q_aval, k_aval, v_aval, config - ) + info, io_dtype, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) - input_bindings = list( - _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) - ) + input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) o = _tensor( graph, name="O", @@ -588,9 +568,7 @@ def build_fp8_bwd_graph( uid=_UID_SEQ_KV, ) options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend( - (GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10)) - ) + input_bindings.extend((GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10))) if _is_dropout(config): seed = _tensor( @@ -669,9 +647,7 @@ def build_fp8_bwd_graph( GraphBinding(_UID_DO_F16, 27), ) ) - padded = mxfp8_padded_sizes( - info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim - ) + padded = mxfp8_padded_sizes(info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim) tensors.update(_mx_bwd_scales(graph, cudnn, info, padded, input_bindings)) else: for name, uid, index in ( @@ -719,9 +695,9 @@ def build_fp8_bwd_graph( ("amax_dv", _UID_AMAX_DV, 8), ("amax_dp", _UID_AMAX_DP, 12), ): - op[name].set_output(True).set_uid(uid).set_data_type( - cudnn.data_type.FLOAT - ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) + op[name].set_output(True).set_uid(uid).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 3, offset)) else: for amax in op["amax"]: @@ -1047,17 +1023,13 @@ def _validate_fp8_support(qkv, quantizers, config, mode): window_size=tuple(int(value) for value in config.window_size), return_max_logit=False, cuda_graph=False, - deterministic=not bool( - int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) - ), + deterministic=not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))), cudnn_version=get_cudnn_version(), sm_arch=_device_arch(), ) ) if not support.supported: - raise ValueError( - f"Unsupported JAX FP8 attention configuration: {support.reason}." - ) + raise ValueError(f"Unsupported JAX FP8 attention configuration: {support.reason}.") if mode == "mxfp8" and (get_cudnn_version() < (9, 21, 0) or _device_arch() < 100): raise ValueError("MXFP8 attention requires cuDNN 9.21 and SM100 or newer.") @@ -1080,15 +1052,11 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): if config.qkv_layout.is_thd(): raise NotImplementedError("FP8 attention does not support THD layouts in JAX.") if mode == "mxfp8" and not config.qkv_layout.is_separate(): - raise NotImplementedError( - "JAX MXFP8 attention currently requires separate BSHD Q/K/V." - ) + raise NotImplementedError("JAX MXFP8 attention currently requires separate BSHD Q/K/V.") if getattr(config.attn_bias_type, "name", "") != "NO_BIAS": raise NotImplementedError("FP8 attention does not support attention bias.") if getattr(config.softmax_type, "name", "") != "VANILLA_SOFTMAX": - raise NotImplementedError( - "JAX FP8 attention currently supports vanilla softmax only." - ) + raise NotImplementedError("JAX FP8 attention currently supports vanilla softmax only.") _validate_fp8_support(qkv, quantizers, config, mode) if mode == "mxfp8" and _is_padding(config): raise NotImplementedError("JAX MXFP8 attention does not support padding masks.") @@ -1129,9 +1097,7 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): if mode == "delayed": quantizers.s.update(amax[0:1]) quantizers.o.update(amax[1:2]) - output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype( - qkv[0].dtype - ) + output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype(qkv[0].dtype) else: output = raw_output return output, ( diff --git a/transformer_engine/jax/cpp_extensions/gemm.py b/transformer_engine/jax/cpp_extensions/gemm.py index 58ef7591f7f..f4a8d2c5eea 100644 --- a/transformer_engine/jax/cpp_extensions/gemm.py +++ b/transformer_engine/jax/cpp_extensions/gemm.py @@ -745,12 +745,8 @@ def impl( # Only perform JAX-based swizzle for MXFP8, NVFP4 swizzle will go though nvte kernel if scaling_mode.is_mxfp8_scaling: - lhs_scale_inv = swizzle_mxfp8_scale( - lhs_scale_inv, lhs_flatten_axis, lhs_transposed - ) - rhs_scale_inv = swizzle_mxfp8_scale( - rhs_scale_inv, rhs_flatten_axis, not rhs_transposed - ) + lhs_scale_inv = swizzle_mxfp8_scale(lhs_scale_inv, lhs_flatten_axis, lhs_transposed) + rhs_scale_inv = swizzle_mxfp8_scale(rhs_scale_inv, rhs_flatten_axis, not rhs_transposed) # Determine if we need to reorder the tensor so that the input/output are in the correct layout for the collective operation need_reorder = not transpose_batch_sequence and not is_outer and not collective_op.is_none diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 130e8a9f663..11a49e0fe80 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -128,8 +128,8 @@ CudnnGraphPtr GetCudnnGraph(cudaStream_t stream, Dictionary &attrs) { auto graph = std::make_shared(); auto status = graph->deserialize(handle, serialized_data); - NVTE_CHECK(status.is_good(), "Failed to deserialize cuDNN frontend graph: ", - status.get_message()); + NVTE_CHECK(status.is_good(), + "Failed to deserialize cuDNN frontend graph: ", status.get_message()); std::lock_guard lock(GetCudnnGraphCacheMutex()); auto &cache = GetCudnnGraphCache(); @@ -176,18 +176,16 @@ Error_Type ExecuteCudnnGraph(cudaStream_t stream, Dictionary &attrs, static_cast(input_buffer_indices[i]) < input_ptrs.size(), "cuDNN graph input binding index is out of range."); NVTE_CHECK(input_byte_offsets[i] >= 0, "cuDNN graph input byte offset must be non-negative."); - auto *ptr = static_cast(input_ptrs[input_buffer_indices[i]]) + - input_byte_offsets[i]; + auto *ptr = static_cast(input_ptrs[input_buffer_indices[i]]) + input_byte_offsets[i]; variant_pack.emplace(input_uids[i], ptr); } for (size_t i = 0; i < output_uids.size(); ++i) { NVTE_CHECK(output_buffer_indices[i] >= 0 && static_cast(output_buffer_indices[i]) < output_ptrs.size(), "cuDNN graph output binding index is out of range."); - NVTE_CHECK(output_byte_offsets[i] >= 0, - "cuDNN graph output byte offset must be non-negative."); - auto *ptr = static_cast(output_ptrs[output_buffer_indices[i]]) + - output_byte_offsets[i]; + NVTE_CHECK(output_byte_offsets[i] >= 0, "cuDNN graph output byte offset must be non-negative."); + auto *ptr = + static_cast(output_ptrs[output_buffer_indices[i]]) + output_byte_offsets[i]; variant_pack.emplace(output_uids[i], ptr); } @@ -215,9 +213,7 @@ void AppendRemainingBuffers(Variadic_Buffer_Type args, std::vector *ptrs } } -size_t BufferBytes(const Buffer_Type &buffer) { - return buffer.size_bytes(); -} +size_t BufferBytes(const Buffer_Type &buffer) { return buffer.size_bytes(); } void MemsetResultAsync(cudaStream_t stream, Result_Type result, int value) { NVTE_CHECK_CUDA(cudaMemsetAsync(result->untyped_data(), value, BufferBytes(*result), stream)); @@ -251,13 +247,15 @@ void PopulateRngStateAsync(cudaStream_t stream, const Buffer_Type &seed, Result_ } // namespace -Error_Type FusedAttnForwardFFI( - cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, Buffer_Type v_buf, - Buffer_Type bias_buf, Buffer_Type softmax_offset_buf, Buffer_Type seed_buf, - Buffer_Type q_seqlens_buf, Buffer_Type kv_seqlens_buf, Buffer_Type q_seq_offsets_buf, - Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type remaining_args, Result_Type output_buf, - Result_Type stats_buf, Result_Type max_buf, Result_Type rng_state_buf, - Result_Type workspace_buf, Dictionary attrs) { +Error_Type FusedAttnForwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, + Buffer_Type v_buf, Buffer_Type bias_buf, + Buffer_Type softmax_offset_buf, Buffer_Type seed_buf, + Buffer_Type q_seqlens_buf, Buffer_Type kv_seqlens_buf, + Buffer_Type q_seq_offsets_buf, Buffer_Type k_seq_offsets_buf, + Variadic_Buffer_Type remaining_args, Result_Type output_buf, + Result_Type stats_buf, Result_Type max_buf, + Result_Type rng_state_buf, Result_Type workspace_buf, + Dictionary attrs) { const bool is_ragged = get_attr_value(attrs, "is_ragged"); const uint64_t rng_increment = static_cast(get_attr_value(attrs, "rng_offset_increment")); @@ -271,17 +269,21 @@ Error_Type FusedAttnForwardFFI( } std::vector input_ptrs = { - q_buf.untyped_data(), k_buf.untyped_data(), - v_buf.untyped_data(), bias_buf.untyped_data(), - softmax_offset_buf.untyped_data(), seed_buf.untyped_data(), - q_seqlens_buf.untyped_data(), kv_seqlens_buf.untyped_data(), - q_seq_offsets_buf.untyped_data(), k_seq_offsets_buf.untyped_data(), + q_buf.untyped_data(), + k_buf.untyped_data(), + v_buf.untyped_data(), + bias_buf.untyped_data(), + softmax_offset_buf.untyped_data(), + seed_buf.untyped_data(), + q_seqlens_buf.untyped_data(), + kv_seqlens_buf.untyped_data(), + q_seq_offsets_buf.untyped_data(), + k_seq_offsets_buf.untyped_data(), }; AppendRemainingBuffers(remaining_args, &input_ptrs); std::vector output_ptrs = {output_buf->untyped_data(), stats_buf->untyped_data(), max_buf->untyped_data(), rng_state_buf->untyped_data()}; - return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, - workspace_buf->untyped_data()); + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, workspace_buf->untyped_data()); } XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnForwardHandler, FusedAttnForwardFFI, @@ -306,35 +308,45 @@ XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnForwardHandler, FusedAttnForwardFFI, .Attrs(), FFI_CudaGraph_Traits); -Error_Type FusedAttnBackwardFFI( - cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, Buffer_Type v_buf, - Buffer_Type bias_buf, Buffer_Type softmax_offset_buf, Buffer_Type stats_buf, - Buffer_Type rng_state_buf, Buffer_Type output_buf, Buffer_Type doutput_buf, - Buffer_Type q_seqlens_buf, Buffer_Type kv_seqlens_buf, Buffer_Type q_seq_offsets_buf, - Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type remaining_args, Result_Type dq_buf, - Result_Type dk_buf, Result_Type dv_buf, Result_Type dbias_buf, - Result_Type dsoftmax_offset_buf, Result_Type workspace_buf, Dictionary attrs) { +Error_Type FusedAttnBackwardFFI(cudaStream_t stream, Buffer_Type q_buf, Buffer_Type k_buf, + Buffer_Type v_buf, Buffer_Type bias_buf, + Buffer_Type softmax_offset_buf, Buffer_Type stats_buf, + Buffer_Type rng_state_buf, Buffer_Type output_buf, + Buffer_Type doutput_buf, Buffer_Type q_seqlens_buf, + Buffer_Type kv_seqlens_buf, Buffer_Type q_seq_offsets_buf, + Buffer_Type k_seq_offsets_buf, Variadic_Buffer_Type remaining_args, + Result_Type dq_buf, Result_Type dk_buf, Result_Type dv_buf, + Result_Type dbias_buf, Result_Type dsoftmax_offset_buf, + Result_Type workspace_buf, Dictionary attrs) { if (get_attr_value(attrs, "is_ragged")) { MemsetResultAsync(stream, dq_buf, 0); MemsetResultAsync(stream, dk_buf, 0); MemsetResultAsync(stream, dv_buf, 0); } std::vector input_ptrs = { - q_buf.untyped_data(), k_buf.untyped_data(), - v_buf.untyped_data(), bias_buf.untyped_data(), - softmax_offset_buf.untyped_data(), stats_buf.untyped_data(), - rng_state_buf.untyped_data(), output_buf.untyped_data(), - doutput_buf.untyped_data(), q_seqlens_buf.untyped_data(), - kv_seqlens_buf.untyped_data(), q_seq_offsets_buf.untyped_data(), + q_buf.untyped_data(), + k_buf.untyped_data(), + v_buf.untyped_data(), + bias_buf.untyped_data(), + softmax_offset_buf.untyped_data(), + stats_buf.untyped_data(), + rng_state_buf.untyped_data(), + output_buf.untyped_data(), + doutput_buf.untyped_data(), + q_seqlens_buf.untyped_data(), + kv_seqlens_buf.untyped_data(), + q_seq_offsets_buf.untyped_data(), k_seq_offsets_buf.untyped_data(), }; AppendRemainingBuffers(remaining_args, &input_ptrs); std::vector output_ptrs = { - dq_buf->untyped_data(), dk_buf->untyped_data(), dv_buf->untyped_data(), - dbias_buf->untyped_data(), dsoftmax_offset_buf->untyped_data(), + dq_buf->untyped_data(), + dk_buf->untyped_data(), + dv_buf->untyped_data(), + dbias_buf->untyped_data(), + dsoftmax_offset_buf->untyped_data(), }; - return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, - workspace_buf->untyped_data()); + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, workspace_buf->untyped_data()); } XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnBackwardHandler, FusedAttnBackwardFFI, @@ -371,8 +383,7 @@ Error_Type FusedAttnScoreModForwardFFI(cudaStream_t stream, Buffer_Type q_buf, B v_buf.untyped_data()}; AppendRemainingBuffers(score_mod_args, &input_ptrs); std::vector output_ptrs = {output_buf->untyped_data(), stats_buf->untyped_data()}; - return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, - workspace_buf->untyped_data()); + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, workspace_buf->untyped_data()); } XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnScoreModForwardHandler, FusedAttnScoreModForwardFFI, @@ -400,8 +411,7 @@ Error_Type FusedAttnScoreModBackwardFFI(cudaStream_t stream, Buffer_Type q_buf, AppendRemainingBuffers(score_mod_args, &input_ptrs); std::vector output_ptrs = {dq_buf->untyped_data(), dk_buf->untyped_data(), dv_buf->untyped_data()}; - return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, - workspace_buf->untyped_data()); + return ExecuteCudnnGraph(stream, attrs, input_ptrs, output_ptrs, workspace_buf->untyped_data()); } XLA_FFI_DEFINE_HANDLER_SYMBOL(FusedAttnScoreModBackwardHandler, FusedAttnScoreModBackwardFFI, diff --git a/transformer_engine/jax/csrc/extensions/attention_kernels.cu b/transformer_engine/jax/csrc/extensions/attention_kernels.cu index a5d32fa8d9c..c2797292f30 100644 --- a/transformer_engine/jax/csrc/extensions/attention_kernels.cu +++ b/transformer_engine/jax/csrc/extensions/attention_kernels.cu @@ -11,7 +11,7 @@ namespace jax { namespace { __global__ void PopulateFusedAttnRngStateKernel(int64_t *rng_state, const int64_t *seed, - uint64_t offset) { + uint64_t offset) { rng_state[0] = seed[0]; rng_state[1] = static_cast(offset); } diff --git a/transformer_engine/jax/flax/transformer.py b/transformer_engine/jax/flax/transformer.py index 27ba24d5a8d..d8ceab1c84d 100644 --- a/transformer_engine/jax/flax/transformer.py +++ b/transformer_engine/jax/flax/transformer.py @@ -376,9 +376,7 @@ def __call__( "JAX FP8 attention currently supports FP16/BF16 DPA boundaries only; " "fp8_mha is not implemented." ) - fused_attn_kwargs["quantizer_set"] = self.generate_attention_quantizer_set( - fp8_recipe - ) + fused_attn_kwargs["quantizer_set"] = self.generate_attention_quantizer_set(fp8_recipe) if self.qkv_layout.is_qkvpacked(): """qkvpacked format, treat diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index d9ee2d6428a..2f7ebb7d056 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -133,14 +133,10 @@ def _quantized_scale_inv(tensor: Any, *, columnwise: bool = False) -> torch.Tens if _is_float8_tensor(tensor): return tensor._scale_inv if _is_mxfp8_tensor(tensor): - scale = ( - tensor._columnwise_scale_inv if columnwise else tensor._rowwise_scale_inv - ) + scale = tensor._columnwise_scale_inv if columnwise else tensor._rowwise_scale_inv if scale is None: orientation = "columnwise" if columnwise else "rowwise" - raise ValueError( - f"MXFP8 attention input has no {orientation} scale-inverse buffer." - ) + raise ValueError(f"MXFP8 attention input has no {orientation} scale-inverse buffer.") return scale raise TypeError(f"Expected an FP8 attention tensor, got {type(tensor).__name__}.") @@ -167,9 +163,7 @@ def _padded_sequence_lengths(cu_seqlens: torch.Tensor, batch: int) -> torch.Tens lengths = _sequence_lengths(cu_seqlens) if lengths.numel() == batch: return lengths - padding = torch.zeros( - batch - lengths.numel(), dtype=lengths.dtype, device=lengths.device - ) + padding = torch.zeros(batch - lengths.numel(), dtype=lengths.dtype, device=lengths.device) return torch.cat((lengths, padding)) @@ -230,9 +224,7 @@ def _allocate_fp8_kernel_output(quantizer, shape, fake_dtype, device): ) if isinstance(quantizer, MXFP8Quantizer): return torch.empty(shape, dtype=fake_dtype, device=device), None - raise TypeError( - f"Unsupported FP8 attention output quantizer {type(quantizer).__name__}." - ) + raise TypeError(f"Unsupported FP8 attention output quantizer {type(quantizer).__name__}.") def _format_from_layout_component(component: str) -> str: @@ -243,11 +235,7 @@ def _q_kv_formats(qkv_layout: str) -> Tuple[str, str]: layout = qkv_layout.removeprefix("paged_kv_") components = layout.split("_") q_format = _format_from_layout_component(components[0]) - kv_format = ( - _format_from_layout_component(components[-1]) - if len(components) > 1 - else q_format - ) + kv_format = _format_from_layout_component(components[-1]) if len(components) > 1 else q_format return q_format, kv_format @@ -343,9 +331,7 @@ def _allocate_output( def _storage_span(tensor: torch.Tensor) -> int: if tensor.numel() == 0: return 0 - return 1 + sum( - (size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride()) - ) + return 1 + sum((size - 1) * stride for size, stride in zip(tensor.shape, tensor.stride())) def _allocate_grad_views( @@ -363,22 +349,17 @@ def _allocate_grad_views( for indices in groups.values(): if len(indices) == 1: inp = inputs[indices[0]] - out = torch.empty_strided( - inp.shape, inp.stride(), dtype=inp.dtype, device=inp.device - ) + out = torch.empty_strided(inp.shape, inp.stride(), dtype=inp.dtype, device=inp.device) out.zero_() outputs[indices[0]] = out continue min_offset = min(inputs[index].storage_offset() for index in indices) max_end = max( - inputs[index].storage_offset() + _storage_span(inputs[index]) - for index in indices + inputs[index].storage_offset() + _storage_span(inputs[index]) for index in indices ) exemplar = inputs[indices[0]] - base = torch.zeros( - max_end - min_offset, dtype=exemplar.dtype, device=exemplar.device - ) + base = torch.zeros(max_end - min_offset, dtype=exemplar.dtype, device=exemplar.device) for index in indices: inp = inputs[index] outputs[index] = torch.as_strided( @@ -406,9 +387,7 @@ def _reserve_philox_state( # Development-tree compatibility before the local extension has been # rebuilt. Installed packages always expose the graph-safe helper above. if rng_gen is None: - index = ( - device.index if device.index is not None else torch.cuda.current_device() - ) + index = device.index if device.index is not None else torch.cuda.current_device() rng_gen = torch.cuda.default_generators[index] seed = rng_gen.initial_seed() offset = rng_gen.get_offset() @@ -426,10 +405,8 @@ def _mask_options( ) -> Dict[str, Any]: options = cudnn_mask_options( causal=attn_mask_type in ("causal", "padding_causal"), - bottom_right=attn_mask_type - in ("causal_bottom_right", "padding_causal_bottom_right"), - padding=attn_mask_type - in ("padding", "padding_causal", "padding_causal_bottom_right"), + bottom_right=attn_mask_type in ("causal_bottom_right", "padding_causal_bottom_right"), + padding=attn_mask_type in ("padding", "padding_causal", "padding_causal_bottom_right"), bottom_right_diagonal=bottom_right_diagonal, window_size=window_size, max_seqlen_q=max_seqlen_q, @@ -522,34 +499,25 @@ def _build_f16_fwd_graph( bottom_right_diagonal: bool, ) -> GraphEntry: cudnn = import_cudnn_frontend() - graph = make_graph( - torch_to_cudnn_dtype(q.dtype), q.device, name="te_fused_attention_fwd" - ) + graph = make_graph(torch_to_cudnn_dtype(q.dtype), q.device, name="te_fused_attention_fwd") q_format, kv_format = _q_kv_formats(qkv_layout) batch = cu_seqlens_q.numel() - 1 is_ragged_q = q_format == "thd" is_ragged_kv = kv_format == "thd" use_ragged_stats = is_ragged_q and cudnn.backend_version() >= 90600 - use_token_buckets = ( - cudnn.backend_version() >= 90600 - and torch.cuda.get_device_capability(q.device) != (12, 0) - ) + use_token_buckets = cudnn.backend_version() >= 90600 and torch.cuda.get_device_capability( + q.device + ) != (12, 0) use_direct_offsets = cudnn.backend_version() >= 92400 and dropout == 0.0 use_legacy_offsets = (is_ragged_q or is_ragged_kv) and not use_direct_offsets - graph_batch = ( - _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch - ) + graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch if not use_token_buckets: use_ragged_stats = False graph_seqlen_q = ( - _max_ragged_tokens(q.shape[0]) - if is_ragged_q and use_token_buckets - else max_seqlen_q + _max_ragged_tokens(q.shape[0]) if is_ragged_q and use_token_buckets else max_seqlen_q ) graph_seqlen_kv = ( - _max_ragged_tokens(k.shape[0]) - if is_ragged_kv and use_token_buckets - else max_seqlen_kv + _max_ragged_tokens(k.shape[0]) if is_ragged_kv and use_token_buckets else max_seqlen_kv ) tensors: Dict[str, Any] = {} @@ -772,16 +740,14 @@ def _build_f16_fwd_graph( data_type=torch.int64 if use_legacy_offsets else None, ) tensors["offset_stats"] = offset_stats - stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( - stats_dim - ).set_stride(stats_stride) + stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim(stats_dim).set_stride( + stats_stride + ) if use_ragged_stats: stats_t.set_ragged_offset(offset_stats).set_ragged_offset_multiplier(stats_mult) tensors["Stats"] = stats_t - return GraphEntry( - graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) - ) + return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) def _f16_forward( @@ -818,17 +784,12 @@ def _f16_forward( heads = q.shape[1] if q_format == "bhsd" else q.shape[-2] total_tokens_q = q.shape[0] if q_format == "thd" else batch * max_seqlen_q cudnn = import_cudnn_frontend() - use_token_buckets = ( - cudnn.backend_version() >= 90600 - and torch.cuda.get_device_capability(q.device) != (12, 0) - ) + use_token_buckets = cudnn.backend_version() >= 90600 and torch.cuda.get_device_capability( + q.device + ) != (12, 0) use_direct_offsets = cudnn.backend_version() >= 92400 and dropout == 0.0 - use_legacy_offsets = ( - q_format == "thd" or kv_format == "thd" - ) and not use_direct_offsets - graph_batch = ( - _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch - ) + use_legacy_offsets = (q_format == "thd" or kv_format == "thd") and not use_direct_offsets + graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch ragged_stats = ( q_format == "thd" and cudnn.backend_version() >= 90600 @@ -850,17 +811,11 @@ def _f16_forward( ) stats = torch.empty(stats_shape, dtype=torch.float32, device=q.device) max_scores = ( - torch.empty(stats_shape, dtype=torch.float32, device=q.device) - if return_max_logit - else None + torch.empty(stats_shape, dtype=torch.float32, device=q.device) if return_max_logit else None ) rng_state = _reserve_philox_state(q.device, rng_gen, _F16_RNG_ELTS_PER_THREAD) - cu_seqlens_q_padded = ( - cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded - ) - cu_seqlens_kv_padded = ( - cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded - ) + cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded + cu_seqlens_kv_padded = cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded key = _f16_fwd_key( is_training=is_training, @@ -889,9 +844,7 @@ def _f16_forward( entry = get_graph_entry(key) if entry is None: if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "cuDNN attention graph must be built before CUDA graph capture." - ) + raise RuntimeError("cuDNN attention graph must be built before CUDA graph capture.") entry = _build_f16_fwd_graph( is_training=is_training, max_seqlen_q=max_seqlen_q, @@ -936,12 +889,8 @@ def _f16_forward( legacy_offsets = tensors["_legacy_offsets"] graph_batch = tensors["_graph_batch"] if "seq_len_q" in tensors: - variant_pack[tensors["seq_len_q"]] = _padded_sequence_lengths( - cu_seqlens_q, graph_batch - ) - variant_pack[tensors["seq_len_kv"]] = _padded_sequence_lengths( - cu_seqlens_kv, graph_batch - ) + variant_pack[tensors["seq_len_q"]] = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + variant_pack[tensors["seq_len_kv"]] = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) if "offset_q" in tensors: if legacy_offsets: variant_pack[tensors["offset_q"]] = _element_ragged_offsets( @@ -994,41 +943,29 @@ def _f16_forward( max_scores_for_reduce = max_scores if q_format == "thd": if max_scores.ndim == 4: - seqlens_q = _sequence_lengths(cu_seqlens_q).to( - device=max_scores.device + seqlens_q = _sequence_lengths(cu_seqlens_q).to(device=max_scores.device) + sq_idx = torch.arange(max_scores.shape[2], device=max_scores.device).view( + 1, 1, -1, 1 ) - sq_idx = torch.arange( - max_scores.shape[2], device=max_scores.device - ).view(1, 1, -1, 1) valid = sq_idx < seqlens_q.view(-1, 1, 1, 1) - max_scores_for_reduce = max_scores.masked_fill( - ~valid, float("-inf") - ) + max_scores_for_reduce = max_scores.masked_fill(~valid, float("-inf")) elif max_scores.ndim == 3: - seqlens_q = _sequence_lengths(cu_seqlens_q).to( - device=max_scores.device - ) + seqlens_q = _sequence_lengths(cu_seqlens_q).to(device=max_scores.device) total_tokens = max_scores.shape[0] starts = cu_seqlens_q_padded[:-1].to(device=max_scores.device) ends = (starts + seqlens_q).clamp(max=total_tokens) - delta = torch.zeros( - total_tokens + 1, dtype=torch.int32, device=max_scores.device - ) + delta = torch.zeros(total_tokens + 1, dtype=torch.int32, device=max_scores.device) updates = torch.ones_like(starts, dtype=torch.int32) delta.scatter_add_(0, starts.clamp(max=total_tokens), updates) delta.scatter_add_(0, ends, -updates) valid = delta[:-1].cumsum(0) > 0 - max_scores_for_reduce = max_scores.masked_fill( - ~valid.view(-1, 1, 1), float("-inf") - ) + max_scores_for_reduce = max_scores.masked_fill(~valid.view(-1, 1, 1), float("-inf")) reduce_dims = (0, 2) if max_scores_for_reduce.ndim == 3 else (0, 2, 3) max_logit = torch.amax(max_scores_for_reduce, dim=reduce_dims).to(output.dtype) return output, aux, max_logit -def _fp8_output_shape( - batch: int, seqlen: int, heads: int, dim: int, tensor_format: str -): +def _fp8_output_shape(batch: int, seqlen: int, heads: int, dim: int, tensor_format: str): if tensor_format == "bshd": return (batch, seqlen, heads, dim) if tensor_format == "sbhd": @@ -1063,14 +1000,8 @@ def _allocate_attention_grad_data( factory = torch.zeros if zero else torch.empty if len(components) == 1: - if ( - head_dim_qk != head_dim_v - or heads != kv_heads - or max_seqlen_q != max_seqlen_kv - ): - raise ValueError( - f"Packed QKV gradient layout {layout!r} requires matching Q/K/V." - ) + if head_dim_qk != head_dim_v or heads != kv_heads or max_seqlen_q != max_seqlen_kv: + raise ValueError(f"Packed QKV gradient layout {layout!r} requires matching Q/K/V.") packed_dim = components[0].index("3") packed_shape = list(q_shape) packed_shape.insert(packed_dim, 3) @@ -1078,9 +1009,7 @@ def _allocate_attention_grad_data( return tuple(packed.select(packed_dim, index) for index in range(3)) if len(components) == 2: if head_dim_qk != head_dim_v: - raise ValueError( - f"Packed KV gradient layout {layout!r} requires dQK == dV." - ) + raise ValueError(f"Packed KV gradient layout {layout!r} requires dQK == dV.") q_out = factory(q_shape, dtype=dtype, device=device) packed_dim = components[1].index("2") packed_shape = list(k_shape) @@ -1088,8 +1017,7 @@ def _allocate_attention_grad_data( packed = factory(packed_shape, dtype=dtype, device=device) return q_out, packed.select(packed_dim, 0), packed.select(packed_dim, 1) return tuple( - factory(shape, dtype=dtype, device=device) - for shape in (q_shape, k_shape, v_shape) + factory(shape, dtype=dtype, device=device) for shape in (q_shape, k_shape, v_shape) ) @@ -1328,17 +1256,13 @@ def _build_fp8_fwd_graph( output_t.set_output(True).set_dim((batch, heads, max_seqlen_q, d_v)).set_stride( _format_stride(batch, heads, max_seqlen_q, d_v, o_format) ) - output_dtype = ( - _fp8_cudnn_dtype(output) if _is_float8_tensor(output) else output.dtype - ) + output_dtype = _fp8_cudnn_dtype(output) if _is_float8_tensor(output) else output.dtype output_t.set_data_type(output_dtype) stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( (batch, heads, max_seqlen_q, 1) ).set_stride((heads * max_seqlen_q, max_seqlen_q, 1, 1)) tensors.update(O=output_t, Stats=stats_t) - return GraphEntry( - graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) - ) + return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) def _fp8_forward( @@ -1367,9 +1291,7 @@ def _fp8_forward( softmax_offset, ): if not isinstance(q, QuantizedTensorStorage): - raise TypeError( - "The FP8 cuDNN attention backend requires quantized Q/K/V tensors." - ) + raise TypeError("The FP8 cuDNN attention backend requires quantized Q/K/V tensors.") if _is_mxfp8_tensor(q) and "padding" in attn_mask_type: # cuDNN Frontend 1.27's Python sdpa_mxfp8 forward binding omits the # seq_len inputs that its C++ graph API exposes. Preserve functional @@ -1414,12 +1336,8 @@ def _fp8_forward( heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] d_v = v.shape[-1] output_shape = _fp8_output_shape(batch, max_seqlen_q, heads, d_v, o_format) - output, amax_o = _allocate_fp8_kernel_output( - o_quantizer, output_shape, fake_dtype, q.device - ) - stats = torch.empty( - (batch, heads, max_seqlen_q, 1), dtype=torch.float32, device=q.device - ) + output, amax_o = _allocate_fp8_kernel_output(o_quantizer, output_shape, fake_dtype, q.device) + stats = torch.empty((batch, heads, max_seqlen_q, 1), dtype=torch.float32, device=q.device) amax_s = ( s_quantizer.amax if isinstance(s_quantizer, Float8Quantizer) @@ -1429,9 +1347,7 @@ def _fp8_forward( else None ) ) - rng_elts = ( - max_seqlen_q * max_seqlen_q + _FP8_THREADS_PER_CTA - 1 - ) // _FP8_THREADS_PER_CTA + rng_elts = (max_seqlen_q * max_seqlen_q + _FP8_THREADS_PER_CTA - 1) // _FP8_THREADS_PER_CTA rng_state = _reserve_philox_state(q.device, rng_gen, rng_elts) q_data = _quantized_data(q) @@ -1461,9 +1377,7 @@ def _fp8_forward( entry = get_graph_entry(key) if entry is None: if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "cuDNN FP8 attention graph must be built before CUDA graph capture." - ) + raise RuntimeError("cuDNN FP8 attention graph must be built before CUDA graph capture.") entry = _build_fp8_fwd_graph( max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, @@ -1570,9 +1484,7 @@ def fused_attn_fwd( del cuda_graph backend = FusedAttnBackend.cast(fused_attention_backend) if backend == FusedAttnBackend.No_Backend: - raise ValueError( - "No cuDNN fused-attention backend supports this configuration." - ) + raise ValueError("No cuDNN fused-attention backend supports this configuration.") if attn_scale is None: attn_scale = 1.0 / math.sqrt(q.size(-1)) if bottom_right_diagonal is None: @@ -1586,9 +1498,7 @@ def fused_attn_fwd( if attn_bias_type != "no_bias" or attn_bias is not None: raise ValueError("FP8 fused attention does not support attention bias.") if return_max_logit: - raise ValueError( - "FP8 fused attention does not support returning maximum logits." - ) + raise ValueError("FP8 fused attention does not support returning maximum logits.") return _fp8_forward( is_training, max_seqlen_q, @@ -1684,34 +1594,25 @@ def _build_f16_bwd_graph( deterministic: bool, ) -> GraphEntry: cudnn = import_cudnn_frontend() - graph = make_graph( - torch_to_cudnn_dtype(q.dtype), q.device, name="te_fused_attention_bwd" - ) + graph = make_graph(torch_to_cudnn_dtype(q.dtype), q.device, name="te_fused_attention_bwd") q_format, kv_format = _q_kv_formats(qkv_layout) dq_format, dkv_format = _q_kv_formats(dqkv_layout) batch = cu_seqlens_q.numel() - 1 is_ragged_q = q_format == "thd" is_ragged_kv = kv_format == "thd" use_ragged_stats = is_ragged_q and cudnn.backend_version() >= 90600 - use_token_buckets = ( - cudnn.backend_version() >= 90600 - and torch.cuda.get_device_capability(q.device) != (12, 0) - ) + use_token_buckets = cudnn.backend_version() >= 90600 and torch.cuda.get_device_capability( + q.device + ) != (12, 0) use_legacy_offsets = is_ragged_q or is_ragged_kv - graph_batch = ( - _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch - ) + graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch if not use_token_buckets: use_ragged_stats = False graph_seqlen_q = ( - _max_ragged_tokens(q.shape[0]) - if is_ragged_q and use_token_buckets - else max_seqlen_q + _max_ragged_tokens(q.shape[0]) if is_ragged_q and use_token_buckets else max_seqlen_q ) graph_seqlen_kv = ( - _max_ragged_tokens(k.shape[0]) - if is_ragged_kv and use_token_buckets - else max_seqlen_kv + _max_ragged_tokens(k.shape[0]) if is_ragged_kv and use_token_buckets else max_seqlen_kv ) tensors: Dict[str, Any] = { @@ -1903,9 +1804,7 @@ def _build_f16_bwd_graph( if softmax_offset is None or d_softmax_offset is None: raise ValueError(f"softmax_type={softmax_type!r} requires sink tensors.") sink_t = graph.tensor_like(softmax_offset, name="softmax_offset") - dsink_t = graph.tensor_like( - d_softmax_offset, name="d_softmax_offset" - ).set_output(True) + dsink_t = graph.tensor_like(d_softmax_offset, name="d_softmax_offset").set_output(True) tensors.update(softmax_offset=sink_t, d_softmax_offset=dsink_t) options.update(sink_token=sink_t, dSink_token=dsink_t) @@ -1938,9 +1837,7 @@ def _build_f16_bwd_graph( dv_t.set_ragged_offset(offset_v).set_ragged_offset_multiplier(1) tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) - return GraphEntry( - graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) - ) + return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) def _build_fp8_bwd_graph( @@ -2087,17 +1984,13 @@ def _build_fp8_bwd_graph( options["dropout"] = (float(dropout), seed_t, offset_t) if softmax_type != "vanilla": sink_t = graph.tensor_like(softmax_offset, name="softmax_offset") - dsink_t = graph.tensor_like( - d_softmax_offset, name="d_softmax_offset" - ).set_output(True) + dsink_t = graph.tensor_like(d_softmax_offset, name="d_softmax_offset").set_output(True) tensors.update(softmax_offset=sink_t, d_softmax_offset=dsink_t) options.update(sink_token=sink_t, dSink_token=dsink_t) if _is_mxfp8_tensor(q): if d_o_f16 is None: - raise ValueError( - "MXFP8 attention backward requires the high-precision dO tensor." - ) + raise ValueError("MXFP8 attention backward requires the high-precision dO tensor.") q_col = _quantized_data(q, columnwise=True) k_col = _quantized_data(k, columnwise=True) do_col = _quantized_data(d_o, columnwise=True) @@ -2156,9 +2049,7 @@ def mx_scale(name, h, s, d, fmt): tensors[name] = tensor return tensor - descale_q = mx_scale( - "descale_q", heads, "s_q_padded", "d_qk_scale_padded", scale_format_q - ) + descale_q = mx_scale("descale_q", heads, "s_q_padded", "d_qk_scale_padded", scale_format_q) descale_q_t = mx_scale( "descale_q_t", heads, "s_q_scale_padded", "d_qk_padded", scale_format_q ) @@ -2215,9 +2106,7 @@ def mx_scale(name, h, s, d, fmt): "descale_o", "descale_do", ) - scalars = { - name: _scalar_graph_tensor(graph, cudnn, name) for name in scalar_names - } + scalars = {name: _scalar_graph_tensor(graph, cudnn, name) for name in scalar_names} tensors.update(scalars) delayed = isinstance(dqkv_quantizer, Float8Quantizer) if isinstance(s_quantizer, Float8Quantizer): @@ -2237,37 +2126,23 @@ def mx_scale(name, h, s, d, fmt): tensors[key] = _scalar_graph_tensor(graph, cudnn, name) descale_o_arg = tensors["descale_o"] descale_s_arg = ( - tensors["descale_s"] - if "descale_s" in tensors - else tensors["constant_descale_s"] + tensors["descale_s"] if "descale_s" in tensors else tensors["constant_descale_s"] ) descale_dp_arg = ( - tensors["descale_dp"] - if "descale_dp" in tensors - else tensors["constant_descale_dp"] - ) - scale_s_arg = ( - tensors["scale_s"] if "scale_s" in tensors else tensors["constant_scale_s"] + tensors["descale_dp"] if "descale_dp" in tensors else tensors["constant_descale_dp"] ) + scale_s_arg = tensors["scale_s"] if "scale_s" in tensors else tensors["constant_scale_s"] scale_dp_arg = ( - tensors["scale_dp"] - if "scale_dp" in tensors - else tensors["constant_scale_dp"] + tensors["scale_dp"] if "scale_dp" in tensors else tensors["constant_scale_dp"] ) scale_dq_arg = ( - tensors["scale_dq"] - if "scale_dq" in tensors - else tensors["constant_scale_dq"] + tensors["scale_dq"] if "scale_dq" in tensors else tensors["constant_scale_dq"] ) scale_dk_arg = ( - tensors["scale_dk"] - if "scale_dk" in tensors - else tensors["constant_scale_dk"] + tensors["scale_dk"] if "scale_dk" in tensors else tensors["constant_scale_dk"] ) scale_dv_arg = ( - tensors["scale_dv"] - if "scale_dv" in tensors - else tensors["constant_scale_dv"] + tensors["scale_dv"] if "scale_dv" in tensors else tensors["constant_scale_dv"] ) op = build_fp8_backward_operation( graph, @@ -2292,9 +2167,7 @@ def mx_scale(name, h, s, d, fmt): "scale_dp": scale_dp_arg, }, options, - FP8AttentionGraphConfig( - "delayed" if delayed else "current", "te_sdpa_fp8_backward" - ), + FP8AttentionGraphConfig("delayed" if delayed else "current", "te_sdpa_fp8_backward"), ) dq_t, dk_t, dv_t = op["dq"], op["dk"], op["dv"] amax_dq_t, amax_dk_t = op["amax_dq"], op["amax_dk"] @@ -2326,9 +2199,7 @@ def mx_scale(name, h, s, d, fmt): _format_stride(batch, kv_heads, max_seqlen_kv, d_value, dkv_format) ) tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) - return GraphEntry( - graph=graph, tensors=tensors, workspace_size=finalize_graph(graph) - ) + return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) def _fp8_backward( @@ -2399,9 +2270,7 @@ def _fp8_backward( batch = cu_seqlens_q.numel() - 1 heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] - output_dtype = ( - torch.uint8 if isinstance(dqkv_quantizer, Float8Quantizer) else fake_dtype - ) + output_dtype = torch.uint8 if isinstance(dqkv_quantizer, Float8Quantizer) else fake_dtype grad_data = _allocate_attention_grad_data( batch=batch, heads=heads, @@ -2422,12 +2291,8 @@ def _fp8_backward( stats, rng_state = aux_ctx_tensors[:2] softmax_offset = aux_ctx_tensors[2] if softmax_type != "vanilla" else None d_o_f16 = aux_ctx_tensors[-1] if _is_mxfp8_tensor(q) else None - d_softmax_offset = ( - torch.empty_like(softmax_offset) if softmax_offset is not None else None - ) - hidden_amax = [ - torch.zeros(1, dtype=torch.float32, device=q.device) for _ in range(4) - ] + d_softmax_offset = torch.empty_like(softmax_offset) if softmax_offset is not None else None + hidden_amax = [torch.zeros(1, dtype=torch.float32, device=q.device) for _ in range(4)] key = ( "fp8_bwd", @@ -2457,9 +2322,7 @@ def _fp8_backward( entry = get_graph_entry(key) if entry is None: if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "cuDNN FP8 attention graph must be built before CUDA graph capture." - ) + raise RuntimeError("cuDNN FP8 attention graph must be built before CUDA graph capture.") entry = _build_fp8_bwd_graph( max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, @@ -2561,9 +2424,7 @@ def _fp8_backward( if isinstance(dqkv_quantizer, Float8Quantizer) else hidden_amax ) - for name, value in zip( - ("amax_dq", "amax_dk", "amax_dv", "amax_dp"), amax_values - ): + for name, value in zip(("amax_dq", "amax_dk", "amax_dv", "amax_dp"), amax_values): variant_pack[t[name]] = value if "seq_len_q" in t: variant_pack[t["seq_len_q"]] = _sequence_lengths(cu_seqlens_q) @@ -2628,9 +2489,7 @@ def fused_attn_bwd( ) if backend == FusedAttnBackend.FP8: if attn_bias_type != "no_bias": - raise ValueError( - "FP8 fused attention backward does not support attention bias." - ) + raise ValueError("FP8 fused attention backward does not support attention bias.") return _fp8_backward( max_seqlen_q, max_seqlen_kv, @@ -2662,9 +2521,7 @@ def fused_attn_bwd( deterministic, ) if backend != FusedAttnBackend.F16_arbitrary_seqlen: - raise ValueError( - "No cuDNN fused-attention backend supports this backward configuration." - ) + raise ValueError("No cuDNN fused-attention backend supports this backward configuration.") stats = aux_ctx_tensors[0] rng_state = aux_ctx_tensors[1] @@ -2702,24 +2559,15 @@ def fused_attn_bwd( # cuDNN does not support the [1,1,1,S] reduction form. if not tuple(attn_bias.shape[:3]) == (1, 1, 1): d_bias = torch.empty_like(attn_bias) - d_softmax_offset = ( - torch.empty_like(softmax_offset) if softmax_type != "vanilla" else None - ) - cu_seqlens_q_padded = ( - cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded - ) - cu_seqlens_kv_padded = ( - cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded - ) + d_softmax_offset = torch.empty_like(softmax_offset) if softmax_type != "vanilla" else None + cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded + cu_seqlens_kv_padded = cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded cudnn = import_cudnn_frontend() - use_token_buckets = ( - cudnn.backend_version() >= 90600 - and torch.cuda.get_device_capability(q.device) != (12, 0) - ) + use_token_buckets = cudnn.backend_version() >= 90600 and torch.cuda.get_device_capability( + q.device + ) != (12, 0) use_legacy_offsets = q_format == "thd" or kv_format == "thd" - graph_batch = ( - _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch - ) + graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch key = ( "f16_bwd", @@ -2752,9 +2600,7 @@ def fused_attn_bwd( entry = get_graph_entry(key) if entry is None: if torch.cuda.is_current_stream_capturing(): - raise RuntimeError( - "cuDNN attention graph must be built before CUDA graph capture." - ) + raise RuntimeError("cuDNN attention graph must be built before CUDA graph capture.") entry = _build_f16_bwd_graph( max_seqlen_q=max_seqlen_q, max_seqlen_kv=max_seqlen_kv, @@ -2808,12 +2654,8 @@ def fused_attn_bwd( variant_pack[tensors["dBias"]] = d_bias graph_batch = tensors["_graph_batch"] if "seq_len_q" in tensors: - variant_pack[tensors["seq_len_q"]] = _padded_sequence_lengths( - cu_seqlens_q, graph_batch - ) - variant_pack[tensors["seq_len_kv"]] = _padded_sequence_lengths( - cu_seqlens_kv, graph_batch - ) + variant_pack[tensors["seq_len_q"]] = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + variant_pack[tensors["seq_len_kv"]] = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) if "offset_q" in tensors: variant_pack[tensors["offset_q"]] = _element_ragged_offsets( cu_seqlens_q_padded, graph_batch, q.stride(0) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 1e2bc101d6c..2f74d46e873 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -45,9 +45,7 @@ def _bhsd_dim_stride( (tensor.shape[0], tensor.shape[2], tensor.shape[1], tensor.shape[3]), (tensor.stride(0), tensor.stride(2), tensor.stride(1), tensor.stride(3)), ) - raise ValueError( - f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}." - ) + raise ValueError(f"Flex Attention only supports SBHD/BSHD tensor formats, got {tensor_format}.") def _bhsd_graph_tensor(graph, tensor: torch.Tensor, tensor_format: str): @@ -92,14 +90,10 @@ def _score_mod_tensor_dict_metadata( """Describe score_mod tensor parameters without including their values.""" if tensors is None: return () - return tuple( - (name, _score_mod_tensor_metadata(tensor)) for name, tensor in tensors.items() - ) + return tuple((name, _score_mod_tensor_metadata(tensor)) for name, tensor in tensors.items()) -def _score_mod_bhsd_tensor_metadata( - tensor: torch.Tensor, tensor_format: str -) -> Tuple[Any, ...]: +def _score_mod_bhsd_tensor_metadata(tensor: torch.Tensor, tensor_format: str) -> Tuple[Any, ...]: """Describe an SBHD/BSHD runtime tensor as a cuDNN BHSD graph tensor.""" dim, stride = _bhsd_dim_stride(tensor, tensor_format) return (dim, stride, tensor.dtype, _score_mod_device_key(tensor.device)) @@ -133,9 +127,7 @@ def _get_cudnn_current_stream_handle(cudnn, device: torch.device): def _build_cudnn_pygraph(dtype: torch.dtype, device: torch.device): """Create a cuDNN frontend Python graph for F16/BF16 SDPA.""" if dtype not in (torch.float16, torch.bfloat16): - raise ValueError( - f"Flex Attention only supports FP16/BF16 tensors, got {dtype}." - ) + raise ValueError(f"Flex Attention only supports FP16/BF16 tensors, got {dtype}.") return make_graph(torch_to_cudnn_dtype(dtype), device, name="te_flex_attention") @@ -256,10 +248,7 @@ def _cudnn_score_mod_bwd_cache_key( """Pre-build cache key for score_mod bprop execution plans.""" score_mod_key = _score_mod_callback_cache_key(score_mod) score_mod_bprop_key = _score_mod_callback_cache_key(score_mod_bprop) - if ( - score_mod_key is _SCORE_MOD_UNCACHEABLE - or score_mod_bprop_key is _SCORE_MOD_UNCACHEABLE - ): + if score_mod_key is _SCORE_MOD_UNCACHEABLE or score_mod_bprop_key is _SCORE_MOD_UNCACHEABLE: return None return ( "bwd", @@ -407,9 +396,7 @@ def _build_cudnn_score_mod_bwd_graph( else {} ) wrapped_score_mod = _wrap_score_mod(score_mod, score_mod_graph_tensors) - wrapped_score_mod_bprop = _wrap_score_mod( - score_mod_bprop, score_mod_bprop_graph_tensors - ) + wrapped_score_mod_bprop = _wrap_score_mod(score_mod_bprop, score_mod_bprop_graph_tensors) dq_layer = torch.empty_like(query_layer) dk_layer = torch.empty_like(key_layer) @@ -520,9 +507,7 @@ def forward( score_mod_tensors = dict(score_mod_tensors or {}) score_mod_bprop_tensors = dict(score_mod_bprop_tensors or {}) output_shape = (*query_layer.shape[:-1], value_layer.shape[-1]) - output_layer = torch.empty( - output_shape, device=query_layer.device, dtype=query_layer.dtype - ) + output_layer = torch.empty(output_shape, device=query_layer.device, dtype=query_layer.dtype) if is_training: stats = torch.empty( (*q_bhsd_dim[:-1], 1), @@ -594,8 +579,7 @@ def backward(ctx, d_out: torch.Tensor): # pylint: disable=missing-function-docstring if not ctx.is_training: raise RuntimeError( - "score_mod backward requires DotProductAttention to be in " - "training mode." + "score_mod backward requires DotProductAttention to be in training mode." ) saved_tensors = ctx.saved_tensors From 478ce7c707e40c3ea53027e4abaf498304bb5f2b Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 20:43:52 +0000 Subject: [PATCH 20/36] Remove obsolete fused attention backend override Signed-off-by: Vladimir Cherepanov --- docs/envvars.rst | 6 ------ docs/examples/attention/attention.ipynb | 8 +------- tests/pytorch/utils.py | 17 ++++------------- .../dot_product_attention.py | 17 ++++++++--------- 4 files changed, 13 insertions(+), 35 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 3ced6a8e873..f7b7f8727a6 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -190,12 +190,6 @@ backend-selection overview. :Default: ``1`` :Description: Enable or disable UnfusedDotProductAttention backend (native PyTorch). When set to ``0``, UnfusedDotProductAttention will not be used. -.. envvar:: NVTE_FUSED_ATTN_BACKEND - - :Type: ``int`` (1 or 2) - :Default: Auto-selected - :Description: Request a cuDNN FusedAttention backend when that request is supported by the active fused-attention path. ``1`` = F16_arbitrary_seqlen (cuDNN, any seq len), ``2`` = FP8 backend. If not set, the backend is automatically selected based on the input configuration. BF16/FP16 attention uses sub-backend ``1`` when eligible. FP8 attention uses sub-backend ``2`` when FP8 DPA is enabled and supported by the architecture, cuDNN version, and input configuration. - .. envvar:: NVTE_FUSED_ATTN_USE_FAv2_BWD :Type: ``int`` (0 or 1) diff --git a/docs/examples/attention/attention.ipynb b/docs/examples/attention/attention.ipynb index 989661b5434..4ffa804401f 100644 --- a/docs/examples/attention/attention.ipynb +++ b/docs/examples/attention/attention.ipynb @@ -346,17 +346,11 @@ "NVTE_FUSED_ATTN = 0 # disables cuDNN attention; default = 1\n", "```\n", "\n", - "**cuDNN attention sub-backends:**\n", - "This environment variable allows users to express their preference of cuDNN attention sub-backends. However, the elected sub-backend will only be used *if* it is eligible, i.e. if it has support for the provided inputs and runtime environment.\n", - "```\n", - "NVTE_FUSED_ATTN_BACKEND = 1/2 # user preference of cuDNN sub-backend\n", - "```\n", - "\n", "```\n", "
\n", "Note\n", " \n", - "Environment variables NVTE_FLASH_ATTN, NVTE_UNFUSED_ATTN, NVTE_FUSED_ATTN_BACKEND, and NVTE_FUSED_ATTN_USE_FAv2_BWD are supported in PyTorch. NVTE_FUSED_ATTN and NVTE_ALLOW_NONDETERMINISTIC_ALGO are supported in both PyTorch and JAX.\n", + "Environment variables NVTE_FLASH_ATTN, NVTE_UNFUSED_ATTN, and NVTE_FUSED_ATTN_USE_FAv2_BWD are supported in PyTorch. NVTE_FUSED_ATTN and NVTE_ALLOW_NONDETERMINISTIC_ALGO are supported in both PyTorch and JAX.\n", "
\n", "\n", "### 2.3 Example Tests\n", diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py index 0002bcef2ca..9ec93a3c603 100644 --- a/tests/pytorch/utils.py +++ b/tests/pytorch/utils.py @@ -27,7 +27,6 @@ AttentionLogging, check_set_window_size, ) -from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend from transformer_engine.pytorch.module.base import get_dummy_wgrad @@ -374,11 +373,6 @@ def get_available_attention_backends( if core_attention_bias_shape != "111s": core_attention_bias_requires_grad = True - fused_attn_backends = [] - available_backends = None - flash_attention_backend = None - fused_attention_backend = None - def test(): attention_params = AttentionParams( qkv_dtype=qkv_dtype, @@ -436,16 +430,13 @@ def test(): _attention_backends["backend_selection_requires_update"] = False return available_backends, flash_attention_backend, fused_attention_backend - backends = {1: "F16_arbitrary_seqlen", 2: "FP8"} if AttentionLogging._is_logging_setup is False: AttentionLogging.setup_logging() - for i in backends: - os.environ["NVTE_FUSED_ATTN_BACKEND"] = str(i) - _attention_backends["backend_selection_requires_update"] = True - available_backends, flash_attention_backend, fused_attention_backend = test() - if fused_attention_backend == FusedAttnBackend[backends[i]]: - fused_attn_backends.append(fused_attention_backend) + available_backends, flash_attention_backend, fused_attention_backend = test() + fused_attn_backends = [] + if fused_attention_backend is not None: + fused_attn_backends.append(fused_attention_backend) return available_backends, flash_attention_backend, fused_attn_backends diff --git a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py index 9c1d510bb5a..e59cb8fb619 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py @@ -1985,15 +1985,14 @@ def forward( .. note:: - Users can use environment variables :attr:`NVTE_FLASH_ATTN`, :attr:`NVTE_FUSED_ATTN`, - and :attr:`NVTE_FUSED_ATTN_BACKEND` to control which DotProductAttention backend, - and FusedAttention backend if applicable, to use. Transformer Engine first filters - backends by support for the runtime environment and input configuration, then applies - a performance-based preference order. On supported pre-Hopper GPUs, FlashAttention is - preferred over FusedAttention and UnfusedDotProductAttention when both optimized - backends are eligible. On Hopper and newer GPUs, including Blackwell, FusedAttention is - preferred over FlashAttention and UnfusedDotProductAttention when both optimized - backends are eligible. + Users can use environment variables :attr:`NVTE_FLASH_ATTN` and + :attr:`NVTE_FUSED_ATTN` to control which DotProductAttention backend to use. + Transformer Engine first filters backends by support for the runtime environment and + input configuration, then applies a performance-based preference order. On supported + pre-Hopper GPUs, FlashAttention is preferred over FusedAttention and + UnfusedDotProductAttention when both optimized backends are eligible. On Hopper and + newer GPUs, including Blackwell, FusedAttention is preferred over FlashAttention and + UnfusedDotProductAttention when both optimized backends are eligible. If FusedAttention is being used, users can also choose to switch to flash-attn's implementation for backward by setting :attr:`NVTE_FUSED_ATTN_USE_FAv2_BWD=1` (default: 0), because of the performance differences between various versions of From 150d5ed5f2852c76721fe841454224d346cb5d9f Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 20:48:26 +0000 Subject: [PATCH 21/36] Restore fused attention rejection diagnostics Signed-off-by: Vladimir Cherepanov --- docs/examples/attention/attention.ipynb | 6 +- tests/jax/test_fused_attn.py | 43 ++++++++++++- tests/pytorch/attention/test_attention.py | 52 +++++++++++++++- tests/pytorch/test_torch_compile.py | 2 +- .../jax/cpp_extensions/attention.py | 61 ++++++++++++++++--- .../jax/cpp_extensions/cudnn_attention.py | 14 +++-- .../dot_product_attention/_cudnn_backend.py | 12 ++-- .../attention/dot_product_attention/utils.py | 38 ++++++------ 8 files changed, 181 insertions(+), 47 deletions(-) diff --git a/docs/examples/attention/attention.ipynb b/docs/examples/attention/attention.ipynb index 4ffa804401f..4ca7f6554f5 100644 --- a/docs/examples/attention/attention.ipynb +++ b/docs/examples/attention/attention.ipynb @@ -249,11 +249,7 @@ "NVTE_DEBUG = 0/1 # disables/enables debugging\n", "NVTE_DEBUG_LEVEL = 0/1/2 # enables logging.WARNING/INFO/DEBUG-level messages\n", "```\n", - "
\n", - "Note:\n", - " \n", - "These flags are supported in PyTorch only as of Transformer Engine 2.0. JAX support is expected to be added in the future.\n", - "
" + "These flags are available in both PyTorch and JAX." ] }, { diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index a251ddbdb07..722eac2bb1f 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -4,7 +4,7 @@ """Tests for fused attention""" import os from enum import Enum, auto -from dataclasses import dataclass, field +from dataclasses import dataclass, field, replace from functools import partial from math import sqrt from typing import Any, Callable, Mapping, Tuple, Optional, Dict @@ -419,6 +419,45 @@ class Config: assert dimensions[4] == (2, 768, 8, 1) +def test_fused_attn_backend_message(monkeypatch): + """The JAX selector returns the shared policy's rejection reason.""" + monkeypatch.setattr(cudnn_attention, "get_cudnn_version", lambda: (9, 25, 0)) + monkeypatch.setattr(cudnn_attention, "_device_arch", lambda: 90) + baseline = FusedAttnHelper( + is_training=True, + q_dtype=jnp.bfloat16, + kv_dtype=jnp.bfloat16, + qkv_layout=QKVLayout.BSHD_BSHD_BSHD, + attn_bias_type=AttnBiasType.NO_BIAS, + attn_mask_type=AttnMaskType.NO_MASK, + softmax_type=AttnSoftmaxType.VANILLA_SOFTMAX, + dropout_probability=0.0, + q_num_heads=8, + kv_num_heads=8, + q_max_seqlen=128, + kv_max_seqlen=128, + head_dim_qk=64, + head_dim_v=64, + window_size=(-1, -1), + ) + + backend, message = baseline.get_fused_attn_backend() + assert backend == NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen + assert message == "" + + backend, message = replace( + baseline, attn_bias_type=AttnBiasType.PRE_SCALE_BIAS + ).get_fused_attn_backend() + assert backend == NVTE_Fused_Attn_Backend.NVTE_No_Backend + assert message == "attention bias is not supported" + + backend, message = replace( + baseline, head_dim_qk=1024, head_dim_v=1024 + ).get_fused_attn_backend() + assert backend == NVTE_Fused_Attn_Backend.NVTE_No_Backend + assert message == "head dimensions are not supported" + + class BiasShape(Enum): """ Enum class to represent the different bias shapes used in the fused attention. @@ -620,7 +659,7 @@ def _check_configs(self): "is either BSHD_BSHD_BSHD or THD_THD_THD" ) - self.backend = FusedAttnHelper( + self.backend, _ = FusedAttnHelper( self.is_training, self.dtype, self.dtype, diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index 6bdbcd023df..0812aeecf55 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -47,7 +47,8 @@ scaled_init_method_normal, ) from transformer_engine.pytorch.utils import get_cudnn_version -from transformer_engine.pytorch.constants import FP8BwdTensorIdx, FP8FwdTensorIdx +from transformer_engine.pytorch.constants import DType, FP8BwdTensorIdx, FP8FwdTensorIdx +from transformer_engine.pytorch.attention.dot_product_attention import _cudnn_backend import transformer_engine_torch as tex from transformer_engine.pytorch.quantized_tensor import ( Quantizer, @@ -129,6 +130,55 @@ def test_flash_attention_supported_version_message(): ) +def test_fused_attn_backend_message(monkeypatch): + """The PyTorch selector returns the shared policy's rejection reason.""" + monkeypatch.setattr(_cudnn_backend, "get_cudnn_version", lambda: (9, 25, 0)) + monkeypatch.setattr(_cudnn_backend, "get_device_compute_capability", lambda: (9, 0)) + baseline = dict( + is_training=True, + q_dtype=DType.kBFloat16, + kv_dtype=DType.kBFloat16, + qkv_layout="bshd_bshd_bshd", + bias_type="no_bias", + attn_mask_type="no_mask", + softmax_type="vanilla", + dropout=0.0, + num_attn_heads=8, + num_gqa_groups=8, + max_seqlen_q=128, + max_seqlen_kv=128, + head_dim_qk=64, + head_dim_v=64, + window_size_left=-1, + window_size_right=-1, + return_max_logit=False, + cuda_graph=False, + deterministic=False, + ) + + backend, message = _cudnn_backend.get_fused_attn_backend(**baseline) + assert backend == FusedAttnBackend.F16_arbitrary_seqlen + assert message == "" + + backend, message = _cudnn_backend.get_fused_attn_backend( + **{**baseline, "bias_type": "pre_scale_bias"} + ) + assert backend == FusedAttnBackend.No_Backend + assert message == "attention bias is not supported" + + backend, message = _cudnn_backend.get_fused_attn_backend( + **{**baseline, "head_dim_qk": 1024, "head_dim_v": 1024} + ) + assert backend == FusedAttnBackend.No_Backend + assert message == "head dimensions are not supported" + + backend, message = _cudnn_backend.get_fused_attn_backend( + **{**baseline, "q_dtype": DType.kFloat16} + ) + assert backend == FusedAttnBackend.No_Backend + assert message == "Q and KV must have the same data type" + + # Define F16 data types to test param_types = [torch.float16] if is_bf16_available(): diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index 2d1f7d0e697..e790f51b701 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -1356,7 +1356,7 @@ def fn(x, params): monkeypatch.setattr( dpa_utils, "get_cudnn_fused_attn_backend", - lambda *args: dpa_utils.FusedAttnBackend["No_Backend"], + lambda *args: (dpa_utils.FusedAttnBackend["No_Backend"], "disabled by test"), ) def fn_no_backend(x, params): diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index 905ef15cd06..4a6aca414d4 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -2,6 +2,7 @@ # # See LICENSE for license information. """JAX/TE custom ops for attention""" +import logging import os import warnings from dataclasses import dataclass, replace @@ -43,7 +44,7 @@ from .cudnn_attention import ( build_bwd_graph, build_fwd_graph, - is_fused_attn_supported, + get_fused_attn_support, ragged_graph_batch_size, ) from .misc import ( @@ -60,6 +61,35 @@ ] +# NVTE_DEBUG = 0/1 # disables/enables debug mode, default = 0 +_NVTE_DEBUG = int(os.getenv("NVTE_DEBUG", "0")) +# NVTE_DEBUG_LEVEL = 0/1/2 # enables increasingly verbose debug messages, default = 0 +_NVTE_DEBUG_LEVEL = int(os.getenv("NVTE_DEBUG_LEVEL", "0")) + + +class AttentionLogging: + """Logging for the JAX attention module.""" + + _log_level = _NVTE_DEBUG * _NVTE_DEBUG_LEVEL + _formatter = logging.Formatter("[%(levelname)-8s | %(name)-19s]: %(message)s") + _stream_handler = logging.StreamHandler() + logger = logging.getLogger(__name__) + _is_logging_setup = False + + @staticmethod + def setup_logging(): + """Set up log levels, logger, and handlers.""" + if AttentionLogging._is_logging_setup: + return + log_levels = {0: logging.WARNING, 1: logging.INFO, 2: logging.DEBUG} + level = AttentionLogging._log_level if AttentionLogging._log_level in log_levels else 2 + AttentionLogging._stream_handler.setFormatter(AttentionLogging._formatter) + AttentionLogging.logger.setLevel(log_levels[level]) + if not AttentionLogging.logger.hasHandlers(): + AttentionLogging.logger.addHandler(AttentionLogging._stream_handler) + AttentionLogging._is_logging_setup = True + + @partial( jax.tree_util.register_dataclass, data_fields=[], @@ -131,15 +161,28 @@ class FusedAttnHelper: def is_fused_attn_kernel_available(self): """Check if there is available fused attention kernel""" - return self.get_fused_attn_backend() != NVTE_Fused_Attn_Backend.NVTE_No_Backend + backend, _ = self.get_fused_attn_backend() + return backend != NVTE_Fused_Attn_Backend.NVTE_No_Backend def get_fused_attn_backend(self): - """Get the fused attention kernel backend""" - return ( + """Get the fused attention backend and a rejection reason when unavailable.""" + support = get_fused_attn_support(self) + backend = ( NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen - if is_fused_attn_supported(self) + if support.supported else NVTE_Fused_Attn_Backend.NVTE_No_Backend ) + message = "" if support.supported else support.reason + + AttentionLogging.setup_logging() + logger = AttentionLogging.logger + logger.debug("Running fused attention backend selection with config=%s", self) + if backend == NVTE_Fused_Attn_Backend.NVTE_No_Backend: + logger.info("No fused attention backend available; falling back to unfused attention.") + logger.debug("Reason fused attention was rejected: %s", message) + else: + logger.info("Selected fused attention backend: %s", backend) + return backend, message @staticmethod def is_non_deterministic_allowed(): @@ -386,7 +429,7 @@ def abstract( output_shape = (*batch_shape, q_max_seqlen, attn_heads, v_head_dim) out_aval = q_aval.update(shape=output_shape, dtype=q_dtype) - backend = FusedAttnHelper( + backend, message = FusedAttnHelper( config.is_training, q_dtype, k_dtype, @@ -405,7 +448,7 @@ def abstract( config.return_max_logit, ).get_fused_attn_backend() if backend != NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: - raise ValueError(f"Unsupported {backend=}") + raise ValueError(f"Unsupported {backend=}: {message}") graph_info = build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) softmax_dtype = dtypes.canonicalize_dtype(jnp.float32) @@ -925,7 +968,7 @@ def abstract( v_head_dim, ) = FusedAttnHelper.parse_qkv_aval(q_aval, k_aval, v_aval, config.qkv_layout) - backend = FusedAttnHelper( + backend, message = FusedAttnHelper( config.is_training, q_dtype, k_dtype, @@ -944,7 +987,7 @@ def abstract( False, ).get_fused_attn_backend() if backend != NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen: - raise ValueError(f"Unsupported {backend=}") + raise ValueError(f"Unsupported {backend=}: {message}") graph_info = build_bwd_graph( q_aval, diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index 1678bfe84b1..ba91656429b 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -17,6 +17,7 @@ from transformer_engine.common.attention.cudnn import ( AttentionLayout, FusedAttentionConfig, + FusedAttentionSupport, check_f16_fused_attention_support, cudnn_mask_options, ragged_batch_bucket, @@ -897,10 +898,10 @@ def _policy_mask_name(mask) -> str: }[mask.name] -def is_fused_attn_supported(helper) -> bool: - """Apply the shared F16/BF16 cuDNN attention compatibility policy.""" +def get_fused_attn_support(helper) -> FusedAttentionSupport: + """Return the shared F16/BF16 cuDNN attention compatibility result.""" - support = check_f16_fused_attention_support( + return check_f16_fused_attention_support( FusedAttentionConfig( is_training=bool(helper.is_training), q_dtype=str(jnp.dtype(helper.q_dtype)), @@ -924,4 +925,9 @@ def is_fused_attn_supported(helper) -> bool: sm_arch=_device_arch(), ) ) - return support.supported + + +def is_fused_attn_supported(helper) -> bool: + """Apply the shared F16/BF16 cuDNN attention compatibility policy.""" + + return get_fused_attn_support(helper).supported diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py index 8d8fffcc91a..d146752c125 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_backend.py @@ -54,16 +54,13 @@ def get_fused_attn_backend( cuda_graph, deterministic, ): - """Return the Python cuDNN SDPA backend value for an attention configuration.""" + """Return the Python cuDNN SDPA backend and rejection reason for a configuration.""" # Import lazily to avoid a circular import through dot_product_attention.utils. from .cudnn_attention import FusedAttnBackend q_dtype = DType.cast(q_dtype) kv_dtype = DType.cast(kv_dtype) - if q_dtype != kv_dtype: - raise ValueError("Q and KV must have the same data type") - major, minor = get_device_compute_capability() sm_arch = major * 10 + minor cudnn_version_tuple = get_cudnn_version() @@ -96,7 +93,8 @@ def get_fused_attn_backend( ) ) if support.supported: - return FusedAttnBackend.FP8 + return FusedAttnBackend.FP8, "" + return FusedAttnBackend.No_Backend, support.reason support = check_f16_fused_attention_support( FusedAttentionConfig( @@ -125,5 +123,5 @@ def get_fused_attn_backend( if support.warning is not None: warnings.warn(support.warning) if support.supported: - return FusedAttnBackend.F16_arbitrary_seqlen - return FusedAttnBackend.No_Backend + return FusedAttnBackend.F16_arbitrary_seqlen, "" + return FusedAttnBackend.No_Backend, support.reason diff --git a/transformer_engine/pytorch/attention/dot_product_attention/utils.py b/transformer_engine/pytorch/attention/dot_product_attention/utils.py index 468341d93ec..01a8d2a237e 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/utils.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/utils.py @@ -375,24 +375,23 @@ def _get_fused_attn_backend( keys and resolved to the pybind enums here, so that every argument is a python literal or a python enum. - Returns a plain int rather than a FusedAttnBackend member: dynamo - reconstructs the result of an assume_constant_result call by re-emitting the - call, which is only valid inside the frame that made it. An int survives a - graph break because it is baked into the graph as a literal, while an enum - member comes out of the reconstruction corrupted (see the cast at the call + Returns a plain int rather than a FusedAttnBackend member alongside the rejection + reason. Dynamo reconstructs the result of an assume_constant_result call by + re-emitting the call, which is only valid inside the frame that made it. An int + survives a graph break because it is baked into the graph as a literal, while an + enum member comes out of the reconstruction corrupted (see the cast at the call site, which restores the enum).""" - return int( - get_cudnn_fused_attn_backend( - is_training, - q_type, - kv_type, - qkv_layout, - bias_type, - attn_mask_type, - softmax_type, - *args, - ) + backend, reason = get_cudnn_fused_attn_backend( + is_training, + q_type, + kv_type, + qkv_layout, + bias_type, + attn_mask_type, + softmax_type, + *args, ) + return int(backend), reason def get_attention_backend( @@ -1505,7 +1504,7 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt # NOTE: under torch.compile the numeric args below must not be symbolic # (assume_constant_result requires concrete values); ints/floats made # dynamic by automatic dynamic currently graph break here. - fused_attention_backend = _get_fused_attn_backend( + fused_attention_backend, fused_attention_rejection_reason = _get_fused_attn_backend( is_training, q_type, kv_type, @@ -1527,7 +1526,10 @@ def _is_fa3_supported(num_heads, num_gqa_groups, head_dim_qk, head_dim_v, qkv_dt deterministic, ) if fused_attention_backend == FusedAttnBackend.No_Backend.value: - logger.debug("Disabling FusedAttention as no backend supports the provided input") + logger.debug( + "Disabling FusedAttention: %s", + fused_attention_rejection_reason or "no backend supports the provided input", + ) use_fused_attention = False fused_attention_backend = None elif ( From 4e5efd1ad33ac7d350d3328a5e08613893ce81f1 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 21:08:21 +0000 Subject: [PATCH 22/36] Restore fused attention graph cache diagnostics Signed-off-by: Vladimir Cherepanov --- docs/envvars.rst | 6 + docs/examples/attention/attention.ipynb | 13 + tests/test_common_attention_helpers.py | 82 ++++++ .../common/attention/cache_debug.py | 277 ++++++++++++++++++ transformer_engine/common/cudnn_frontend.py | 36 ++- .../jax/cpp_extensions/cudnn_attention.py | 36 ++- .../jax/cpp_extensions/cudnn_graph.py | 57 +++- .../jax/cpp_extensions/flex_attention.py | 43 ++- .../jax/cpp_extensions/fp8_attention.py | 30 +- .../jax/csrc/extensions/attention.cpp | 6 + .../csrc/extensions/attention_cache_debug.h | 248 ++++++++++++++++ .../jax/csrc/extensions/pybind.cpp | 10 + .../dot_product_attention/_cudnn_graph.py | 41 ++- .../dot_product_attention/cudnn_attention.py | 24 +- .../dot_product_attention/flex_attention.py | 19 +- 15 files changed, 888 insertions(+), 40 deletions(-) create mode 100644 transformer_engine/common/attention/cache_debug.py create mode 100644 transformer_engine/jax/csrc/extensions/attention_cache_debug.h diff --git a/docs/envvars.rst b/docs/envvars.rst index f7b7f8727a6..e6433f5733a 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -196,6 +196,12 @@ backend-selection overview. :Default: ``0`` :Description: When using FusedAttention, use FlashAttention-2 implementation for the backward pass instead of the cuDNN implementation. This can be useful due to performance differences between various versions of flash-attn and FusedAttention. +.. envvar:: NVTE_FUSED_ATTN_CACHE_DEBUG + + :Type: ``int`` (0, 1 or 2), optionally followed by ``:`` + :Default: ``0`` + :Description: Log FusedAttention Python graph-cache activity to stderr, prefixed with ``[FUSED-ATTN-CACHE]``. ``1`` prints an end-of-run summary of cache counters and the mean CPU wall time of each cuDNN graph-build stage. ``2`` additionally traces every cache event, including the cache key on hits and misses. When the launcher exports a rank, only rank 0 logs by default; append ``:`` to override this, for example ``1:all`` or ``2:0,3``. Supported by PyTorch and JAX. + .. envvar:: NVTE_ALLOW_NONDETERMINISTIC_ALGO :Type: ``int`` (0 or 1) diff --git a/docs/examples/attention/attention.ipynb b/docs/examples/attention/attention.ipynb index 4ca7f6554f5..c7d75daabd8 100644 --- a/docs/examples/attention/attention.ipynb +++ b/docs/examples/attention/attention.ipynb @@ -326,6 +326,19 @@ "!NVTE_DEBUG=1 NVTE_DEBUG_LEVEL=2 python example_attention.py" ] }, + { + "cell_type": "markdown", + "id": "fc0c9e3d", + "metadata": {}, + "source": [ + "For the FusedAttention backend, cuDNN Python graphs are cached and reused. `NVTE_FUSED_ATTN_CACHE_DEBUG` reports activity in these caches in both PyTorch and JAX.\n", + "```\n", + "NVTE_FUSED_ATTN_CACHE_DEBUG = 0 # disable cache diagnostics\n", + "NVTE_FUSED_ATTN_CACHE_DEBUG = 1/2:ranks # summary/trace level and optional ranks\n", + "```\n", + "Level 1 prints an end-of-run summary with graph creation, insertion, hit, miss, plan-build, and execution counters for each backend and pass. It also reports the mean CPU wall time of `validate`, `build_operation_graph`, `create_execution_plans`, `check_support`, and `build_plans`. Level 2 additionally prints each event as it occurs and includes the normalized cache key on hits and misses.\n" + ] + }, { "cell_type": "markdown", "id": "611d8fdb", diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py index db0912e60ba..af190580114 100644 --- a/tests/test_common_attention_helpers.py +++ b/tests/test_common_attention_helpers.py @@ -9,6 +9,7 @@ import pytest +from transformer_engine.common.attention import cache_debug from transformer_engine.common.attention.cudnn import ( AttentionLayout, FusedAttentionConfig, @@ -372,6 +373,87 @@ def test_shared_cudnn_graph_creation_and_finalization(): ] +def test_shared_cudnn_graph_reports_build_diagnostics(): + events = [] + graph = _FakeGraph() + build_cudnn_graph( + _FakeCudnn(), + graph, + description="attention", + debug_callback=lambda event, elapsed_ns: events.append((event, elapsed_ns)), + ) + assert [event for event, _ in events] == [ + "CREATE_GRAPH", + "validate", + "build_operation_graph", + "create_execution_plans", + "check_support", + "build_plans", + "BUILD_PLANS", + ] + assert all(elapsed_ns >= 0 for _, elapsed_ns in events) + + def test_shared_cudnn_graph_support_error_has_context(): with pytest.raises(RuntimeError, match="cuDNN test graph is not supported"): build_cudnn_graph(_FakeCudnn(), _FakeGraph(unsupported=True), description="test") + + +@pytest.fixture +def cache_debug_environment(monkeypatch): + for variable in ( + "NVTE_FUSED_ATTN_CACHE_DEBUG", + "RANK", + "LOCAL_RANK", + "OMPI_COMM_WORLD_RANK", + "SLURM_PROCID", + ): + monkeypatch.delenv(variable, raising=False) + cache_debug._reset_for_tests() + yield monkeypatch + cache_debug._reset_for_tests() + + +def test_cache_debug_rank_selection(cache_debug_environment): + monkeypatch = cache_debug_environment + monkeypatch.setenv("NVTE_FUSED_ATTN_CACHE_DEBUG", "2") + monkeypatch.setenv("RANK", "1") + cache_debug._reset_for_tests() + assert not cache_debug.enabled() + + monkeypatch.setenv("NVTE_FUSED_ATTN_CACHE_DEBUG", "2:1,3") + cache_debug._reset_for_tests() + assert cache_debug.enabled(trace=True) + + monkeypatch.setenv("NVTE_FUSED_ATTN_CACHE_DEBUG", "1:all") + cache_debug._reset_for_tests() + assert cache_debug.enabled() + assert not cache_debug.enabled(trace=True) + + +def test_cache_debug_summary(cache_debug_environment): + cache_debug_environment.setenv("NVTE_FUSED_ATTN_CACHE_DEBUG", "1") + cache_debug.record_lookup("f16", "fwd", hit=False, key=("graph", 1)) + cache_debug.record_event("f16", "fwd", "create_graph") + cache_debug.record_event("f16", "fwd", "cache_graph") + cache_debug.record_event("f16", "fwd", "execute", device=0) + cache_debug.record_lookup("f16", "fwd", hit=True, key=("graph", 1)) + cache_debug.record_build_time("f16", "fwd", "validate", 2_000_000) + + summary = cache_debug.render_summary() + assert "summary begin" in summary + assert "f16 fwd" in summary + assert "hit= 1" in summary + assert "miss= 1" in summary + assert "execute= 1" in summary + assert "validate" in summary + assert "2.000 ms/call" in summary + + +def test_cache_debug_level_two_traces_lookup_key(cache_debug_environment, capsys): + cache_debug_environment.setenv("NVTE_FUSED_ATTN_CACHE_DEBUG", "2") + cache_debug.record_lookup("fp8", "bwd", hit=False, key=("shape", 128)) + trace = capsys.readouterr().err + assert "[FUSED-ATTN-CACHE]" in trace + assert "fp8 bwd MISS" in trace + assert "('shape', 128)" in trace diff --git a/transformer_engine/common/attention/cache_debug.py b/transformer_engine/common/attention/cache_debug.py new file mode 100644 index 00000000000..f4f384b5221 --- /dev/null +++ b/transformer_engine/common/attention/cache_debug.py @@ -0,0 +1,277 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Fused-attention graph-cache diagnostics for Python graph runtimes.""" + +from __future__ import annotations + +import atexit +import os +import sys +import threading +from collections import defaultdict +from functools import lru_cache +from time import perf_counter_ns +from typing import Callable, Optional + +_PREFIX = "[FUSED-ATTN-CACHE]" +_RANK_ENV_VARS = ("RANK", "LOCAL_RANK", "OMPI_COMM_WORLD_RANK", "SLURM_PROCID") +_COUNTER_NAMES = ("hit", "miss", "create_graph", "cache_graph", "build_plans", "execute") +_BUILD_STAGES = ( + "validate", + "build_operation_graph", + "create_execution_plans", + "check_support", + "build_plans", +) + +_lock = threading.Lock() +_counters = defaultdict(lambda: defaultdict(int)) +_thread_counters = defaultdict(lambda: defaultdict(int)) +_thread_devices = defaultdict(set) +_stage_timings = defaultdict(lambda: [0, 0]) +_thread_ids: dict[int, int] = {} +_summary_registered = False + + +@lru_cache(maxsize=1) +def _configuration() -> tuple[int, bool, Optional[int]]: + """Return ``(level, selected, rank)`` for the current process.""" + + value = os.getenv("NVTE_FUSED_ATTN_CACHE_DEBUG", "0") + level_text, separator, rank_text = value.partition(":") + try: + level = int(level_text) + except ValueError: + level = 1 if level_text else 0 + if level <= 0: + return 0, False, None + + rank = None + for variable in _RANK_ENV_VARS: + rank_value = os.getenv(variable) + if rank_value: + try: + rank = int(rank_value) + except ValueError: + rank = 0 + break + if rank is None: + return level, True, None + if not separator: + return level, rank == 0, rank + if rank_text == "all": + return level, True, rank + selected_ranks = set() + for token in rank_text.split(","): + try: + selected_ranks.add(int(token)) + except ValueError: + continue + return level, rank in selected_ranks, rank + + +def enabled(*, trace: bool = False) -> bool: + """Whether cache diagnostics are enabled for this process and level.""" + + level, selected, _ = _configuration() + return selected and level >= (2 if trace else 1) + + +def _thread_id() -> int: + native_id = threading.get_ident() + with _lock: + thread_id = _thread_ids.get(native_id) + if thread_id is None: + thread_id = len(_thread_ids) + _thread_ids[native_id] = thread_id + return thread_id + + +def _rank_tag() -> str: + _, _, rank = _configuration() + return "" if rank is None else f"rank={rank} | " + + +def _write(text: str) -> None: + sys.stderr.write(text) + sys.stderr.flush() + + +def _counter_line( + thread_field: str, + device_field: str, + backend: str, + direction: str, + counters: dict[str, int], + event: str = "", +) -> str: + label = f"{backend} {direction}" + if event: + label += f" {event}" + values = ", ".join(f"{name}={counters.get(name, 0):4d}" for name in _COUNTER_NAMES) + return ( + f"{_PREFIX} {_rank_tag()}{thread_field:<7} {device_field:<9} | " + f"{label:<24} | {values}\n" + ) + + +def _register_summary() -> None: + global _summary_registered + with _lock: + if _summary_registered: + return + atexit.register(print_summary) + _summary_registered = True + + +def record_event( + backend: str, + direction: str, + event: str, + *, + device: Optional[int] = None, + key=None, +) -> None: + """Record one graph cache event and optionally emit its level-2 trace.""" + + if not enabled(): + return + event = event.lower() + if event not in _COUNTER_NAMES: + raise ValueError(f"Unknown fused-attention cache event {event!r}.") + _register_summary() + thread_id = _thread_id() + site = (backend, direction) + thread_site = (thread_id, backend, direction) + with _lock: + _counters[site][event] += 1 + _thread_counters[thread_site][event] += 1 + if device is not None: + _thread_devices[thread_id].add(device) + snapshot = dict(_counters[site]) + if not enabled(trace=True): + return + if event in ("hit", "miss") and key is not None: + _write( + f"{_PREFIX} {_rank_tag()}tid={thread_id:<3} dev={str(device):<3} | " + f"{backend} {direction} {event.upper():<12} | {key!r}\n" + ) + else: + _write( + _counter_line( + f"tid={thread_id}", + f"dev={device}", + backend, + direction, + snapshot, + event.upper(), + ) + ) + + +def record_lookup( + backend: str, + direction: str, + *, + hit: bool, + device: Optional[int] = None, + key=None, +) -> None: + """Record a cache hit or miss.""" + + record_event(backend, direction, "hit" if hit else "miss", device=device, key=key) + + +def record_build_time(backend: str, direction: str, stage: str, elapsed_ns: int) -> None: + """Accumulate CPU wall time for a cuDNN graph build stage.""" + + if not enabled(): + return + if stage not in _BUILD_STAGES: + raise ValueError(f"Unknown fused-attention graph build stage {stage!r}.") + _register_summary() + with _lock: + timing = _stage_timings[(backend, direction, stage)] + timing[0] += 1 + timing[1] += elapsed_ns + + +def build_recorder(backend: str, direction: str) -> Callable[[str, int], None]: + """Create the callback consumed by ``build_cudnn_graph``.""" + + def record(name: str, elapsed_ns: int) -> None: + if name in _BUILD_STAGES: + record_build_time(backend, direction, name, elapsed_ns) + else: + record_event(backend, direction, name.lower()) + + return record + + +def render_summary() -> str: + """Render the current diagnostic summary without writing it.""" + + if not enabled(): + return "" + with _lock: + counters = {site: dict(values) for site, values in _counters.items()} + thread_counters = {site: dict(values) for site, values in _thread_counters.items()} + thread_devices = {thread_id: set(values) for thread_id, values in _thread_devices.items()} + stage_timings = {site: tuple(values) for site, values in _stage_timings.items()} + + marker = f"{_PREFIX} {_rank_tag()}===== summary" + lines = [marker + " begin =====\n"] + for (thread_id, backend, direction), values in sorted(thread_counters.items()): + devices = thread_devices.get(thread_id, set()) + if not devices: + device_field = "dev=None" + elif len(devices) == 1: + device_field = f"dev={next(iter(devices))}" + else: + device_field = "dev=mixed" + lines.append( + _counter_line(f"tid={thread_id}", device_field, backend, direction, values) + ) + for (backend, direction), values in sorted(counters.items()): + lines.append(_counter_line("tid=all", "dev=all", backend, direction, values)) + for (backend, direction, stage), (calls, elapsed_ns) in sorted(stage_timings.items()): + lines.append( + f"{_PREFIX} {_rank_tag()}{backend:<3} {direction:<3} {stage:<22} | " + f"calls={calls} | time={elapsed_ns / calls / 1e6:9.3f} ms/call\n" + ) + lines.append(marker + " end =====\n") + return "".join(lines) + + +def print_summary() -> None: + """Write the current summary to stderr if diagnostics are enabled.""" + + summary = render_summary() + if summary: + _write(summary) + + +def time_call(callback: Optional[Callable[[str, int], None]], stage: str, function): + """Call a graph-build stage and report its elapsed CPU wall time.""" + + if callback is None: + return function() + start = perf_counter_ns() + try: + return function() + finally: + callback(stage, perf_counter_ns() - start) + + +def _reset_for_tests() -> None: + """Clear diagnostic state. Intended only for GPU-independent tests.""" + + with _lock: + _counters.clear() + _thread_counters.clear() + _thread_devices.clear() + _stage_timings.clear() + _thread_ids.clear() + _configuration.cache_clear() diff --git a/transformer_engine/common/cudnn_frontend.py b/transformer_engine/common/cudnn_frontend.py index c6a5ac5c5fc..634fe986588 100644 --- a/transformer_engine/common/cudnn_frontend.py +++ b/transformer_engine/common/cudnn_frontend.py @@ -7,7 +7,9 @@ from __future__ import annotations import importlib -from typing import Any +from typing import Any, Callable, Optional + +from transformer_engine.common.attention.cache_debug import time_call def import_cudnn_frontend(*, feature: str, requirement: str): @@ -42,15 +44,35 @@ def make_cudnn_graph( return cudnn.pygraph(**kwargs) -def build_cudnn_graph(cudnn, graph, *, description: str) -> int: +def build_cudnn_graph( + cudnn, + graph, + *, + description: str, + debug_callback: Optional[Callable[[str, int], None]] = None, +) -> int: """Validate and plan a graph, returning a nonzero workspace size.""" - graph.validate() - graph.build_operation_graph() + if debug_callback is not None: + debug_callback("CREATE_GRAPH", 0) + time_call(debug_callback, "validate", graph.validate) + time_call(debug_callback, "build_operation_graph", graph.build_operation_graph) try: - graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]) - graph.check_support() + time_call( + debug_callback, + "create_execution_plans", + lambda: graph.create_execution_plans( + [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] + ), + ) + time_call(debug_callback, "check_support", graph.check_support) except cudnn.cudnnGraphNotSupportedError as exc: raise RuntimeError(f"cuDNN {description} graph is not supported: {exc}") from exc - graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE) + time_call( + debug_callback, + "build_plans", + lambda: graph.build_plans(cudnn.build_plan_policy.HEURISTICS_CHOICE), + ) + if debug_callback is not None: + debug_callback("BUILD_PLANS", 0) return max(int(graph.get_workspace_size()), 1) diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index ba91656429b..2eb9168161c 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -32,6 +32,8 @@ finalize_graph, import_cudnn, make_graph, + record_cache_event, + record_cache_lookup, serialized_graph, ) from .misc import get_all_device_compute_capability, get_cudnn_version @@ -342,9 +344,13 @@ def _cache_key(direction: str, q_aval, k_aval, v_aval, bias_aval, config, *extra def build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGraphInfo: """Build or retrieve the standard fused-attention forward graph.""" key = _cache_key("fwd", q_aval, k_aval, v_aval, bias_aval, config) - if key not in _graph_cache: + graph_info = _graph_cache.get(key) + record_cache_lookup(("f16", "fwd"), hit=graph_info is not None, key=key) + if graph_info is None: _graph_cache[key] = _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) - return _graph_cache[key] + record_cache_event(("f16", "fwd"), "cache_graph") + graph_info = _graph_cache[key] + return graph_info def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGraphInfo: @@ -563,11 +569,18 @@ def _build_fwd_graph(q_aval, k_aval, v_aval, bias_aval, config) -> AttentionGrap else: stats.set_stride((info.q_heads * graph_sq, graph_sq, 1, 1)) - workspace, data, version = finalize_graph(cudnn, graph, description="fused-attention forward") + cache_site = ("f16", "fwd") + workspace, data, version = finalize_graph( + cudnn, + graph, + description="fused-attention forward", + cache_site=cache_site, + ) result = serialized_graph( serialized_graph_data=data, cudnn_frontend_version=version, workspace_size=workspace, + cache_site=cache_site, input_bindings=input_bindings, output_bindings=output_bindings, scalar_uids=scalar_uids, @@ -598,7 +611,9 @@ def build_bwd_graph( output_aval, doutput_aval, ) - if key not in _graph_cache: + graph_info = _graph_cache.get(key) + record_cache_lookup(("f16", "bwd"), hit=graph_info is not None, key=key) + if graph_info is None: _graph_cache[key] = _build_bwd_graph( q_aval, k_aval, @@ -609,7 +624,9 @@ def build_bwd_graph( doutput_aval, config, ) - return _graph_cache[key] + record_cache_event(("f16", "bwd"), "cache_graph") + graph_info = _graph_cache[key] + return graph_info def _build_bwd_graph( @@ -852,11 +869,18 @@ def io_tensor(name, dim, stride, uid, dtype=io_dtype): dk.set_ragged_offset(offset_k) dv.set_ragged_offset(offset_v) - workspace, data, version = finalize_graph(cudnn, graph, description="fused-attention backward") + cache_site = ("f16", "bwd") + workspace, data, version = finalize_graph( + cudnn, + graph, + description="fused-attention backward", + cache_site=cache_site, + ) result = serialized_graph( serialized_graph_data=data, cudnn_frontend_version=version, workspace_size=workspace, + cache_site=cache_site, input_bindings=input_bindings, output_bindings=output_bindings, scalar_uids=scalar_uids, diff --git a/transformer_engine/jax/cpp_extensions/cudnn_graph.py b/transformer_engine/jax/cpp_extensions/cudnn_graph.py index b5e3a6d4d62..15a2a01b31b 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_graph.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_graph.py @@ -22,6 +22,7 @@ import numpy as np import transformer_engine_jax +from transformer_engine.common.attention.cache_debug import enabled as cache_debug_enabled from transformer_engine.common.cudnn_frontend import ( build_cudnn_graph, import_cudnn_frontend, @@ -46,6 +47,8 @@ class SerializedGraph: graph_hash: tuple[int, int] cudnn_frontend_version: int workspace_size: int + attention_backend: str + attention_direction: str input_uids: np.ndarray input_buffer_indices: np.ndarray input_byte_offsets: np.ndarray @@ -63,6 +66,8 @@ def ffi_attrs(self) -> dict[str, Any]: "graph_hash0": self.graph_hash[0], "graph_hash1": self.graph_hash[1], "cudnn_frontend_version": self.cudnn_frontend_version, + "attention_backend": self.attention_backend, + "attention_direction": self.attention_direction, "input_uids": self.input_uids, "input_buffer_indices": self.input_buffer_indices, "input_byte_offsets": self.input_byte_offsets, @@ -218,6 +223,7 @@ def serialized_graph( serialized_graph_data: bytes, cudnn_frontend_version: int, workspace_size: int, + cache_site: tuple[str, str], input_bindings: Sequence[GraphBinding], output_bindings: Sequence[GraphBinding], scalar_uids: Sequence[int] = (), @@ -234,6 +240,8 @@ def binding_array(bindings, field): graph_hash=graph_hash(serialized_graph_data), cudnn_frontend_version=int(cudnn_frontend_version), workspace_size=max(int(workspace_size), 1), + attention_backend=cache_site[0], + attention_direction=cache_site[1], input_uids=binding_array(input_bindings, "uid"), input_buffer_indices=binding_array(input_bindings, "buffer_index"), input_byte_offsets=binding_array(input_bindings, "byte_offset"), @@ -246,9 +254,54 @@ def binding_array(bindings, field): ) -def finalize_graph(cudnn, graph, *, description: str) -> tuple[int, bytes, int]: +def record_cache_event( + cache_site: tuple[str, str], + event: str, + *, + key=None, + elapsed_ns: int = 0, +) -> None: + """Record an event in the JAX native cache-diagnostic state.""" + + if not cache_debug_enabled(): + return + native_event = "plans_built" if event == "BUILD_PLANS" else event.lower() + transformer_engine_jax.record_fused_attn_cache_event( + *cache_site, + native_event, + -1, + "" if key is None else repr(key), + elapsed_ns, + ) + + +def record_cache_lookup(cache_site: tuple[str, str], *, hit: bool, key=None) -> None: + """Record a JAX Python graph-cache lookup.""" + + record_cache_event(cache_site, "hit" if hit else "miss", key=key) + + +def _build_recorder(cache_site: tuple[str, str]): + def record(event: str, elapsed_ns: int) -> None: + record_cache_event(cache_site, event, elapsed_ns=elapsed_ns) + + return record + + +def finalize_graph( + cudnn, + graph, + *, + description: str, + cache_site: tuple[str, str], +) -> tuple[int, bytes, int]: """Validate, plan and serialize a cuDNN frontend graph.""" - workspace_size = build_cudnn_graph(cudnn, graph, description=description) + workspace_size = build_cudnn_graph( + cudnn, + graph, + description=description, + debug_callback=_build_recorder(cache_site), + ) return ( workspace_size, bytes(graph.serialize()), diff --git a/transformer_engine/jax/cpp_extensions/flex_attention.py b/transformer_engine/jax/cpp_extensions/flex_attention.py index de366c8b279..15ccde37b4f 100644 --- a/transformer_engine/jax/cpp_extensions/flex_attention.py +++ b/transformer_engine/jax/cpp_extensions/flex_attention.py @@ -24,6 +24,8 @@ finalize_graph, import_cudnn, make_graph, + record_cache_event, + record_cache_lookup, ) from .cudnn_graph import ( bshd_as_bhsd_dim_stride as _bshd_as_bhsd_dim_stride, @@ -363,11 +365,13 @@ def _serialized_score_mod_graph( output_uids: Sequence[int], scalar_uids: Sequence[int], scalar_values: Sequence[bytes], + cache_site: Tuple[str, str], ) -> _SerializedScoreModGraph: return make_serialized_graph( serialized_graph_data=serialized_graph, cudnn_frontend_version=int(cudnn_frontend_version), workspace_size=int(workspace_size), + cache_site=cache_site, input_bindings=[ GraphBinding(uid=int(uid), buffer_index=index) for index, uid in enumerate(input_uids) ], @@ -389,8 +393,15 @@ def wrapped_score_mod(sdpa_graph, score_tensor): return wrapped_score_mod -def _finalize_score_mod_graph(cudnn, graph) -> Tuple[int, bytes, int]: - return finalize_graph(cudnn, graph, description="score_mod SDPA") +def _finalize_score_mod_graph( + cudnn, graph, cache_site: Tuple[str, str] +) -> Tuple[int, bytes, int]: + return finalize_graph( + cudnn, + graph, + description="score_mod SDPA", + cache_site=cache_site, + ) def _graph_cache_key( @@ -468,11 +479,15 @@ def _build_score_mod_fwd_graph(q_aval, k_aval, v_aval, score_mod_avals, config): stats.set_data_type(cudnn.data_type.FLOAT) output_uids.append(_SCORE_MOD_UID_STATS) - workspace_size, serialized_graph, frontend_version = _finalize_score_mod_graph(cudnn, graph) + cache_site = ("f16", "fwd") + workspace_size, serialized_graph, frontend_version = _finalize_score_mod_graph( + cudnn, graph, cache_site + ) return _serialized_score_mod_graph( serialized_graph=serialized_graph, cudnn_frontend_version=frontend_version, workspace_size=workspace_size, + cache_site=cache_site, input_uids=[_SCORE_MOD_UID_Q, _SCORE_MOD_UID_K, _SCORE_MOD_UID_V, *tensor_uids], output_uids=output_uids, scalar_uids=scalar_uids, @@ -567,11 +582,15 @@ def _build_score_mod_bwd_graph( dk.set_output(True).set_uid(_SCORE_MOD_UID_DK).set_dim(k_dim).set_stride(k_stride) dv.set_output(True).set_uid(_SCORE_MOD_UID_DV).set_dim(v_dim).set_stride(v_stride) - workspace_size, serialized_graph, frontend_version = _finalize_score_mod_graph(cudnn, graph) + cache_site = ("f16", "bwd") + workspace_size, serialized_graph, frontend_version = _finalize_score_mod_graph( + cudnn, graph, cache_site + ) return _serialized_score_mod_graph( serialized_graph=serialized_graph, cudnn_frontend_version=frontend_version, workspace_size=workspace_size, + cache_site=cache_site, input_uids=[ _SCORE_MOD_UID_Q, _SCORE_MOD_UID_K, @@ -599,13 +618,17 @@ def _fused_attn_score_mod_fwd( score_mod_avals = tuple(_shape_dtype(arg) for arg in score_mod_tensors) key = _graph_cache_key("fwd", config, (q_aval, k_aval, v_aval, *score_mod_avals)) if key is None: + record_cache_lookup(("f16", "fwd"), hit=False, key="uncacheable score_mod") graph = _build_score_mod_fwd_graph(q_aval, k_aval, v_aval, score_mod_avals, config) else: - if key not in _score_mod_graph_cache: + graph = _score_mod_graph_cache.get(key) + record_cache_lookup(("f16", "fwd"), hit=graph is not None, key=key) + if graph is None: _score_mod_graph_cache[key] = _build_score_mod_fwd_graph( q_aval, k_aval, v_aval, score_mod_avals, config ) - graph = _score_mod_graph_cache[key] + record_cache_event(("f16", "fwd"), "cache_graph") + graph = _score_mod_graph_cache[key] batch, q_seqlen, q_heads, _ = q.shape _, _, _, v_head_dim = v.shape @@ -645,6 +668,7 @@ def _fused_attn_score_mod_bwd( avals = tuple(_shape_dtype(arg) for arg in all_inputs) key = _graph_cache_key("bwd", config, avals) if key is None: + record_cache_lookup(("f16", "bwd"), hit=False, key="uncacheable score_mod") graph = _build_score_mod_bwd_graph( *avals[:6], avals[6 : 6 + len(score_mod_tensors)], @@ -652,14 +676,17 @@ def _fused_attn_score_mod_bwd( config, ) else: - if key not in _score_mod_graph_cache: + graph = _score_mod_graph_cache.get(key) + record_cache_lookup(("f16", "bwd"), hit=graph is not None, key=key) + if graph is None: _score_mod_graph_cache[key] = _build_score_mod_bwd_graph( *avals[:6], avals[6 : 6 + len(score_mod_tensors)], avals[6 + len(score_mod_tensors) :], config, ) - graph = _score_mod_graph_cache[key] + record_cache_event(("f16", "bwd"), "cache_graph") + graph = _score_mod_graph_cache[key] dq = jax.ShapeDtypeStruct(q.shape, q.dtype) dk = jax.ShapeDtypeStruct(k.shape, k.dtype) diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index 669c7aafb29..778654f153c 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -63,6 +63,8 @@ finalize_graph, import_cudnn, make_graph, + record_cache_event, + record_cache_lookup, serialized_graph, ) from .misc import get_cudnn_version @@ -251,8 +253,10 @@ def build_fp8_fwd_graph( avals = (q_aval, k_aval, v_aval, q_scale_aval, k_scale_aval, v_scale_aval) key = _cache_key("fwd", mode, avals, config, output_dtype) - if key in _graph_cache: - return _graph_cache[key] + graph_info = _graph_cache.get(key) + record_cache_lookup(("fp8", "fwd"), hit=graph_info is not None, key=key) + if graph_info is not None: + return graph_info cudnn = import_cudnn() graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) @@ -401,18 +405,24 @@ def build_fp8_fwd_graph( (1, 1, 1, 1) ).set_stride((1, 1, 1, 1)) + cache_site = ("fp8", "fwd") workspace, data, version = finalize_graph( - cudnn, graph, description=f"JAX {mode} FP8 attention forward" + cudnn, + graph, + description=f"JAX {mode} FP8 attention forward", + cache_site=cache_site, ) result = serialized_graph( serialized_graph_data=data, cudnn_frontend_version=version, workspace_size=workspace, + cache_site=cache_site, input_bindings=input_bindings, output_bindings=output_bindings, ) info_result = FP8AttentionGraphInfo(result, o_shape, q_shape, k_shape, v_shape) _graph_cache[key] = info_result + record_cache_event(cache_site, "cache_graph") return info_result @@ -506,8 +516,10 @@ def build_fp8_bwd_graph( avals = (q_aval, k_aval, v_aval, stats_aval, output_aval, doutput_aval) key = _cache_key("bwd", mode, avals, config, grad_dtype) - if key in _graph_cache: - return _graph_cache[key] + graph_info = _graph_cache.get(key) + record_cache_lookup(("fp8", "bwd"), hit=graph_info is not None, key=key) + if graph_info is not None: + return graph_info cudnn = import_cudnn() graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) @@ -705,18 +717,24 @@ def build_fp8_bwd_graph( (1, 1, 1, 1) ).set_stride((1, 1, 1, 1)) + cache_site = ("fp8", "bwd") workspace, data, version = finalize_graph( - cudnn, graph, description=f"JAX {mode} FP8 attention backward" + cudnn, + graph, + description=f"JAX {mode} FP8 attention backward", + cache_site=cache_site, ) result = serialized_graph( serialized_graph_data=data, cudnn_frontend_version=version, workspace_size=workspace, + cache_site=cache_site, input_bindings=input_bindings, output_bindings=output_bindings, ) graph_info = FP8AttentionGraphInfo(result, o_shape, q_shape, k_shape, v_shape) _graph_cache[key] = graph_info + record_cache_event(cache_site, "cache_graph") return graph_info diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 11a49e0fe80..5d842a8d06c 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -18,6 +18,7 @@ #include #include "../extensions.h" +#include "attention_cache_debug.h" namespace transformer_engine { namespace jax { @@ -199,6 +200,11 @@ Error_Type ExecuteCudnnGraph(cudaStream_t stream, Dictionary &attrs, auto handle = GetCudnnHandle(); NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); + int device_id = 0; + NVTE_CHECK_CUDA(cudaGetDevice(&device_id)); + attention_cache_debug::Record( + get_attr_value(attrs, "attention_backend"), + get_attr_value(attrs, "attention_direction"), "execute", device_id); auto status = graph->execute(handle, variant_pack, workspace); NVTE_CHECK(status.is_good(), "cuDNN frontend graph execution failed: ", status.get_message()); return ffi_with_cuda_error_check(); diff --git a/transformer_engine/jax/csrc/extensions/attention_cache_debug.h b/transformer_engine/jax/csrc/extensions/attention_cache_debug.h new file mode 100644 index 00000000000..757592fd3ca --- /dev/null +++ b/transformer_engine/jax/csrc/extensions/attention_cache_debug.h @@ -0,0 +1,248 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +#ifndef TRANSFORMER_ENGINE_JAX_CSRC_EXTENSIONS_ATTENTION_CACHE_DEBUG_H_ +#define TRANSFORMER_ENGINE_JAX_CSRC_EXTENSIONS_ATTENTION_CACHE_DEBUG_H_ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace transformer_engine { +namespace jax { +namespace attention_cache_debug { +namespace detail { + +constexpr size_t kSiteCount = 4; +constexpr size_t kStageCount = 5; +inline constexpr std::array kStageNames = { + "validate", "build_operation_graph", "create_execution_plans", "check_support", + "build_plans"}; + +inline int DebugLevel() { + static const int level = [] { + const char *value = std::getenv("NVTE_FUSED_ATTN_CACHE_DEBUG"); + if (value == nullptr || value[0] == '\0' || value[0] == '0') return 0; + const int parsed = std::atoi(value); + return parsed > 0 ? parsed : 1; + }(); + return level; +} + +inline int LauncherRank() { + static const int rank = [] { + for (const char *name : {"RANK", "LOCAL_RANK", "OMPI_COMM_WORLD_RANK", "SLURM_PROCID"}) { + const char *value = std::getenv(name); + if (value != nullptr && value[0] != '\0') return std::atoi(value); + } + return -1; + }(); + return rank; +} + +inline bool Enabled() { + static const bool enabled = [] { + if (DebugLevel() < 1) return false; + const int rank = LauncherRank(); + if (rank < 0) return true; + const char *value = std::getenv("NVTE_FUSED_ATTN_CACHE_DEBUG"); + const char *separator = value == nullptr ? nullptr : std::strchr(value, ':'); + if (separator == nullptr) return rank == 0; + const std::string ranks(separator + 1); + if (ranks == "all") return true; + for (size_t start = 0; start <= ranks.size();) { + const size_t end = ranks.find(',', start); + const std::string token = ranks.substr(start, end - start); + if (!token.empty() && std::atoi(token.c_str()) == rank) return true; + if (end == std::string::npos) break; + start = end + 1; + } + return false; + }(); + return enabled; +} + +inline bool TraceEnabled() { return Enabled() && DebugLevel() >= 2; } + +inline size_t SiteIndex(std::string_view backend, std::string_view direction) { + const size_t backend_index = backend == "f16" ? 0 : backend == "fp8" ? 1 : kSiteCount; + const size_t direction_index = direction == "fwd" ? 0 : direction == "bwd" ? 1 : 2; + if (backend_index > 1 || direction_index > 1) { + throw std::invalid_argument("Invalid fused-attention cache diagnostic site: " + + std::string(backend) + " " + std::string(direction)); + } + return backend_index * 2 + direction_index; +} + +struct Counters { + std::atomic hit{0}; + std::atomic miss{0}; + std::atomic create_graph{0}; + std::atomic cache_graph{0}; + std::atomic build_plans{0}; + std::atomic execute{0}; +}; + +struct Timing { + std::atomic calls{0}; + std::atomic elapsed_ns{0}; +}; + +inline std::array &AllCounters() { + static auto *counters = new std::array(); + return *counters; +} + +inline std::array &AllTimings() { + static auto *timings = new std::array(); + return *timings; +} + +inline std::mutex &OutputMutex() { + static auto *mutex = new std::mutex(); + return *mutex; +} + +inline const char *Backend(size_t site) { return site < 2 ? "f16" : "fp8"; } +inline const char *Direction(size_t site) { return site % 2 == 0 ? "fwd" : "bwd"; } + +inline std::string RankTag() { + const int rank = LauncherRank(); + return rank < 0 ? "" : "rank=" + std::to_string(rank) + " | "; +} + +inline void Write(const std::string &message) { + std::lock_guard lock(OutputMutex()); + std::fwrite(message.data(), 1, message.size(), stderr); + std::fflush(stderr); +} + +inline uint64_t Load(const std::atomic &value) { + return value.load(std::memory_order_relaxed); +} + +inline std::string CounterLine(size_t site, const char *event = nullptr, int device = -1, + bool all_devices = false) { + const Counters &counters = AllCounters()[site]; + const std::string device_name = all_devices ? "all" : std::to_string(device); + char line[640]; + std::snprintf( + line, sizeof(line), + "[FUSED-ATTN-CACHE] %sdev=%-3s | %s %s %-12s | hit=%4" PRIu64 + ", miss=%4" PRIu64 ", create_graph=%4" PRIu64 ", cache_graph=%4" PRIu64 + ", build_plans=%4" PRIu64 ", execute=%4" PRIu64 "\n", + RankTag().c_str(), device_name.c_str(), Backend(site), Direction(site), + event == nullptr ? "" : event, + Load(counters.hit), Load(counters.miss), Load(counters.create_graph), + Load(counters.cache_graph), Load(counters.build_plans), Load(counters.execute)); + return line; +} + +inline void PrintSummary() { + if (!Enabled()) return; + const std::string marker = "[FUSED-ATTN-CACHE] " + RankTag() + "===== summary "; + std::string output = marker + "begin =====\n"; + for (size_t site = 0; site < kSiteCount; ++site) { + const Counters &counters = AllCounters()[site]; + if ((Load(counters.hit) | Load(counters.miss) | Load(counters.create_graph) | + Load(counters.cache_graph) | Load(counters.build_plans) | Load(counters.execute)) != 0) { + output += CounterLine(site, nullptr, -1, true); + } + } + for (size_t site = 0; site < kSiteCount; ++site) { + for (size_t stage = 0; stage < kStageCount; ++stage) { + const Timing &timing = AllTimings()[site * kStageCount + stage]; + const uint64_t calls = Load(timing.calls); + if (calls == 0) continue; + const double milliseconds = static_cast(Load(timing.elapsed_ns)) / calls / 1e6; + char line[320]; + std::snprintf(line, sizeof(line), + "[FUSED-ATTN-CACHE] %s%s %-3s %-22s | calls=%" PRIu64 + " | time=%9.3f ms/call\n", + RankTag().c_str(), Backend(site), Direction(site), kStageNames[stage], calls, + milliseconds); + output += line; + } + } + output += marker + "end =====\n"; + Write(output); +} + +inline void RegisterSummary() { + static const bool registered = [] { + std::atexit(PrintSummary); + return true; + }(); + (void)registered; +} + +inline std::atomic *EventCounter(Counters &counters, std::string_view event) { + if (event == "hit") return &counters.hit; + if (event == "miss") return &counters.miss; + if (event == "create_graph") return &counters.create_graph; + if (event == "cache_graph") return &counters.cache_graph; + if (event == "plans_built") return &counters.build_plans; + if (event == "execute") return &counters.execute; + return nullptr; +} + +inline size_t StageIndex(std::string_view event) { + for (size_t index = 0; index < kStageCount; ++index) { + if (event == kStageNames[index]) return index; + } + return kStageCount; +} + +} // namespace detail + +inline void Record(std::string_view backend, std::string_view direction, std::string_view event, + int device = -1, std::string_view key = {}, uint64_t elapsed_ns = 0) { + if (!detail::Enabled()) return; + detail::RegisterSummary(); + const size_t site = detail::SiteIndex(backend, direction); + const size_t stage = detail::StageIndex(event); + if (stage < detail::kStageCount) { + detail::Timing &timing = detail::AllTimings()[site * detail::kStageCount + stage]; + timing.calls.fetch_add(1, std::memory_order_relaxed); + timing.elapsed_ns.fetch_add(elapsed_ns, std::memory_order_relaxed); + return; + } + + detail::Counters &counters = detail::AllCounters()[site]; + std::atomic *counter = detail::EventCounter(counters, event); + if (counter == nullptr) { + throw std::invalid_argument("Invalid fused-attention cache diagnostic event: " + + std::string(event)); + } + counter->fetch_add(1, std::memory_order_relaxed); + if (!detail::TraceEnabled()) return; + if ((event == "hit" || event == "miss") && !key.empty()) { + detail::Write("[FUSED-ATTN-CACHE] " + detail::RankTag() + "dev=" + + std::to_string(device) + " | " + std::string(backend) + " " + + std::string(direction) + " " + std::string(event) + " | " + std::string(key) + + "\n"); + } else { + std::string uppercase(event == "plans_built" ? "build_plans" : event); + for (char &character : uppercase) { + if (character >= 'a' && character <= 'z') character -= 'a' - 'A'; + } + detail::Write(detail::CounterLine(site, uppercase.c_str(), device)); + } +} + +} // namespace attention_cache_debug +} // namespace jax +} // namespace transformer_engine + +#endif // TRANSFORMER_ENGINE_JAX_CSRC_EXTENSIONS_ATTENTION_CACHE_DEBUG_H_ diff --git a/transformer_engine/jax/csrc/extensions/pybind.cpp b/transformer_engine/jax/csrc/extensions/pybind.cpp index 6cf57543896..0ba8755207f 100644 --- a/transformer_engine/jax/csrc/extensions/pybind.cpp +++ b/transformer_engine/jax/csrc/extensions/pybind.cpp @@ -5,6 +5,7 @@ ************************************************************************/ #include "../extensions.h" +#include "attention_cache_debug.h" #include "cgemm_helper.h" #include "common/util/cuda_runtime.h" #include "transformer_engine/gemm.h" @@ -137,6 +138,15 @@ PYBIND11_MODULE(transformer_engine_jax, m) { m.def("get_cuda_version", &GetCudaRuntimeVersion); m.def("get_cudnn_version", &GetCudnnRuntimeVersion); m.def("get_cudnn_frontend_version", &GetCudnnFrontendVersion); + m.def( + "record_fused_attn_cache_event", + [](const std::string &backend, const std::string &direction, const std::string &event, + int device, const std::string &key, uint64_t elapsed_ns) { + attention_cache_debug::Record(backend, direction, event, device, key, elapsed_ns); + }, + pybind11::arg("backend"), pybind11::arg("direction"), pybind11::arg("event"), + pybind11::arg("device") = -1, pybind11::arg("key") = "", + pybind11::arg("elapsed_ns") = 0); m.def("get_device_compute_capability", &GetDeviceComputeCapability); m.def("get_num_compute_streams", &nvte_get_num_compute_streams); m.def("get_cublasLt_version", &cublasLtGetVersion); diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py index 7ad60182bb8..718493c9667 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py @@ -12,6 +12,11 @@ import torch +from transformer_engine.common.attention.cache_debug import ( + build_recorder, + record_event, + record_lookup, +) from transformer_engine.common.cudnn_frontend import ( build_cudnn_graph, make_cudnn_graph, @@ -100,11 +105,16 @@ def make_graph(io_dtype: Any, device: torch.device, *, name: str): ) -def finalize_graph(graph) -> int: +def finalize_graph(graph, *, cache_site: Tuple[str, str]) -> int: """Build a cuDNN graph and return its required workspace size.""" cudnn = import_cudnn_frontend() - return build_cudnn_graph(cudnn, graph, description="attention") + return build_cudnn_graph( + cudnn, + graph, + description="attention", + debug_callback=build_recorder(*cache_site), + ) @dataclass @@ -114,6 +124,7 @@ class GraphEntry: graph: Any tensors: Dict[str, Any] workspace_size: int + cache_site: Optional[Tuple[str, str]] = None _workspaces: Dict[int, torch.Tensor] = field(default_factory=dict, repr=False) def workspace(self, device: torch.device) -> torch.Tensor: @@ -143,6 +154,8 @@ def workspace(self, device: torch.device) -> torch.Tensor: def execute(self, variant_pack: Dict[Any, Any], device: torch.device) -> None: """Execute the graph on PyTorch's current stream.""" + if self.cache_site is not None: + record_event(*self.cache_site, "execute", device=_device_key(device)[1]) self.graph.execute( variant_pack, self.workspace(device), @@ -163,16 +176,38 @@ def graph_cache() -> Dict[Hashable, GraphEntry]: def get_graph_entry(key: Hashable) -> Optional[GraphEntry]: """Look up a graph in this thread's cache.""" - return graph_cache().get(key) + entry = graph_cache().get(key) + cache_site = _cache_site(key) + if cache_site is not None: + record_lookup(*cache_site, hit=entry is not None, key=key) + return entry def put_graph_entry(key: Hashable, entry: GraphEntry) -> GraphEntry: """Insert and return a graph cache entry.""" graph_cache()[key] = entry + cache_site = _cache_site(key) + if cache_site is not None: + entry.cache_site = cache_site + record_event(*cache_site, "cache_graph") return entry +def _cache_site(key: Hashable) -> Optional[Tuple[str, str]]: + """Extract a diagnostic build site from an attention graph cache key.""" + + if not isinstance(key, tuple) or not key or not isinstance(key[0], str): + return None + try: + backend, direction = key[0].split("_", 1) + except ValueError: + return None + if backend not in ("f16", "fp8") or direction not in ("fwd", "bwd"): + return None + return backend, direction + + def clear_graph_cache() -> None: """Clear thread-local handles and graphs. Intended for tests.""" diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 2f7ebb7d056..a65c0f9689c 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -747,7 +747,11 @@ def _build_f16_fwd_graph( stats_t.set_ragged_offset(offset_stats).set_ragged_offset_multiplier(stats_mult) tensors["Stats"] = stats_t - return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) + return GraphEntry( + graph=graph, + tensors=tensors, + workspace_size=finalize_graph(graph, cache_site=("f16", "fwd")), + ) def _f16_forward( @@ -1262,7 +1266,11 @@ def _build_fp8_fwd_graph( (batch, heads, max_seqlen_q, 1) ).set_stride((heads * max_seqlen_q, max_seqlen_q, 1, 1)) tensors.update(O=output_t, Stats=stats_t) - return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) + return GraphEntry( + graph=graph, + tensors=tensors, + workspace_size=finalize_graph(graph, cache_site=("fp8", "fwd")), + ) def _fp8_forward( @@ -1837,7 +1845,11 @@ def _build_f16_bwd_graph( dv_t.set_ragged_offset(offset_v).set_ragged_offset_multiplier(1) tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) - return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) + return GraphEntry( + graph=graph, + tensors=tensors, + workspace_size=finalize_graph(graph, cache_site=("f16", "bwd")), + ) def _build_fp8_bwd_graph( @@ -2199,7 +2211,11 @@ def mx_scale(name, h, s, d, fmt): _format_stride(batch, kv_heads, max_seqlen_kv, d_value, dkv_format) ) tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) - return GraphEntry(graph=graph, tensors=tensors, workspace_size=finalize_graph(graph)) + return GraphEntry( + graph=graph, + tensors=tensors, + workspace_size=finalize_graph(graph, cache_site=("fp8", "bwd")), + ) def _fp8_backward( diff --git a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py index 2f74d46e873..4cd4adf225a 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/flex_attention.py @@ -9,6 +9,7 @@ import torch +from transformer_engine.common.attention.cache_debug import record_event, record_lookup from transformer_engine.common.attention.score_mod import ( UNCACHEABLE_SCORE_MOD, score_mod_callback_cache_key, @@ -164,9 +165,9 @@ class _CudnnScoreModBwdGraphEntry: workspace_size: int -def _finalize_cudnn_graph(graph) -> int: +def _finalize_cudnn_graph(graph, cache_site: Tuple[str, str]) -> int: """Compatibility wrapper around shared graph finalization.""" - return finalize_graph(graph) + return finalize_graph(graph, cache_site=cache_site) def _execute_cudnn_graph( @@ -174,6 +175,7 @@ def _execute_cudnn_graph( variant_pack: Dict[Any, torch.Tensor], workspace_size: int, device: torch.device, + cache_site: Tuple[str, str], ): """Execute a built cuDNN frontend Python graph.""" cudnn = _import_cudnn_frontend() @@ -185,6 +187,7 @@ def _execute_cudnn_graph( device=device, dtype=torch.uint8, ) + record_event(*cache_site, "execute", device=_score_mod_device_key(device)[1]) graph.execute( variant_pack, workspace, @@ -314,7 +317,7 @@ def _build_cudnn_score_mod_fwd_graph( else: stats_tensor = None - workspace_size = _finalize_cudnn_graph(graph) + workspace_size = _finalize_cudnn_graph(graph, ("f16", "fwd")) return _CudnnScoreModFwdGraphEntry( graph=graph, q=q, @@ -356,11 +359,14 @@ def _get_cudnn_score_mod_fwd_graph( ) key = _cudnn_score_mod_fwd_cache_key(*build_args) if key is None: + record_lookup("f16", "fwd", hit=False, key="uncacheable score_mod") return _build_cudnn_score_mod_fwd_graph(*build_args) entry = _cudnn_score_mod_graph_cache.get(key) + record_lookup("f16", "fwd", hit=entry is not None, key=key) if entry is None: entry = _build_cudnn_score_mod_fwd_graph(*build_args) _cudnn_score_mod_graph_cache[key] = entry + record_event("f16", "fwd", "cache_graph") return entry @@ -422,7 +428,7 @@ def _build_cudnn_score_mod_bwd_graph( dk.set_output(True).set_dim(dk_dim).set_stride(dk_stride) dv.set_output(True).set_dim(dv_dim).set_stride(dv_stride) - workspace_size = _finalize_cudnn_graph(graph) + workspace_size = _finalize_cudnn_graph(graph, ("f16", "bwd")) return _CudnnScoreModBwdGraphEntry( graph=graph, q=q, @@ -475,11 +481,14 @@ def _get_cudnn_score_mod_bwd_graph( ) key = _cudnn_score_mod_bwd_cache_key(*build_args) if key is None: + record_lookup("f16", "bwd", hit=False, key="uncacheable score_mod") return _build_cudnn_score_mod_bwd_graph(*build_args) entry = _cudnn_score_mod_graph_cache.get(key) + record_lookup("f16", "bwd", hit=entry is not None, key=key) if entry is None: entry = _build_cudnn_score_mod_bwd_graph(*build_args) _cudnn_score_mod_graph_cache[key] = entry + record_event("f16", "bwd", "cache_graph") return entry @@ -546,6 +555,7 @@ def forward( variant_pack, entry.workspace_size, query_layer.device, + ("f16", "fwd"), ) ctx.is_training = is_training @@ -633,6 +643,7 @@ def backward(ctx, d_out: torch.Tensor): variant_pack, entry.workspace_size, query_layer.device, + ("f16", "bwd"), ) return ( From 984ddbdf8ab9b95d8465c19ace8758ac8bd0af96 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 23:07:39 +0000 Subject: [PATCH 23/36] Restore PyTorch FP8 THD attention support Signed-off-by: Vladimir Cherepanov --- tests/test_common_attention_helpers.py | 46 ++ transformer_engine/common/attention/cudnn.py | 26 +- .../dot_product_attention/cudnn_attention.py | 485 ++++++++++++++++-- 3 files changed, 497 insertions(+), 60 deletions(-) diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py index af190580114..0906cf910e0 100644 --- a/tests/test_common_attention_helpers.py +++ b/tests/test_common_attention_helpers.py @@ -135,6 +135,52 @@ def test_shared_fp8_policy(): assert not check_fp8_fused_attention_support(replace(fp8, head_dim_qk=200)).supported +def test_shared_fp8_thd_policy(): + thd = _attention_config( + q_dtype="float8_e4m3", + kv_dtype="float8_e4m3", + layout=AttentionLayout("thd", "thd", "thd", "separate"), + mask_type="padding", + cudnn_version=(9, 23, 0), + sm_arch=100, + ) + assert check_fp8_fused_attention_support(thd).supported + assert check_fp8_fused_attention_support(replace(thd, sm_arch=90)).supported + assert check_fp8_fused_attention_support( + replace(thd, mask_type="padding_causal_bottom_right") + ).supported + assert not check_fp8_fused_attention_support( + replace(thd, cudnn_version=(9, 22, 9)) + ).supported + assert not check_fp8_fused_attention_support(replace(thd, mask_type="no_mask")).supported + assert not check_fp8_fused_attention_support( + replace(thd, is_training=True, sm_arch=90) + ).supported + assert not check_fp8_fused_attention_support(replace(thd, head_dim_qk=144)).supported + + sink_backward = replace(thd, is_training=True, softmax_type="learnable") + assert not check_fp8_fused_attention_support( + replace(sink_backward, cudnn_version=(9, 25, 1)) + ).supported + assert check_fp8_fused_attention_support( + replace(sink_backward, cudnn_version=(9, 26, 0)) + ).supported + + +def test_shared_fp8_thd_policy_allows_64bit_offsets(): + thd = _attention_config( + q_dtype="float8_e4m3", + kv_dtype="float8_e4m3", + layout=AttentionLayout("thd", "thd", "thd", "separate"), + mask_type="padding", + max_seqlen_q=1_048_577, + max_seqlen_kv=1_048_577, + cudnn_version=(9, 23, 0), + sm_arch=100, + ) + assert check_fp8_fused_attention_support(thd).supported + + @pytest.mark.parametrize( "layout, expected", [ diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index 9d0b19e74dd..3f2bd3efc4c 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -606,8 +606,21 @@ def check_fp8_fused_attention_support( skv, dqk, dv, - ): - return _unsupported("FP8 attention does not support 64-bit ragged offsets") + ) and version < 90500: + return _unsupported("FP8 attention requires cuDNN 9.5 for 64-bit ragged offsets") + + is_thd = layout.qkv_format == "thd" + if is_thd: + if version < 92300: + return _unsupported("FP8 THD attention requires cuDNN 9.23 or newer") + if mask not in ("padding", "padding_causal", "padding_causal_bottom_right"): + return _unsupported("FP8 THD attention requires a padding mask") + if config.is_training and arch < 100: + return _unsupported("FP8 THD attention backward requires SM100 or newer") + if config.is_training and config.softmax_type != "vanilla" and version < 92600: + return _unsupported("FP8 THD sink-token backward requires cuDNN 9.26 or newer") + if arch >= 100 and (dqk > 128 or dv > 128): + return _unsupported("FP8 THD attention supports head dimensions up to 128 on SM100+") shape_mask_ok = ( ( @@ -628,7 +641,10 @@ def check_fp8_fused_attention_support( ) and dqk % 16 == 0 and dv % 16 == 0 - and mask in ("no_mask", "causal", "padding", "padding_causal") + and ( + mask in ("no_mask", "causal", "padding", "padding_causal") + or (arch >= 100 and mask == "padding_causal_bottom_right") + ) ) or ( version >= 92100 @@ -647,7 +663,9 @@ def check_fp8_fused_attention_support( version < 92100 and layout.qkv_format in ("bshd", "sbhd") and config.softmax_type == "vanilla" - ) or (version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd")) + ) or ( + version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd") + ) or is_thd if not format_softmax_ok: return _unsupported("FP8 attention layout or softmax type is not supported") return FusedAttentionSupport(True) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index a65c0f9689c..45eeb7f5427 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -239,6 +239,19 @@ def _q_kv_formats(qkv_layout: str) -> Tuple[str, str]: return q_format, kv_format +def _kv_ragged_offsets_source( + qkv_layout: str, + cu_seqlens_q_padded: torch.Tensor, + cu_seqlens_kv_padded: torch.Tensor, +) -> torch.Tensor: + """Choose the token offsets backing K/V for a ragged storage layout.""" + + layout = qkv_layout.removeprefix("paged_kv_") + if len(layout.split("_")) == 1 and "3" in layout: + return cu_seqlens_q_padded + return cu_seqlens_kv_padded + + def _is_paged_layout(qkv_layout: str) -> bool: return qkv_layout.startswith("paged_kv_") @@ -969,13 +982,25 @@ def _f16_forward( return output, aux, max_logit -def _fp8_output_shape(batch: int, seqlen: int, heads: int, dim: int, tensor_format: str): +def _fp8_output_shape( + batch: int, + seqlen: int, + heads: int, + dim: int, + tensor_format: str, + *, + total_tokens: Optional[int] = None, +): if tensor_format == "bshd": return (batch, seqlen, heads, dim) if tensor_format == "sbhd": return (seqlen, batch, heads, dim) if tensor_format == "bhsd": return (batch, heads, seqlen, dim) + if tensor_format == "thd": + if total_tokens is None: + raise ValueError("FP8 THD output allocation requires the total token count.") + return (total_tokens, heads, dim) raise ValueError(f"FP8 attention does not support output format {tensor_format!r}.") @@ -992,15 +1017,38 @@ def _allocate_attention_grad_data( dtype: torch.dtype, device: torch.device, zero: bool, + total_tokens_q: Optional[int] = None, + total_tokens_kv: Optional[int] = None, ): """Allocate dQ/dK/dV buffers with the requested packed storage layout.""" layout = dqkv_layout.removeprefix("paged_kv_") components = layout.split("_") q_format, kv_format = _q_kv_formats(layout) - q_shape = _fp8_output_shape(batch, max_seqlen_q, heads, head_dim_qk, q_format) - k_shape = _fp8_output_shape(batch, max_seqlen_kv, kv_heads, head_dim_qk, kv_format) - v_shape = _fp8_output_shape(batch, max_seqlen_kv, kv_heads, head_dim_v, kv_format) + q_shape = _fp8_output_shape( + batch, + max_seqlen_q, + heads, + head_dim_qk, + q_format, + total_tokens=total_tokens_q, + ) + k_shape = _fp8_output_shape( + batch, + max_seqlen_kv, + kv_heads, + head_dim_qk, + kv_format, + total_tokens=total_tokens_kv, + ) + v_shape = _fp8_output_shape( + batch, + max_seqlen_kv, + kv_heads, + head_dim_v, + kv_format, + total_tokens=total_tokens_kv, + ) factory = torch.zeros if zero else torch.empty if len(components) == 1: @@ -1061,45 +1109,117 @@ def _build_fp8_fwd_graph( softmax_offset, cu_seqlens_q, cu_seqlens_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, ): cudnn = import_cudnn_frontend() q_data = _quantized_data(q) k_data = _quantized_data(k) v_data = _quantized_data(v, columnwise=_is_mxfp8_tensor(v)) + output_data = _quantized_data(output) if _is_float8_tensor(output) else output q_format, kv_format = _q_kv_formats(qkv_layout) batch = cu_seqlens_q.numel() - 1 - heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] - kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] - d_qk = q.shape[-1] - d_v = v.shape[-1] + heads = q_data.shape[-2] if q_format != "bhsd" else q_data.shape[1] + kv_heads = k_data.shape[-2] if kv_format != "bhsd" else k_data.shape[1] + d_qk = q_data.shape[-1] + d_v = v_data.shape[-1] + is_ragged_q = q_format == "thd" + is_ragged_kv = kv_format == "thd" + use_token_buckets = cudnn.backend_version() >= 90600 and torch.cuda.get_device_capability( + q.device + ) != (12, 0) + use_ragged_stats = is_ragged_q and use_token_buckets + use_legacy_offsets = is_ragged_q or is_ragged_kv + graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch + graph_seqlen_q = ( + _max_ragged_tokens(q_data.shape[0]) + if is_ragged_q and use_token_buckets + else max_seqlen_q + ) + graph_seqlen_kv = ( + _max_ragged_tokens(k_data.shape[0]) + if is_ragged_kv and use_token_buckets + else max_seqlen_kv + ) graph = make_graph(_fp8_cudnn_dtype(q), q.device, name="te_fp8_sdpa_fwd") - tensors: Dict[str, Any] = {} + tensors: Dict[str, Any] = { + "_legacy_offsets": use_legacy_offsets, + "_graph_batch": graph_batch, + } + + offset_q = offset_o = offset_k = offset_v = offset_stats = None + if is_ragged_q: + offset_q, q_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_q", + length=graph_batch + 1, + data_type=torch.int64, + ) + offset_o, o_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_o", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors.update(offset_q=offset_q, offset_o=offset_o) + else: + q_mult = o_mult = 1 + if is_ragged_kv: + offset_k, k_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1, + name="offset_k", + length=graph_batch + 1, + data_type=torch.int64, + ) + offset_v, v_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1, + name="offset_v", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors.update(offset_k=offset_k, offset_v=offset_v) + else: + k_mult = v_mult = 1 q_t = _make_bhsd_graph_tensor( graph, q_data, q_format, - batch=batch, - max_seqlen=max_seqlen_q, + batch=graph_batch, + max_seqlen=graph_seqlen_q, data_type=_fp8_cudnn_dtype(q), + ragged_offset=offset_q, + ragged_offset_multiplier=q_mult, name="Q", ) k_t = _make_bhsd_graph_tensor( graph, k_data, kv_format, - batch=batch, - max_seqlen=max_seqlen_kv, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, data_type=_fp8_cudnn_dtype(k), + ragged_offset=offset_k, + ragged_offset_multiplier=k_mult, name="K", ) v_t = _make_bhsd_graph_tensor( graph, v_data, kv_format, - batch=batch, - max_seqlen=max_seqlen_kv, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, data_type=_fp8_cudnn_dtype(v), + ragged_offset=offset_v, + ragged_offset_multiplier=v_mult, name="V", ) tensors.update(Q=q_t, K=k_t, V=v_t) @@ -1119,8 +1239,8 @@ def _build_fp8_fwd_graph( use_padding_mask=is_padding, ) if is_padding: - seq_q = _sequence_lengths(cu_seqlens_q) - seq_kv = _sequence_lengths(cu_seqlens_kv) + seq_q = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + seq_kv = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) seq_q_t = graph.tensor_like(seq_q, name="seq_len_q") seq_kv_t = graph.tensor_like(seq_kv, name="seq_len_kv") tensors.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) @@ -1257,14 +1377,41 @@ def _build_fp8_fwd_graph( ).set_stride((1, 1, 1, 1)) tensors.update(amax_s=amax_s_t, amax_o=amax_o_t) - output_t.set_output(True).set_dim((batch, heads, max_seqlen_q, d_v)).set_stride( - _format_stride(batch, heads, max_seqlen_q, d_v, o_format) + output_dim, output_stride = _logical_bhsd_desc( + output_data, + o_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, ) + output_t.set_output(True).set_dim(output_dim).set_stride(output_stride) + if is_ragged_q: + output_t.set_ragged_offset(offset_o).set_ragged_offset_multiplier(o_mult) output_dtype = _fp8_cudnn_dtype(output) if _is_float8_tensor(output) else output.dtype output_t.set_data_type(output_dtype) - stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim( - (batch, heads, max_seqlen_q, 1) - ).set_stride((heads * max_seqlen_q, max_seqlen_q, 1, 1)) + _, stats_dim, stats_stride = _stats_layout( + batch=graph_batch, + heads=heads, + max_seqlen_q=graph_seqlen_q, + total_tokens_q=q_data.shape[0] if is_ragged_q else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + if use_ragged_stats: + offset_stats, stats_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_stats", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors["offset_stats"] = offset_stats + else: + stats_mult = 1 + stats_t.set_output(True).set_data_type(cudnn.data_type.FLOAT).set_dim(stats_dim).set_stride( + stats_stride + ) + if use_ragged_stats: + stats_t.set_ragged_offset(offset_stats).set_ragged_offset_multiplier(stats_mult) tensors.update(O=output_t, Stats=stats_t) return GraphEntry( graph=graph, @@ -1279,6 +1426,8 @@ def _fp8_forward( max_seqlen_kv, cu_seqlens_q, cu_seqlens_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, q, k, v, @@ -1338,14 +1487,39 @@ def _fp8_forward( False, ) return output, aux - del is_training, fast_zero_fill + del is_training q_format, _ = _q_kv_formats(qkv_layout) batch = cu_seqlens_q.numel() - 1 heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] d_v = v.shape[-1] - output_shape = _fp8_output_shape(batch, max_seqlen_q, heads, d_v, o_format) + output_shape = _fp8_output_shape( + batch, + max_seqlen_q, + heads, + d_v, + o_format, + total_tokens=q.shape[0] if o_format == "thd" else None, + ) output, amax_o = _allocate_fp8_kernel_output(o_quantizer, output_shape, fake_dtype, q.device) - stats = torch.empty((batch, heads, max_seqlen_q, 1), dtype=torch.float32, device=q.device) + if fast_zero_fill and o_format == "thd": + output_data = _quantized_data(output) if _is_float8_tensor(output) else output + output_data.zero_() + if amax_o is not None: + amax_o.zero_() + cudnn = import_cudnn_frontend() + use_ragged_stats = ( + q_format == "thd" + and cudnn.backend_version() >= 90600 + and torch.cuda.get_device_capability(q.device) != (12, 0) + ) + stats_shape, _, _ = _stats_layout( + batch=batch, + heads=heads, + max_seqlen_q=max_seqlen_q, + total_tokens_q=q.shape[0] if q_format == "thd" else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + stats = torch.empty(stats_shape, dtype=torch.float32, device=q.device) amax_s = ( s_quantizer.amax if isinstance(s_quantizer, Float8Quantizer) @@ -1357,6 +1531,15 @@ def _fp8_forward( ) rng_elts = (max_seqlen_q * max_seqlen_q + _FP8_THREADS_PER_CTA - 1) // _FP8_THREADS_PER_CTA rng_state = _reserve_philox_state(q.device, rng_gen, rng_elts) + cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded + cu_seqlens_kv_padded = ( + cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded + ) + kv_offsets_source = _kv_ragged_offsets_source( + qkv_layout, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, + ) q_data = _quantized_data(q) k_data = _quantized_data(k) @@ -1410,6 +1593,8 @@ def _fp8_forward( softmax_offset=softmax_offset, cu_seqlens_q=cu_seqlens_q, cu_seqlens_kv=cu_seqlens_kv, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, ) put_graph_entry(key, entry) @@ -1439,8 +1624,28 @@ def _fp8_forward( variant_pack[t["amax_s"]] = amax_s variant_pack[t["amax_o"]] = amax_o if "seq_len_q" in t: - variant_pack[t["seq_len_q"]] = _sequence_lengths(cu_seqlens_q) - variant_pack[t["seq_len_kv"]] = _sequence_lengths(cu_seqlens_kv) + graph_batch = t["_graph_batch"] + variant_pack[t["seq_len_q"]] = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + variant_pack[t["seq_len_kv"]] = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) + graph_batch = t["_graph_batch"] + if "offset_q" in t: + variant_pack[t["offset_q"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, q_data.stride(0) + ) + variant_pack[t["offset_o"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, output_data.stride(0) + ) + if "offset_k" in t: + variant_pack[t["offset_k"]] = _element_ragged_offsets( + kv_offsets_source, graph_batch, k_data.stride(0) + ) + variant_pack[t["offset_v"]] = _element_ragged_offsets( + kv_offsets_source, graph_batch, v_data.stride(0) + ) + if "offset_stats" in t: + variant_pack[t["offset_stats"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, stats.stride(0) + ) if "dropout_seed" in t: variant_pack[t["dropout_seed"]] = rng_state[:1] variant_pack[t["dropout_offset"]] = rng_state[1:] @@ -1513,6 +1718,8 @@ def fused_attn_fwd( max_seqlen_kv, cu_seqlens_q, cu_seqlens_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, q, k, v, @@ -1886,14 +2093,13 @@ def _build_fp8_bwd_graph( d_softmax_offset, cu_seqlens_q, cu_seqlens_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, ): cudnn = import_cudnn_frontend() q_format, kv_format = _q_kv_formats(qkv_layout) dq_format, dkv_format = _q_kv_formats(dqkv_layout) batch = cu_seqlens_q.numel() - 1 - heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] - kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] - d_qk, d_value = q.shape[-1], v.shape[-1] q_data = _quantized_data(q) k_data = _quantized_data(k) v_data = _quantized_data(v) @@ -1902,34 +2108,106 @@ def _build_fp8_bwd_graph( dq_data = _quantized_data(d_q) if _is_float8_tensor(d_q) else d_q dk_data = _quantized_data(d_k) if _is_float8_tensor(d_k) else d_k dv_data = _quantized_data(d_v) if _is_float8_tensor(d_v) else d_v + heads = q_data.shape[-2] if q_format != "bhsd" else q_data.shape[1] + kv_heads = k_data.shape[-2] if kv_format != "bhsd" else k_data.shape[1] + d_qk, d_value = q_data.shape[-1], v_data.shape[-1] + is_ragged_q = q_format == "thd" + is_ragged_kv = kv_format == "thd" + use_token_buckets = cudnn.backend_version() >= 90600 and torch.cuda.get_device_capability( + q.device + ) != (12, 0) + use_ragged_stats = is_ragged_q and use_token_buckets + use_legacy_offsets = is_ragged_q or is_ragged_kv + graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch + graph_seqlen_q = ( + _max_ragged_tokens(q_data.shape[0]) + if is_ragged_q and use_token_buckets + else max_seqlen_q + ) + graph_seqlen_kv = ( + _max_ragged_tokens(k_data.shape[0]) + if is_ragged_kv and use_token_buckets + else max_seqlen_kv + ) graph = make_graph(_fp8_cudnn_dtype(q), q.device, name="te_fp8_sdpa_bwd") - tensors: Dict[str, Any] = {} + tensors: Dict[str, Any] = { + "_legacy_offsets": use_legacy_offsets, + "_graph_batch": graph_batch, + } + + offset_q = offset_o = offset_k = offset_v = offset_stats = None + if is_ragged_q: + offset_q, q_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_q", + length=graph_batch + 1, + data_type=torch.int64, + ) + offset_o, o_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_o", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors.update(offset_q=offset_q, offset_o=offset_o) + else: + q_mult = o_mult = 1 + if is_ragged_kv: + offset_k, k_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1, + name="offset_k", + length=graph_batch + 1, + data_type=torch.int64, + ) + offset_v, v_mult = _ragged_offset_tensor( + graph, + cu_seqlens_kv_padded, + multiplier=1, + name="offset_v", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors.update(offset_k=offset_k, offset_v=offset_v) + else: + k_mult = v_mult = 1 q_t = _make_bhsd_graph_tensor( graph, q_data, q_format, - batch=batch, - max_seqlen=max_seqlen_q, + batch=graph_batch, + max_seqlen=graph_seqlen_q, data_type=_fp8_cudnn_dtype(q), + ragged_offset=offset_q, + ragged_offset_multiplier=q_mult, name="Q", ) k_t = _make_bhsd_graph_tensor( graph, k_data, kv_format, - batch=batch, - max_seqlen=max_seqlen_kv, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, data_type=_fp8_cudnn_dtype(k), + ragged_offset=offset_k, + ragged_offset_multiplier=k_mult, name="K", ) v_t = _make_bhsd_graph_tensor( graph, v_data, kv_format, - batch=batch, - max_seqlen=max_seqlen_kv, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, data_type=_fp8_cudnn_dtype(v), + ragged_offset=offset_v, + ragged_offset_multiplier=v_mult, name="V", ) o_dtype = _fp8_cudnn_dtype(o) if _is_float8_tensor(o) else o.dtype @@ -1937,21 +2215,51 @@ def _build_fp8_bwd_graph( graph, o_data, o_format, - batch=batch, - max_seqlen=max_seqlen_q, + batch=graph_batch, + max_seqlen=graph_seqlen_q, data_type=o_dtype, + ragged_offset=offset_o, + ragged_offset_multiplier=o_mult, name="O", ) do_t = _make_bhsd_graph_tensor( graph, do_data, do_format, - batch=batch, - max_seqlen=max_seqlen_q, + batch=graph_batch, + max_seqlen=graph_seqlen_q, data_type=_fp8_cudnn_dtype(d_o), + ragged_offset=offset_o, + ragged_offset_multiplier=o_mult, name="dO", ) - stats_t = graph.tensor_like(stats, name="Stats") + _, stats_dim, stats_stride = _stats_layout( + batch=graph_batch, + heads=heads, + max_seqlen_q=graph_seqlen_q, + total_tokens_q=q_data.shape[0] if is_ragged_q else batch * max_seqlen_q, + ragged=use_ragged_stats, + ) + if use_ragged_stats: + offset_stats, stats_mult = _ragged_offset_tensor( + graph, + cu_seqlens_q_padded, + multiplier=1, + name="offset_stats", + length=graph_batch + 1, + data_type=torch.int64, + ) + tensors["offset_stats"] = offset_stats + else: + stats_mult = 1 + stats_t = graph.tensor( + name="Stats", + dim=stats_dim, + stride=stats_stride, + data_type=cudnn.data_type.FLOAT, + ragged_offset=offset_stats, + ragged_offset_multiplier=stats_mult, + ) tensors.update(Q=q_t, K=k_t, V=v_t, O=o_t, dO=do_t, Stats=stats_t) options = _mask_options( @@ -1973,8 +2281,8 @@ def _build_fp8_bwd_graph( use_deterministic_algorithm=deterministic, ) if is_padding: - seq_q = _sequence_lengths(cu_seqlens_q) - seq_kv = _sequence_lengths(cu_seqlens_kv) + seq_q = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + seq_kv = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) seq_q_t = graph.tensor_like(seq_q, name="seq_len_q") seq_kv_t = graph.tensor_like(seq_kv, name="seq_len_kv") tensors.update(seq_len_q=seq_q_t, seq_len_kv=seq_kv_t) @@ -2195,21 +2503,38 @@ def mx_scale(name, h, s, d, fmt): ).set_stride((1, 1, 1, 1)) tensors[name] = amax_t + dq_dim, dq_stride = _logical_bhsd_desc( + dq_data, + dq_format, + batch=graph_batch, + max_seqlen=graph_seqlen_q, + ) + dk_dim, dk_stride = _logical_bhsd_desc( + dk_data, + dkv_format, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, + ) + dv_dim, dv_stride = _logical_bhsd_desc( + dv_data, + dkv_format, + batch=graph_batch, + max_seqlen=graph_seqlen_kv, + ) dq_t.set_output(True).set_data_type( _fp8_cudnn_dtype(d_q) if _is_float8_tensor(d_q) else d_q.dtype - ).set_dim((batch, heads, max_seqlen_q, d_qk)).set_stride( - _format_stride(batch, heads, max_seqlen_q, d_qk, dq_format) - ) + ).set_dim(dq_dim).set_stride(dq_stride) dk_t.set_output(True).set_data_type( _fp8_cudnn_dtype(d_k) if _is_float8_tensor(d_k) else d_k.dtype - ).set_dim((batch, kv_heads, max_seqlen_kv, d_qk)).set_stride( - _format_stride(batch, kv_heads, max_seqlen_kv, d_qk, dkv_format) - ) + ).set_dim(dk_dim).set_stride(dk_stride) dv_t.set_output(True).set_data_type( _fp8_cudnn_dtype(d_v) if _is_float8_tensor(d_v) else d_v.dtype - ).set_dim((batch, kv_heads, max_seqlen_kv, d_value)).set_stride( - _format_stride(batch, kv_heads, max_seqlen_kv, d_value, dkv_format) - ) + ).set_dim(dv_dim).set_stride(dv_stride) + if is_ragged_q: + dq_t.set_ragged_offset(offset_q).set_ragged_offset_multiplier(q_mult) + if is_ragged_kv: + dk_t.set_ragged_offset(offset_k).set_ragged_offset_multiplier(k_mult) + dv_t.set_ragged_offset(offset_v).set_ragged_offset_multiplier(v_mult) tensors.update(dQ=dq_t, dK=dk_t, dV=dv_t) return GraphEntry( graph=graph, @@ -2223,6 +2548,8 @@ def _fp8_backward( max_seqlen_kv, cu_seqlens_q, cu_seqlens_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, q, k, v, @@ -2283,6 +2610,7 @@ def _fp8_backward( deterministic=deterministic, ) q_format, kv_format = _q_kv_formats(qkv_layout) + dq_format, dkv_format = _q_kv_formats(dqkv_layout) batch = cu_seqlens_q.numel() - 1 heads = q.shape[-2] if q_format != "bhsd" else q.shape[1] kv_heads = k.shape[-2] if kv_format != "bhsd" else k.shape[1] @@ -2298,10 +2626,18 @@ def _fp8_backward( dqkv_layout=dqkv_layout, dtype=output_dtype, device=q.device, - zero=fast_zero_fill, + zero=fast_zero_fill + or ( + not isinstance(dqkv_quantizer, Float8Quantizer) + and (dq_format == "thd" or dkv_format == "thd") + ), + total_tokens_q=q.shape[0] if dq_format == "thd" else None, + total_tokens_kv=k.shape[0] if dkv_format == "thd" else None, ) if isinstance(dqkv_quantizer, Float8Quantizer): d_q, d_k, d_v = _wrap_float8_grad_outputs(dqkv_quantizer, grad_data, fake_dtype) + if fast_zero_fill and (dq_format == "thd" or dkv_format == "thd"): + dqkv_quantizer.amax.zero_() else: d_q, d_k, d_v = grad_data stats, rng_state = aux_ctx_tensors[:2] @@ -2309,6 +2645,15 @@ def _fp8_backward( d_o_f16 = aux_ctx_tensors[-1] if _is_mxfp8_tensor(q) else None d_softmax_offset = torch.empty_like(softmax_offset) if softmax_offset is not None else None hidden_amax = [torch.zeros(1, dtype=torch.float32, device=q.device) for _ in range(4)] + cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded + cu_seqlens_kv_padded = ( + cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded + ) + kv_offsets_source = _kv_ragged_offsets_source( + qkv_layout, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, + ) key = ( "fp8_bwd", @@ -2372,6 +2717,8 @@ def _fp8_backward( d_softmax_offset=d_softmax_offset, cu_seqlens_q=cu_seqlens_q, cu_seqlens_kv=cu_seqlens_kv, + cu_seqlens_q_padded=cu_seqlens_q_padded, + cu_seqlens_kv_padded=cu_seqlens_kv_padded, ) put_graph_entry(key, entry) @@ -2443,8 +2790,32 @@ def _fp8_backward( for name, value in zip(("amax_dq", "amax_dk", "amax_dv", "amax_dp"), amax_values): variant_pack[t[name]] = value if "seq_len_q" in t: - variant_pack[t["seq_len_q"]] = _sequence_lengths(cu_seqlens_q) - variant_pack[t["seq_len_kv"]] = _sequence_lengths(cu_seqlens_kv) + graph_batch = t["_graph_batch"] + variant_pack[t["seq_len_q"]] = _padded_sequence_lengths(cu_seqlens_q, graph_batch) + variant_pack[t["seq_len_kv"]] = _padded_sequence_lengths(cu_seqlens_kv, graph_batch) + graph_batch = t["_graph_batch"] + q_data = _quantized_data(q) + k_data = _quantized_data(k) + v_data = _quantized_data(v) + o_data = _quantized_data(o) if _is_float8_tensor(o) else o + if "offset_q" in t: + variant_pack[t["offset_q"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, q_data.stride(0) + ) + variant_pack[t["offset_o"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, o_data.stride(0) + ) + if "offset_k" in t: + variant_pack[t["offset_k"]] = _element_ragged_offsets( + kv_offsets_source, graph_batch, k_data.stride(0) + ) + variant_pack[t["offset_v"]] = _element_ragged_offsets( + kv_offsets_source, graph_batch, v_data.stride(0) + ) + if "offset_stats" in t: + variant_pack[t["offset_stats"]] = _element_ragged_offsets( + cu_seqlens_q_padded, graph_batch, stats.stride(0) + ) if "dropout_seed" in t: variant_pack[t["dropout_seed"]] = rng_state[:1] variant_pack[t["dropout_offset"]] = rng_state[1:] @@ -2511,6 +2882,8 @@ def fused_attn_bwd( max_seqlen_kv, cu_seqlens_q, cu_seqlens_kv, + cu_seqlens_q_padded, + cu_seqlens_kv_padded, q, k, v, From 4c6ecc5c6f853aa144344d75c97ebd61c6839068 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 23:10:08 +0000 Subject: [PATCH 24/36] Reject affected cuDNN versions for JAX MXFP8 attention Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fp8_fused_attn.py | 18 +++++++++++++++ .../jax/cpp_extensions/fp8_attention.py | 22 +++++++++++++++---- 2 files changed, 36 insertions(+), 4 deletions(-) diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py index d8844659a33..0b06230e0c3 100644 --- a/tests/jax/test_fp8_fused_attn.py +++ b/tests/jax/test_fp8_fused_attn.py @@ -26,6 +26,7 @@ from transformer_engine.jax.cpp_extensions.fp8_attention import ( _mx_scale, _mxfp8_scale_inv, + _validate_mxfp8_runtime, ) from transformer_engine.jax.flax import DotProductAttention from transformer_engine.jax.quantize import ( @@ -115,6 +116,8 @@ def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): """Each supported recipe executes FP8 DPA behind FP16/BF16 module boundaries.""" _require_gpu(min_arch, min_cudnn) + if fp8_recipe.mxfp8() and get_cudnn_version() in (92300, 92301): + pytest.skip("cuDNN 9.23.0 and 9.23.1 have known MXFP8 SDPA correctness issues.") batch, seqlen, heads, dim = 2, 128, 8, 128 q_key, k_key, v_key, do_key = jax.random.split(jax.random.PRNGKey(1234), 4) shape = (batch, seqlen, heads, dim) @@ -211,6 +214,21 @@ class tensor_reordering: assert tensor.reordering == FakeCudnn.tensor_reordering.F8_128x4 +@pytest.mark.parametrize("cudnn_version", ((9, 23, 0), (9, 23, 1))) +def test_mxfp8_attention_rejects_affected_cudnn(cudnn_version): + """Known-bad cuDNN 9.23 patch releases are rejected before graph execution.""" + + with pytest.raises(ValueError, match="known SDPA correctness issues"): + _validate_mxfp8_runtime(cudnn_version, 100) + + +@pytest.mark.parametrize("cudnn_version", ((9, 22, 9), (9, 23, 2))) +def test_mxfp8_attention_accepts_neighboring_cudnn(cudnn_version): + """The cuDNN exclusion remains limited to the affected patch releases.""" + + _validate_mxfp8_runtime(cudnn_version, 100) + + @pytest.mark.parametrize("feature", ("alibi", "bottom_right")) def test_fused_attention_parity_features(feature): """JAX executes ALiBi and explicit bottom-right diagonal attention.""" diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index 778654f153c..e56d37e106e 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -1008,6 +1008,18 @@ def _sequence_lengths(sequence_descriptor, config): return q_seqlen.flatten(), kv_seqlen.flatten() +def _validate_mxfp8_runtime(cudnn_version, device_arch): + """Reject runtimes that cannot safely execute MXFP8 attention.""" + + if cudnn_version < (9, 21, 0) or device_arch < 100: + raise ValueError("MXFP8 attention requires cuDNN 9.21 and SM100 or newer.") + if cudnn_version in ((9, 23, 0), (9, 23, 1)): + raise ValueError( + "MXFP8 attention is disabled with cuDNN 9.23.0 and 9.23.1 due to known " + "SDPA correctness issues." + ) + + def _validate_fp8_support(qkv, quantizers, config, mode): if config.qkv_layout.is_qkvpacked(): q = k = v = qkv[0] @@ -1022,6 +1034,8 @@ def _validate_fp8_support(qkv, quantizers, config, mode): jnp.dtype(jnp.float8_e4m3fn): "float8_e4m3", jnp.dtype(jnp.float8_e5m2): "float8_e5m2", }.get(q_dtype, str(q_dtype)) + cudnn_version = get_cudnn_version() + device_arch = _device_arch() support = check_fp8_fused_attention_support( FusedAttentionConfig( is_training=bool(config.is_training), @@ -1042,14 +1056,14 @@ def _validate_fp8_support(qkv, quantizers, config, mode): return_max_logit=False, cuda_graph=False, deterministic=not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))), - cudnn_version=get_cudnn_version(), - sm_arch=_device_arch(), + cudnn_version=cudnn_version, + sm_arch=device_arch, ) ) if not support.supported: raise ValueError(f"Unsupported JAX FP8 attention configuration: {support.reason}.") - if mode == "mxfp8" and (get_cudnn_version() < (9, 21, 0) or _device_arch() < 100): - raise ValueError("MXFP8 attention requires cuDNN 9.21 and SM100 or newer.") + if mode == "mxfp8": + _validate_mxfp8_runtime(cudnn_version, device_arch) def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): From 138dc2ba73f6e146bf118b64d93b2c47ccce5457 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 23:39:30 +0000 Subject: [PATCH 25/36] Gate JAX FP8 attention by recipe support Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fp8_fused_attn.py | 62 +++++------ tests/test_common_attention_helpers.py | 81 ++++++++++++-- transformer_engine/common/attention/cudnn.py | 89 +++++++++++---- .../jax/cpp_extensions/fp8_attention.py | 101 +++++++++++------- 4 files changed, 233 insertions(+), 100 deletions(-) diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py index 0b06230e0c3..62a8bb56187 100644 --- a/tests/jax/test_fp8_fused_attn.py +++ b/tests/jax/test_fp8_fused_attn.py @@ -26,7 +26,6 @@ from transformer_engine.jax.cpp_extensions.fp8_attention import ( _mx_scale, _mxfp8_scale_inv, - _validate_mxfp8_runtime, ) from transformer_engine.jax.flax import DotProductAttention from transformer_engine.jax.quantize import ( @@ -37,7 +36,7 @@ from transformer_engine.jax.sharding import MeshResource -def _require_gpu(min_arch=90, min_cudnn=90700): +def _require_gpu(min_arch=90, min_cudnn=90700, max_arch=None): try: if not any(device.platform == "gpu" for device in jax.devices()): pytest.skip("A CUDA device is required.") @@ -46,6 +45,10 @@ def _require_gpu(min_arch=90, min_cudnn=90700): pytest.skip(f"A usable CUDA device is required: {exc}") if arch < min_arch: pytest.skip(f"This test requires SM{min_arch} or newer, found SM{arch}.") + if max_arch is not None and arch >= max_arch: + pytest.skip( + f"This test requires an architecture older than SM{max_arch}, found SM{arch}." + ) cudnn_version = get_cudnn_version() if cudnn_version < min_cudnn: pytest.skip(f"This test requires cuDNN {min_cudnn}, found {cudnn_version}.") @@ -56,9 +59,9 @@ def _require_gpu(min_arch=90, min_cudnn=90700): def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): q_seqlen, kv_seqlen = q.shape[1], k.shape[1] - scores = jnp.einsum("bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32)) / sqrt( - q.shape[-1] - ) + scores = jnp.einsum( + "bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32) + ) / sqrt(q.shape[-1]) q_pos = jnp.arange(q_seqlen)[:, None] kv_pos = jnp.arange(kv_seqlen)[None, :] shift = kv_seqlen - q_seqlen if bottom_right else 0 @@ -78,7 +81,9 @@ def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): allowed = kv_pos <= q_pos + shift scores = jnp.where(allowed[None, None, :, :], scores, -jnp.inf) probabilities = jax.nn.softmax(scores, axis=-1) - return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype(q.dtype) + return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype( + q.dtype + ) def _assert_fp8_close(actual, expected): @@ -99,8 +104,8 @@ def _assert_fp8_close(actual, expected): ), pytest.param( recipe.Float8CurrentScaling(fp8_dpa=True), - 90, - 90700, + 100, + 91400, id="current", ), pytest.param( @@ -111,11 +116,13 @@ def _assert_fp8_close(actual, expected): ), ), ) -@pytest.mark.parametrize("input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16")) +@pytest.mark.parametrize( + "input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16") +) def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): """Each supported recipe executes FP8 DPA behind FP16/BF16 module boundaries.""" - _require_gpu(min_arch, min_cudnn) + _require_gpu(min_arch, min_cudnn, max_arch=120) if fp8_recipe.mxfp8() and get_cudnn_version() in (92300, 92301): pytest.skip("cuDNN 9.23.0 and 9.23.1 have known MXFP8 SDPA correctness issues.") batch, seqlen, heads, dim = 2, 128, 8, 128 @@ -137,7 +144,9 @@ def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): ) def loss_fn(variables, query, key, value): - output = module.apply(variables, query, key, value, descriptor, deterministic=True) + output = module.apply( + variables, query, key, value, descriptor, deterministic=False + ) loss = jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)) return loss, output @@ -147,7 +156,9 @@ def reference_loss(query, key, value): return loss, output with autocast(enabled=True, recipe=fp8_recipe, mesh_resource=MeshResource()): - variables = module.init(jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=True) + variables = module.init( + jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=False + ) (_, output), (_, dq, dk, dv) = jax.value_and_grad( loss_fn, argnums=(0, 1, 2, 3), has_aux=True )(variables, q, k, v) @@ -170,7 +181,9 @@ def test_mxfp8_attention_scale_layout(): q_layout=QuantizeLayout.ROWWISE_COLWISE, data_layout="NN", ) - tensor = quantizer.quantize(jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2) + tensor = quantizer.quantize( + jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2 + ) assert tensor.rowwise_tensor.scale_inv.shape == (2, 64, 8, 2) assert tensor.colwise_tensor.scale_inv.shape == (2, 2, 8, 64) assert _mxfp8_scale_inv(tensor).shape == (2, 8, 128, 4) @@ -214,21 +227,6 @@ class tensor_reordering: assert tensor.reordering == FakeCudnn.tensor_reordering.F8_128x4 -@pytest.mark.parametrize("cudnn_version", ((9, 23, 0), (9, 23, 1))) -def test_mxfp8_attention_rejects_affected_cudnn(cudnn_version): - """Known-bad cuDNN 9.23 patch releases are rejected before graph execution.""" - - with pytest.raises(ValueError, match="known SDPA correctness issues"): - _validate_mxfp8_runtime(cudnn_version, 100) - - -@pytest.mark.parametrize("cudnn_version", ((9, 22, 9), (9, 23, 2))) -def test_mxfp8_attention_accepts_neighboring_cudnn(cudnn_version): - """The cuDNN exclusion remains limited to the affected patch releases.""" - - _validate_mxfp8_runtime(cudnn_version, 100) - - @pytest.mark.parametrize("feature", ("alibi", "bottom_right")) def test_fused_attention_parity_features(feature): """JAX executes ALiBi and explicit bottom-right diagonal attention.""" @@ -240,7 +238,9 @@ def test_fused_attention_parity_features(feature): q = jax.random.normal(q_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 k = jax.random.normal(k_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 v = jax.random.normal(v_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 - doutput = jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 + doutput = ( + jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 + ) q_lengths = jnp.full((batch,), q_seqlen, dtype=jnp.int32) kv_lengths = jnp.full((batch,), kv_seqlen, dtype=jnp.int32) descriptor = SequenceDescriptor.from_seqlens((q_lengths, kv_lengths)) @@ -292,7 +292,9 @@ def reference_loss(query, key, value): ) return jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)), output - (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)(q, k, v) + (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)( + q, k, v + ) (_, reference), reference_grads = jax.value_and_grad( reference_loss, argnums=(0, 1, 2), has_aux=True )(q, k, v) diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py index 0906cf910e0..87889f1c063 100644 --- a/tests/test_common_attention_helpers.py +++ b/tests/test_common_attention_helpers.py @@ -128,11 +128,70 @@ def test_shared_f16_policy_framework_parity_features(): def test_shared_fp8_policy(): fp8 = _attention_config(q_dtype="float8_e4m3", kv_dtype="float8_e4m3", sm_arch=100) assert check_fp8_fused_attention_support(fp8).supported - assert check_fp8_fused_attention_support(replace(fp8, cudnn_version=(9, 10, 1))).supported - assert not check_fp8_fused_attention_support(replace(fp8, cudnn_version=(9, 10, 0))).supported - assert not check_fp8_fused_attention_support(replace(fp8, bias_type="alibi")).supported - assert not check_fp8_fused_attention_support(replace(fp8, return_max_logit=True)).supported - assert not check_fp8_fused_attention_support(replace(fp8, head_dim_qk=200)).supported + assert check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 10, 1)) + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 10, 0)) + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, bias_type="alibi") + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, return_max_logit=True) + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, head_dim_qk=200) + ).supported + + +def test_shared_fp8_scaling_mode_policy(): + fp8 = _attention_config(q_dtype="float8_e4m3", kv_dtype="float8_e4m3", sm_arch=100) + + assert check_fp8_fused_attention_support(fp8, scaling_mode="delayed").supported + assert check_fp8_fused_attention_support(fp8, scaling_mode="current").supported + assert check_fp8_fused_attention_support(fp8, scaling_mode="mxfp8").supported + + hopper = replace(fp8, sm_arch=90) + assert check_fp8_fused_attention_support(hopper, scaling_mode="delayed").supported + assert not check_fp8_fused_attention_support( + hopper, scaling_mode="current" + ).supported + assert not check_fp8_fused_attention_support(hopper, scaling_mode="mxfp8").supported + + assert not check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 13, 9)), scaling_mode="current" + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 20, 9)), scaling_mode="mxfp8" + ).supported + assert not check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 23, 0)), scaling_mode="mxfp8" + ).supported + assert check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 23, 2)), scaling_mode="mxfp8" + ).supported + + assert not check_fp8_fused_attention_support( + replace(fp8, sm_arch=120), scaling_mode="delayed" + ).supported + + +def test_shared_fp8_deterministic_backward_policy(): + fp8 = _attention_config( + is_training=True, + deterministic=True, + q_dtype="float8_e4m3", + kv_dtype="float8_e4m3", + sm_arch=100, + ) + + assert not check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 18, 9)), scaling_mode="delayed" + ).supported + assert check_fp8_fused_attention_support( + replace(fp8, cudnn_version=(9, 19, 0)), scaling_mode="delayed" + ).supported def test_shared_fp8_thd_policy(): @@ -152,11 +211,15 @@ def test_shared_fp8_thd_policy(): assert not check_fp8_fused_attention_support( replace(thd, cudnn_version=(9, 22, 9)) ).supported - assert not check_fp8_fused_attention_support(replace(thd, mask_type="no_mask")).supported + assert not check_fp8_fused_attention_support( + replace(thd, mask_type="no_mask") + ).supported assert not check_fp8_fused_attention_support( replace(thd, is_training=True, sm_arch=90) ).supported - assert not check_fp8_fused_attention_support(replace(thd, head_dim_qk=144)).supported + assert not check_fp8_fused_attention_support( + replace(thd, head_dim_qk=144) + ).supported sink_backward = replace(thd, is_training=True, softmax_type="learnable") assert not check_fp8_fused_attention_support( @@ -442,7 +505,9 @@ def test_shared_cudnn_graph_reports_build_diagnostics(): def test_shared_cudnn_graph_support_error_has_context(): with pytest.raises(RuntimeError, match="cuDNN test graph is not supported"): - build_cudnn_graph(_FakeCudnn(), _FakeGraph(unsupported=True), description="test") + build_cudnn_graph( + _FakeCudnn(), _FakeGraph(unsupported=True), description="test" + ) @pytest.fixture diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index 3f2bd3efc4c..2257d4409ef 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -191,7 +191,9 @@ def cudnn_mask_options( ) version = encode_cudnn_version(cudnn_version) options: dict[str, bool | int | str] = { - "diagonal_alignment": ("bottom_right" if mask.bottom_right_diagonal else "top_left"), + "diagonal_alignment": ( + "bottom_right" if mask.bottom_right_diagonal else "top_left" + ), "is_padding": mask.padding, } if version < 90600: @@ -354,7 +356,9 @@ def check_f16_fused_attention_support( and (dqk, dv) != (192, 128) and dqk != dv ): - return _unsupported("this Hopper backward head-dimension combination is unsupported") + return _unsupported( + "this Hopper backward head-dimension combination is unsupported" + ) alibi_supported = ( bias == "alibi" @@ -434,7 +438,10 @@ def check_f16_fused_attention_support( and bias == "no_bias" and dropout == 0.0 ) - or (mask in ("causal_bottom_right", "padding_causal_bottom_right") and sq <= skv) + or ( + mask in ("causal_bottom_right", "padding_causal_bottom_right") + and sq <= skv + ) ) mask_ok = mask_ok or modern_mask_ok if not mask_ok: @@ -466,7 +473,10 @@ def check_f16_fused_attention_support( or ( left >= -1 and right == 0 - and (mask in ("no_mask", "causal") or (mask == "causal_bottom_right" and sq == skv)) + and ( + mask in ("no_mask", "causal") + or (mask == "causal_bottom_right" and sq == skv) + ) and sq <= skv and dropout == 0.0 and bias == "no_bias" @@ -573,6 +583,8 @@ def check_f16_fused_attention_support( def check_fp8_fused_attention_support( config: FusedAttentionConfig, + *, + scaling_mode: str | None = None, ) -> FusedAttentionSupport: """Check the shared cuDNN FP8/MXFP8 fused-attention compatibility policy.""" @@ -590,24 +602,51 @@ def check_fp8_fused_attention_support( dv = int(config.head_dim_v) mask = config.mask_type + if scaling_mode not in (None, "delayed", "current", "mxfp8"): + return _unsupported(f"unknown FP8 attention scaling mode {scaling_mode!r}") if arch < 90: return _unsupported("FP8 attention requires SM90 or newer") + if arch >= 120: + return _unsupported("FP8 attention is not supported on SM120 or newer") + if config.is_training and config.deterministic and version < 91900: + return _unsupported( + "deterministic FP8 attention backward requires cuDNN 9.19 or newer" + ) + if scaling_mode == "current": + if arch < 100: + return _unsupported("FP8 current-scaling attention requires SM100 or newer") + if version < 91400: + return _unsupported( + "FP8 current-scaling attention requires cuDNN 9.14 or newer" + ) + if scaling_mode == "mxfp8": + if arch < 100: + return _unsupported("MXFP8 attention requires SM100 or newer") + if version < 92100: + return _unsupported("MXFP8 attention requires cuDNN 9.21 or newer") + if version in (92300, 92301): + return _unsupported("cuDNN 9.23.0 and 9.23.1 have known MXFP8 SDPA issues") if config.bias_type != "no_bias": return _unsupported("FP8 attention does not support attention bias") if config.return_max_logit: return _unsupported("FP8 attention does not support returning max logits") if version == 91000: return _unsupported("cuDNN 9.10.0 has known SDPA issues") - if requires_64bit_ragged_offset( - layout, - config.num_attn_heads, - config.num_gqa_groups, - sq, - skv, - dqk, - dv, - ) and version < 90500: - return _unsupported("FP8 attention requires cuDNN 9.5 for 64-bit ragged offsets") + if ( + requires_64bit_ragged_offset( + layout, + config.num_attn_heads, + config.num_gqa_groups, + sq, + skv, + dqk, + dv, + ) + and version < 90500 + ): + return _unsupported( + "FP8 attention requires cuDNN 9.5 for 64-bit ragged offsets" + ) is_thd = layout.qkv_format == "thd" if is_thd: @@ -618,9 +657,13 @@ def check_fp8_fused_attention_support( if config.is_training and arch < 100: return _unsupported("FP8 THD attention backward requires SM100 or newer") if config.is_training and config.softmax_type != "vanilla" and version < 92600: - return _unsupported("FP8 THD sink-token backward requires cuDNN 9.26 or newer") + return _unsupported( + "FP8 THD sink-token backward requires cuDNN 9.26 or newer" + ) if arch >= 100 and (dqk > 128 or dv > 128): - return _unsupported("FP8 THD attention supports head dimensions up to 128 on SM100+") + return _unsupported( + "FP8 THD attention supports head dimensions up to 128 on SM100+" + ) shape_mask_ok = ( ( @@ -660,12 +703,14 @@ def check_fp8_fused_attention_support( return _unsupported("FP8 attention shape or mask is not supported") format_softmax_ok = ( - version < 92100 - and layout.qkv_format in ("bshd", "sbhd") - and config.softmax_type == "vanilla" - ) or ( - version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd") - ) or is_thd + ( + version < 92100 + and layout.qkv_format in ("bshd", "sbhd") + and config.softmax_type == "vanilla" + ) + or (version >= 92100 and layout.qkv_format in ("bshd", "sbhd", "bhsd")) + or is_thd + ) if not format_softmax_ok: return _unsupported("FP8 attention layout or softmax type is not supported") return FusedAttentionSupport(True) diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index e56d37e106e..27dbd8a022a 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -201,7 +201,9 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="Q", dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.qk_dim), - stride=_matrix_stride(info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen), + stride=_matrix_stride( + info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen + ), dtype=io_dtype, uid=_UID_Q, ) @@ -209,7 +211,9 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="K", dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.qk_dim), - stride=_matrix_stride(info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen), + stride=_matrix_stride( + info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen + ), dtype=io_dtype, uid=_UID_K, ) @@ -217,7 +221,9 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="V", dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.v_dim), - stride=_matrix_stride(info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen), + stride=_matrix_stride( + info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen + ), dtype=io_dtype, uid=_UID_V, ) @@ -262,7 +268,9 @@ def build_fp8_fwd_graph( graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) info, _, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) - input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) + input_bindings = list( + _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) + ) tensors = {"q": q, "k": k, "v": v} options, is_padding = _fp8_options(cudnn, info, config) options["generate_stats"] = True @@ -276,7 +284,9 @@ def build_fp8_fwd_graph( options["diagonal_band_left_bound"] = options.pop("left_bound") if "right_bound" in options: options["diagonal_band_right_bound"] = options.pop("right_bound") - padded = mxfp8_padded_sizes(info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim) + padded = mxfp8_padded_sizes( + info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim + ) scale_specs = ( ( "descale_q", @@ -345,7 +355,9 @@ def build_fp8_fwd_graph( uid=_UID_SEQ_KV, ) options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend((GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11))) + input_bindings.extend( + (GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11)) + ) output_bindings = [GraphBinding(_UID_O, 0), GraphBinding(_UID_STATS, 1)] if _is_dropout(config): @@ -382,23 +394,27 @@ def build_fp8_fwd_graph( output = op["output"] output.set_output(True).set_uid(_UID_O).set_data_type( cudnn_data_type(cudnn, output_dtype) - ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim)).set_stride( + ).set_dim( + (info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim) + ).set_stride( attention_format_stride( info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" ) ) stats = op["stats"] - stats.set_output(True).set_uid(_UID_STATS).set_data_type(cudnn.data_type.FLOAT).set_dim( - (info.input_batch, info.q_heads, info.q_max_seqlen, 1) - ).set_stride((info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1)) + stats.set_output(True).set_uid(_UID_STATS).set_data_type( + cudnn.data_type.FLOAT + ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, 1)).set_stride( + (info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1) + ) if mode != "mxfp8": for name, uid, offset in ( ("amax_s", _UID_AMAX_S, 0), ("amax_o", _UID_AMAX_O, 4), ): - op[name].set_output(True).set_uid(uid).set_data_type(cudnn.data_type.FLOAT).set_dim( - (1, 1, 1, 1) - ).set_stride((1, 1, 1, 1)) + op[name].set_output(True).set_uid(uid).set_data_type( + cudnn.data_type.FLOAT + ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 2, offset)) else: op["amax_o"].set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( @@ -523,9 +539,13 @@ def build_fp8_bwd_graph( cudnn = import_cudnn() graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) - info, io_dtype, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) + info, io_dtype, q, k, v = _graph_io_tensors( + graph, cudnn, q_aval, k_aval, v_aval, config + ) q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) - input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) + input_bindings = list( + _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) + ) o = _tensor( graph, name="O", @@ -580,7 +600,9 @@ def build_fp8_bwd_graph( uid=_UID_SEQ_KV, ) options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend((GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10))) + input_bindings.extend( + (GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10)) + ) if _is_dropout(config): seed = _tensor( @@ -659,7 +681,9 @@ def build_fp8_bwd_graph( GraphBinding(_UID_DO_F16, 27), ) ) - padded = mxfp8_padded_sizes(info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim) + padded = mxfp8_padded_sizes( + info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim + ) tensors.update(_mx_bwd_scales(graph, cudnn, info, padded, input_bindings)) else: for name, uid, index in ( @@ -707,9 +731,9 @@ def build_fp8_bwd_graph( ("amax_dv", _UID_AMAX_DV, 8), ("amax_dp", _UID_AMAX_DP, 12), ): - op[name].set_output(True).set_uid(uid).set_data_type(cudnn.data_type.FLOAT).set_dim( - (1, 1, 1, 1) - ).set_stride((1, 1, 1, 1)) + op[name].set_output(True).set_uid(uid).set_data_type( + cudnn.data_type.FLOAT + ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 3, offset)) else: for amax in op["amax"]: @@ -1008,18 +1032,6 @@ def _sequence_lengths(sequence_descriptor, config): return q_seqlen.flatten(), kv_seqlen.flatten() -def _validate_mxfp8_runtime(cudnn_version, device_arch): - """Reject runtimes that cannot safely execute MXFP8 attention.""" - - if cudnn_version < (9, 21, 0) or device_arch < 100: - raise ValueError("MXFP8 attention requires cuDNN 9.21 and SM100 or newer.") - if cudnn_version in ((9, 23, 0), (9, 23, 1)): - raise ValueError( - "MXFP8 attention is disabled with cuDNN 9.23.0 and 9.23.1 due to known " - "SDPA correctness issues." - ) - - def _validate_fp8_support(qkv, quantizers, config, mode): if config.qkv_layout.is_qkvpacked(): q = k = v = qkv[0] @@ -1055,15 +1067,18 @@ def _validate_fp8_support(qkv, quantizers, config, mode): window_size=tuple(int(value) for value in config.window_size), return_max_logit=False, cuda_graph=False, - deterministic=not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))), + deterministic=not bool( + int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) + ), cudnn_version=cudnn_version, sm_arch=device_arch, - ) + ), + scaling_mode=mode, ) if not support.supported: - raise ValueError(f"Unsupported JAX FP8 attention configuration: {support.reason}.") - if mode == "mxfp8": - _validate_mxfp8_runtime(cudnn_version, device_arch) + raise ValueError( + f"Unsupported JAX FP8 attention configuration: {support.reason}." + ) def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): @@ -1084,11 +1099,15 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): if config.qkv_layout.is_thd(): raise NotImplementedError("FP8 attention does not support THD layouts in JAX.") if mode == "mxfp8" and not config.qkv_layout.is_separate(): - raise NotImplementedError("JAX MXFP8 attention currently requires separate BSHD Q/K/V.") + raise NotImplementedError( + "JAX MXFP8 attention currently requires separate BSHD Q/K/V." + ) if getattr(config.attn_bias_type, "name", "") != "NO_BIAS": raise NotImplementedError("FP8 attention does not support attention bias.") if getattr(config.softmax_type, "name", "") != "VANILLA_SOFTMAX": - raise NotImplementedError("JAX FP8 attention currently supports vanilla softmax only.") + raise NotImplementedError( + "JAX FP8 attention currently supports vanilla softmax only." + ) _validate_fp8_support(qkv, quantizers, config, mode) if mode == "mxfp8" and _is_padding(config): raise NotImplementedError("JAX MXFP8 attention does not support padding masks.") @@ -1129,7 +1148,9 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): if mode == "delayed": quantizers.s.update(amax[0:1]) quantizers.o.update(amax[1:2]) - output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype(qkv[0].dtype) + output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype( + qkv[0].dtype + ) else: output = raw_output return output, ( From 51471fac606f75bf84c693c90161c8b449a6bdf1 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Fri, 11 Sep 2026 23:43:28 +0000 Subject: [PATCH 26/36] Match JAX current-scaling attention quantizers Signed-off-by: Vladimir Cherepanov --- tests/jax/test_fp8_fused_attn.py | 27 +++++++++ .../jax/cpp_extensions/fp8_attention.py | 57 +++++++++++++------ transformer_engine/jax/flax/module.py | 24 +++++++- 3 files changed, 89 insertions(+), 19 deletions(-) diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py index 62a8bb56187..33fd36aa4bf 100644 --- a/tests/jax/test_fp8_fused_attn.py +++ b/tests/jax/test_fp8_fused_attn.py @@ -5,6 +5,7 @@ """Coverage for JAX FP8 DPA and attention features shared with PyTorch.""" from math import sqrt +from types import SimpleNamespace import jax import jax.numpy as jnp @@ -26,9 +27,11 @@ from transformer_engine.jax.cpp_extensions.fp8_attention import ( _mx_scale, _mxfp8_scale_inv, + _validate_quantizer_modes, ) from transformer_engine.jax.flax import DotProductAttention from transformer_engine.jax.quantize import ( + AttentionQuantizerSet, BlockScaleQuantizer, QuantizeLayout, ScalingMode, @@ -93,6 +96,30 @@ def _assert_fp8_close(actual, expected): assert np.sqrt(np.mean(np.square(actual - expected))) < 0.11 +def _mode_quantizer(mode): + return SimpleNamespace(scaling_mode=mode) + + +def test_fp8_attention_quantizer_mode_assignments(): + """Current-scaling DPA uses delayed scaling for its internal S and dP tensors.""" + + current = _mode_quantizer(ScalingMode.CURRENT_TENSOR_SCALING) + delayed = _mode_quantizer(ScalingMode.DELAYED_TENSOR_SCALING) + quantizers = AttentionQuantizerSet( + qkv=current, + s=delayed, + o=current, + do=current, + dp=delayed, + dqkv=current, + ) + assert _validate_quantizer_modes(quantizers) == "current" + + quantizers.s = current + with pytest.raises(ValueError, match=r"s=current \(expected delayed\)"): + _validate_quantizer_modes(quantizers) + + @pytest.mark.parametrize( "fp8_recipe,min_arch,min_cudnn", ( diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index 27dbd8a022a..72947f079b2 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -927,6 +927,36 @@ def _scaling_mode(quantizer) -> str: raise ValueError(f"FP8 attention does not support scaling mode {mode}.") +def _validate_quantizer_modes(quantizers) -> str: + """Validate the supported scaling-mode assignment for each DPA tensor role.""" + + roles = ("qkv", "s", "o", "do", "dp", "dqkv") + mode = _scaling_mode(quantizers.qkv) + expected = { + "delayed": {role: "delayed" for role in roles}, + "current": { + "qkv": "current", + "s": "delayed", + "o": "current", + "do": "current", + "dp": "delayed", + "dqkv": "current", + }, + "mxfp8": {role: "mxfp8" for role in roles}, + }[mode] + actual = {role: _scaling_mode(getattr(quantizers, role)) for role in roles} + mismatches = [ + f"{role}={actual[role]} (expected {expected[role]})" + for role in roles + if actual[role] != expected[role] + ] + if mismatches: + raise ValueError( + "Unsupported FP8 attention quantizer modes: " + ", ".join(mismatches) + ) + return mode + + def _rowwise(tensor): return tensor.get_tensor(TensorUsage.LHS) @@ -1084,18 +1114,7 @@ def _validate_fp8_support(qkv, quantizers, config, mode): def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): """Quantize high-precision inputs and execute dense FP8 attention forward.""" - mode = _scaling_mode(quantizers.qkv) - if any( - _scaling_mode(q) != mode - for q in ( - quantizers.s, - quantizers.o, - quantizers.do, - quantizers.dp, - quantizers.dqkv, - ) - ): - raise ValueError("All FP8 attention quantizers must use the same scaling mode.") + mode = _validate_quantizer_modes(quantizers) if config.qkv_layout.is_thd(): raise NotImplementedError("FP8 attention does not support THD layouts in JAX.") if mode == "mxfp8" and not config.qkv_layout.is_separate(): @@ -1124,7 +1143,8 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): seed = _FusedAttnRNGStateChecker().check_seed( seed, config.dropout_probability, config.is_training ) - output_dtype = quantizers.o.q_dtype if mode == "delayed" else qkv[0].dtype + output_is_delayed = _scaling_mode(quantizers.o) == "delayed" + output_dtype = quantizers.o.q_dtype if output_is_delayed else qkv[0].dtype s_scale = _tensor_scale(quantizers.s) s_scale_inv = jnp.reciprocal(s_scale) output_scale_inv = _tensor_scale_inv(quantizers.o) @@ -1145,8 +1165,9 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): mode=mode, output_dtype=output_dtype, ) - if mode == "delayed": + if _scaling_mode(quantizers.s) == "delayed": quantizers.s.update(amax[0:1]) + if output_is_delayed: quantizers.o.update(amax[1:2]) output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype( qkv[0].dtype @@ -1217,7 +1238,8 @@ def fused_attn_fp8_bwd(ctx, doutput, config): q_t = k_t = do_t = empty_data q_scale_t = k_scale_t = do_scale_t = empty_scale - grad_dtype = quantizers.dqkv.q_dtype if mode == "delayed" else input_dtype + grad_is_delayed = _scaling_mode(quantizers.dqkv) == "delayed" + grad_dtype = quantizers.dqkv.q_dtype if grad_is_delayed else input_dtype grad_scale_inv = _tensor_scale_inv(quantizers.dqkv) dq, dk, dv, amax, _, _ = execute_fp8_bwd( q_data, @@ -1252,11 +1274,12 @@ def fused_attn_fp8_bwd(ctx, doutput, config): mode=mode, grad_dtype=grad_dtype, ) - if mode == "delayed": + if grad_is_delayed: quantizers.dqkv.update(jnp.max(amax[:3])) - quantizers.dp.update(amax[3:4]) dq, dk, dv = ( (tensor.astype(input_dtype) * grad_scale_inv).astype(input_dtype) for tensor in (dq, dk, dv) ) + if _scaling_mode(quantizers.dp) == "delayed": + quantizers.dp.update(amax[3:4]) return _split_gradient_outputs(dq, dk, dv, config.qkv_layout), quantizers diff --git a/transformer_engine/jax/flax/module.py b/transformer_engine/jax/flax/module.py index 5f6bb268d15..e0012bf5926 100644 --- a/transformer_engine/jax/flax/module.py +++ b/transformer_engine/jax/flax/module.py @@ -17,6 +17,7 @@ from jax.ad_checkpoint import checkpoint_name from transformer_engine.common.recipe import ( + DelayedScaling, MXFP8BlockScaling, ) @@ -426,12 +427,31 @@ def generate_attention_quantizer_set(self, fp8_recipe=None): first = self.generate_quantizer_set(postfix="_attention_qkv_s_do", fp8_recipe=fp8_recipe) second = self.generate_quantizer_set(postfix="_attention_o_dp", fp8_recipe=fp8_recipe) third = self.generate_quantizer_set(postfix="_attention_dqkv", fp8_recipe=fp8_recipe) + s_quantizer = first.kernel + dp_quantizer = second.dgrad + if fp8_recipe.float8_current_scaling(): + # S and dP are internal to the fused attention graph, so their current + # amax is not available before they are quantized. Match the PyTorch DPA + # recipe by maintaining one-step delayed scales for these two roles. + delayed_recipe = DelayedScaling( + fp8_format=fp8_recipe.fp8_format, + amax_history_len=1, + amax_compute_algo="most_recent", + fp8_dpa=fp8_recipe.fp8_dpa, + fp8_mha=fp8_recipe.fp8_mha, + ) + internal = self.generate_quantizer_set( + postfix="_attention_s_dp", + fp8_recipe=delayed_recipe, + ) + s_quantizer = internal.kernel + dp_quantizer = internal.dgrad return AttentionQuantizerSet( qkv=first.x, - s=first.kernel, + s=s_quantizer, o=second.x, do=first.dgrad, - dp=second.dgrad, + dp=dp_quantizer, dqkv=third.dgrad, ) From 13cec2401f5709cf84ef48c5fa37a1604cae8807 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 23:56:13 +0000 Subject: [PATCH 27/36] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- tests/jax/test_fp8_fused_attn.py | 38 +++----- tests/jax/test_fused_attn.py | 4 +- tests/test_common_attention_helpers.py | 40 +++------ .../common/attention/cache_debug.py | 9 +- transformer_engine/common/attention/cudnn.py | 38 ++------ transformer_engine/common/cudnn_frontend.py | 4 +- .../jax/cpp_extensions/flex_attention.py | 4 +- .../jax/cpp_extensions/fp8_attention.py | 88 ++++++------------- .../jax/csrc/extensions/attention.cpp | 6 +- .../csrc/extensions/attention_cache_debug.h | 30 +++---- .../jax/csrc/extensions/pybind.cpp | 3 +- .../dot_product_attention/cudnn_attention.py | 24 ++--- 12 files changed, 84 insertions(+), 204 deletions(-) diff --git a/tests/jax/test_fp8_fused_attn.py b/tests/jax/test_fp8_fused_attn.py index 33fd36aa4bf..6d689cbd36e 100644 --- a/tests/jax/test_fp8_fused_attn.py +++ b/tests/jax/test_fp8_fused_attn.py @@ -49,9 +49,7 @@ def _require_gpu(min_arch=90, min_cudnn=90700, max_arch=None): if arch < min_arch: pytest.skip(f"This test requires SM{min_arch} or newer, found SM{arch}.") if max_arch is not None and arch >= max_arch: - pytest.skip( - f"This test requires an architecture older than SM{max_arch}, found SM{arch}." - ) + pytest.skip(f"This test requires an architecture older than SM{max_arch}, found SM{arch}.") cudnn_version = get_cudnn_version() if cudnn_version < min_cudnn: pytest.skip(f"This test requires cuDNN {min_cudnn}, found {cudnn_version}.") @@ -62,9 +60,9 @@ def _require_gpu(min_arch=90, min_cudnn=90700, max_arch=None): def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): q_seqlen, kv_seqlen = q.shape[1], k.shape[1] - scores = jnp.einsum( - "bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32) - ) / sqrt(q.shape[-1]) + scores = jnp.einsum("bqhd,bkhd->bhqk", q.astype(jnp.float32), k.astype(jnp.float32)) / sqrt( + q.shape[-1] + ) q_pos = jnp.arange(q_seqlen)[:, None] kv_pos = jnp.arange(kv_seqlen)[None, :] shift = kv_seqlen - q_seqlen if bottom_right else 0 @@ -84,9 +82,7 @@ def _reference_attention(q, k, v, *, bottom_right=False, alibi=False): allowed = kv_pos <= q_pos + shift scores = jnp.where(allowed[None, None, :, :], scores, -jnp.inf) probabilities = jax.nn.softmax(scores, axis=-1) - return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype( - q.dtype - ) + return jnp.einsum("bhqk,bkhd->bqhd", probabilities, v.astype(jnp.float32)).astype(q.dtype) def _assert_fp8_close(actual, expected): @@ -143,9 +139,7 @@ def test_fp8_attention_quantizer_mode_assignments(): ), ), ) -@pytest.mark.parametrize( - "input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16") -) +@pytest.mark.parametrize("input_dtype", (jnp.float16, jnp.bfloat16), ids=("float16", "bfloat16")) def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): """Each supported recipe executes FP8 DPA behind FP16/BF16 module boundaries.""" @@ -171,9 +165,7 @@ def test_fp8_dpa_forward_backward(fp8_recipe, min_arch, min_cudnn, input_dtype): ) def loss_fn(variables, query, key, value): - output = module.apply( - variables, query, key, value, descriptor, deterministic=False - ) + output = module.apply(variables, query, key, value, descriptor, deterministic=False) loss = jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)) return loss, output @@ -183,9 +175,7 @@ def reference_loss(query, key, value): return loss, output with autocast(enabled=True, recipe=fp8_recipe, mesh_resource=MeshResource()): - variables = module.init( - jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=False - ) + variables = module.init(jax.random.PRNGKey(0), q, k, v, descriptor, deterministic=False) (_, output), (_, dq, dk, dv) = jax.value_and_grad( loss_fn, argnums=(0, 1, 2, 3), has_aux=True )(variables, q, k, v) @@ -208,9 +198,7 @@ def test_mxfp8_attention_scale_layout(): q_layout=QuantizeLayout.ROWWISE_COLWISE, data_layout="NN", ) - tensor = quantizer.quantize( - jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2 - ) + tensor = quantizer.quantize(jnp.ones((2, 64, 8, 64), dtype=jnp.bfloat16), flatten_axis=-2) assert tensor.rowwise_tensor.scale_inv.shape == (2, 64, 8, 2) assert tensor.colwise_tensor.scale_inv.shape == (2, 2, 8, 64) assert _mxfp8_scale_inv(tensor).shape == (2, 8, 128, 4) @@ -265,9 +253,7 @@ def test_fused_attention_parity_features(feature): q = jax.random.normal(q_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 k = jax.random.normal(k_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 v = jax.random.normal(v_key, (batch, kv_seqlen, heads, dim), jnp.bfloat16) * 0.25 - doutput = ( - jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 - ) + doutput = jax.random.normal(do_key, (batch, q_seqlen, heads, dim), jnp.bfloat16) * 0.25 q_lengths = jnp.full((batch,), q_seqlen, dtype=jnp.int32) kv_lengths = jnp.full((batch,), kv_seqlen, dtype=jnp.int32) descriptor = SequenceDescriptor.from_seqlens((q_lengths, kv_lengths)) @@ -319,9 +305,7 @@ def reference_loss(query, key, value): ) return jnp.sum(output.astype(jnp.float32) * doutput.astype(jnp.float32)), output - (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)( - q, k, v - ) + (_, output), grads = jax.value_and_grad(te_loss, argnums=(0, 1, 2), has_aux=True)(q, k, v) (_, reference), reference_grads = jax.value_and_grad( reference_loss, argnums=(0, 1, 2), has_aux=True )(q, k, v) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 722eac2bb1f..bf74cbbfdc1 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -451,9 +451,7 @@ def test_fused_attn_backend_message(monkeypatch): assert backend == NVTE_Fused_Attn_Backend.NVTE_No_Backend assert message == "attention bias is not supported" - backend, message = replace( - baseline, head_dim_qk=1024, head_dim_v=1024 - ).get_fused_attn_backend() + backend, message = replace(baseline, head_dim_qk=1024, head_dim_v=1024).get_fused_attn_backend() assert backend == NVTE_Fused_Attn_Backend.NVTE_No_Backend assert message == "head dimensions are not supported" diff --git a/tests/test_common_attention_helpers.py b/tests/test_common_attention_helpers.py index 87889f1c063..8b19acbc6c4 100644 --- a/tests/test_common_attention_helpers.py +++ b/tests/test_common_attention_helpers.py @@ -128,21 +128,11 @@ def test_shared_f16_policy_framework_parity_features(): def test_shared_fp8_policy(): fp8 = _attention_config(q_dtype="float8_e4m3", kv_dtype="float8_e4m3", sm_arch=100) assert check_fp8_fused_attention_support(fp8).supported - assert check_fp8_fused_attention_support( - replace(fp8, cudnn_version=(9, 10, 1)) - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, cudnn_version=(9, 10, 0)) - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, bias_type="alibi") - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, return_max_logit=True) - ).supported - assert not check_fp8_fused_attention_support( - replace(fp8, head_dim_qk=200) - ).supported + assert check_fp8_fused_attention_support(replace(fp8, cudnn_version=(9, 10, 1))).supported + assert not check_fp8_fused_attention_support(replace(fp8, cudnn_version=(9, 10, 0))).supported + assert not check_fp8_fused_attention_support(replace(fp8, bias_type="alibi")).supported + assert not check_fp8_fused_attention_support(replace(fp8, return_max_logit=True)).supported + assert not check_fp8_fused_attention_support(replace(fp8, head_dim_qk=200)).supported def test_shared_fp8_scaling_mode_policy(): @@ -154,9 +144,7 @@ def test_shared_fp8_scaling_mode_policy(): hopper = replace(fp8, sm_arch=90) assert check_fp8_fused_attention_support(hopper, scaling_mode="delayed").supported - assert not check_fp8_fused_attention_support( - hopper, scaling_mode="current" - ).supported + assert not check_fp8_fused_attention_support(hopper, scaling_mode="current").supported assert not check_fp8_fused_attention_support(hopper, scaling_mode="mxfp8").supported assert not check_fp8_fused_attention_support( @@ -208,18 +196,12 @@ def test_shared_fp8_thd_policy(): assert check_fp8_fused_attention_support( replace(thd, mask_type="padding_causal_bottom_right") ).supported - assert not check_fp8_fused_attention_support( - replace(thd, cudnn_version=(9, 22, 9)) - ).supported - assert not check_fp8_fused_attention_support( - replace(thd, mask_type="no_mask") - ).supported + assert not check_fp8_fused_attention_support(replace(thd, cudnn_version=(9, 22, 9))).supported + assert not check_fp8_fused_attention_support(replace(thd, mask_type="no_mask")).supported assert not check_fp8_fused_attention_support( replace(thd, is_training=True, sm_arch=90) ).supported - assert not check_fp8_fused_attention_support( - replace(thd, head_dim_qk=144) - ).supported + assert not check_fp8_fused_attention_support(replace(thd, head_dim_qk=144)).supported sink_backward = replace(thd, is_training=True, softmax_type="learnable") assert not check_fp8_fused_attention_support( @@ -505,9 +487,7 @@ def test_shared_cudnn_graph_reports_build_diagnostics(): def test_shared_cudnn_graph_support_error_has_context(): with pytest.raises(RuntimeError, match="cuDNN test graph is not supported"): - build_cudnn_graph( - _FakeCudnn(), _FakeGraph(unsupported=True), description="test" - ) + build_cudnn_graph(_FakeCudnn(), _FakeGraph(unsupported=True), description="test") @pytest.fixture diff --git a/transformer_engine/common/attention/cache_debug.py b/transformer_engine/common/attention/cache_debug.py index f4f384b5221..7a0eb11522c 100644 --- a/transformer_engine/common/attention/cache_debug.py +++ b/transformer_engine/common/attention/cache_debug.py @@ -111,10 +111,7 @@ def _counter_line( if event: label += f" {event}" values = ", ".join(f"{name}={counters.get(name, 0):4d}" for name in _COUNTER_NAMES) - return ( - f"{_PREFIX} {_rank_tag()}{thread_field:<7} {device_field:<9} | " - f"{label:<24} | {values}\n" - ) + return f"{_PREFIX} {_rank_tag()}{thread_field:<7} {device_field:<9} | {label:<24} | {values}\n" def _register_summary() -> None: @@ -231,9 +228,7 @@ def render_summary() -> str: device_field = f"dev={next(iter(devices))}" else: device_field = "dev=mixed" - lines.append( - _counter_line(f"tid={thread_id}", device_field, backend, direction, values) - ) + lines.append(_counter_line(f"tid={thread_id}", device_field, backend, direction, values)) for (backend, direction), values in sorted(counters.items()): lines.append(_counter_line("tid=all", "dev=all", backend, direction, values)) for (backend, direction, stage), (calls, elapsed_ns) in sorted(stage_timings.items()): diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index 2257d4409ef..4cca6daf98b 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -191,9 +191,7 @@ def cudnn_mask_options( ) version = encode_cudnn_version(cudnn_version) options: dict[str, bool | int | str] = { - "diagonal_alignment": ( - "bottom_right" if mask.bottom_right_diagonal else "top_left" - ), + "diagonal_alignment": ("bottom_right" if mask.bottom_right_diagonal else "top_left"), "is_padding": mask.padding, } if version < 90600: @@ -356,9 +354,7 @@ def check_f16_fused_attention_support( and (dqk, dv) != (192, 128) and dqk != dv ): - return _unsupported( - "this Hopper backward head-dimension combination is unsupported" - ) + return _unsupported("this Hopper backward head-dimension combination is unsupported") alibi_supported = ( bias == "alibi" @@ -438,10 +434,7 @@ def check_f16_fused_attention_support( and bias == "no_bias" and dropout == 0.0 ) - or ( - mask in ("causal_bottom_right", "padding_causal_bottom_right") - and sq <= skv - ) + or (mask in ("causal_bottom_right", "padding_causal_bottom_right") and sq <= skv) ) mask_ok = mask_ok or modern_mask_ok if not mask_ok: @@ -473,10 +466,7 @@ def check_f16_fused_attention_support( or ( left >= -1 and right == 0 - and ( - mask in ("no_mask", "causal") - or (mask == "causal_bottom_right" and sq == skv) - ) + and (mask in ("no_mask", "causal") or (mask == "causal_bottom_right" and sq == skv)) and sq <= skv and dropout == 0.0 and bias == "no_bias" @@ -609,16 +599,12 @@ def check_fp8_fused_attention_support( if arch >= 120: return _unsupported("FP8 attention is not supported on SM120 or newer") if config.is_training and config.deterministic and version < 91900: - return _unsupported( - "deterministic FP8 attention backward requires cuDNN 9.19 or newer" - ) + return _unsupported("deterministic FP8 attention backward requires cuDNN 9.19 or newer") if scaling_mode == "current": if arch < 100: return _unsupported("FP8 current-scaling attention requires SM100 or newer") if version < 91400: - return _unsupported( - "FP8 current-scaling attention requires cuDNN 9.14 or newer" - ) + return _unsupported("FP8 current-scaling attention requires cuDNN 9.14 or newer") if scaling_mode == "mxfp8": if arch < 100: return _unsupported("MXFP8 attention requires SM100 or newer") @@ -644,9 +630,7 @@ def check_fp8_fused_attention_support( ) and version < 90500 ): - return _unsupported( - "FP8 attention requires cuDNN 9.5 for 64-bit ragged offsets" - ) + return _unsupported("FP8 attention requires cuDNN 9.5 for 64-bit ragged offsets") is_thd = layout.qkv_format == "thd" if is_thd: @@ -657,13 +641,9 @@ def check_fp8_fused_attention_support( if config.is_training and arch < 100: return _unsupported("FP8 THD attention backward requires SM100 or newer") if config.is_training and config.softmax_type != "vanilla" and version < 92600: - return _unsupported( - "FP8 THD sink-token backward requires cuDNN 9.26 or newer" - ) + return _unsupported("FP8 THD sink-token backward requires cuDNN 9.26 or newer") if arch >= 100 and (dqk > 128 or dv > 128): - return _unsupported( - "FP8 THD attention supports head dimensions up to 128 on SM100+" - ) + return _unsupported("FP8 THD attention supports head dimensions up to 128 on SM100+") shape_mask_ok = ( ( diff --git a/transformer_engine/common/cudnn_frontend.py b/transformer_engine/common/cudnn_frontend.py index 634fe986588..c735025b5dc 100644 --- a/transformer_engine/common/cudnn_frontend.py +++ b/transformer_engine/common/cudnn_frontend.py @@ -61,9 +61,7 @@ def build_cudnn_graph( time_call( debug_callback, "create_execution_plans", - lambda: graph.create_execution_plans( - [cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK] - ), + lambda: graph.create_execution_plans([cudnn.heur_mode.A, cudnn.heur_mode.FALLBACK]), ) time_call(debug_callback, "check_support", graph.check_support) except cudnn.cudnnGraphNotSupportedError as exc: diff --git a/transformer_engine/jax/cpp_extensions/flex_attention.py b/transformer_engine/jax/cpp_extensions/flex_attention.py index 15ccde37b4f..b1700c27514 100644 --- a/transformer_engine/jax/cpp_extensions/flex_attention.py +++ b/transformer_engine/jax/cpp_extensions/flex_attention.py @@ -393,9 +393,7 @@ def wrapped_score_mod(sdpa_graph, score_tensor): return wrapped_score_mod -def _finalize_score_mod_graph( - cudnn, graph, cache_site: Tuple[str, str] -) -> Tuple[int, bytes, int]: +def _finalize_score_mod_graph(cudnn, graph, cache_site: Tuple[str, str]) -> Tuple[int, bytes, int]: return finalize_graph( cudnn, graph, diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index 72947f079b2..c8f2ddff656 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -201,9 +201,7 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="Q", dim=(info.input_batch, info.q_heads, info.q_max_seqlen, info.qk_dim), - stride=_matrix_stride( - info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen - ), + stride=_matrix_stride(info, config.qkv_layout, "q", info.q_max_seqlen, info.kv_max_seqlen), dtype=io_dtype, uid=_UID_Q, ) @@ -211,9 +209,7 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="K", dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.qk_dim), - stride=_matrix_stride( - info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen - ), + stride=_matrix_stride(info, config.qkv_layout, "k", info.q_max_seqlen, info.kv_max_seqlen), dtype=io_dtype, uid=_UID_K, ) @@ -221,9 +217,7 @@ def _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config): graph, name="V", dim=(info.input_batch, info.kv_heads, info.kv_max_seqlen, info.v_dim), - stride=_matrix_stride( - info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen - ), + stride=_matrix_stride(info, config.qkv_layout, "v", info.q_max_seqlen, info.kv_max_seqlen), dtype=io_dtype, uid=_UID_V, ) @@ -268,9 +262,7 @@ def build_fp8_fwd_graph( graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) info, _, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) - input_bindings = list( - _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) - ) + input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) tensors = {"q": q, "k": k, "v": v} options, is_padding = _fp8_options(cudnn, info, config) options["generate_stats"] = True @@ -284,9 +276,7 @@ def build_fp8_fwd_graph( options["diagonal_band_left_bound"] = options.pop("left_bound") if "right_bound" in options: options["diagonal_band_right_bound"] = options.pop("right_bound") - padded = mxfp8_padded_sizes( - info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim - ) + padded = mxfp8_padded_sizes(info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim) scale_specs = ( ( "descale_q", @@ -355,9 +345,7 @@ def build_fp8_fwd_graph( uid=_UID_SEQ_KV, ) options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend( - (GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11)) - ) + input_bindings.extend((GraphBinding(_UID_SEQ_Q, 10), GraphBinding(_UID_SEQ_KV, 11))) output_bindings = [GraphBinding(_UID_O, 0), GraphBinding(_UID_STATS, 1)] if _is_dropout(config): @@ -394,27 +382,23 @@ def build_fp8_fwd_graph( output = op["output"] output.set_output(True).set_uid(_UID_O).set_data_type( cudnn_data_type(cudnn, output_dtype) - ).set_dim( - (info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim) - ).set_stride( + ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim)).set_stride( attention_format_stride( info.input_batch, info.q_heads, info.q_max_seqlen, info.v_dim, "bshd" ) ) stats = op["stats"] - stats.set_output(True).set_uid(_UID_STATS).set_data_type( - cudnn.data_type.FLOAT - ).set_dim((info.input_batch, info.q_heads, info.q_max_seqlen, 1)).set_stride( - (info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1) - ) + stats.set_output(True).set_uid(_UID_STATS).set_data_type(cudnn.data_type.FLOAT).set_dim( + (info.input_batch, info.q_heads, info.q_max_seqlen, 1) + ).set_stride((info.q_heads * info.q_max_seqlen, info.q_max_seqlen, 1, 1)) if mode != "mxfp8": for name, uid, offset in ( ("amax_s", _UID_AMAX_S, 0), ("amax_o", _UID_AMAX_O, 4), ): - op[name].set_output(True).set_uid(uid).set_data_type( - cudnn.data_type.FLOAT - ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) + op[name].set_output(True).set_uid(uid).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 2, offset)) else: op["amax_o"].set_output(False).set_data_type(cudnn.data_type.FLOAT).set_dim( @@ -539,13 +523,9 @@ def build_fp8_bwd_graph( cudnn = import_cudnn() graph = make_graph(cudnn, cudnn_data_type(cudnn, q_aval.dtype)) - info, io_dtype, q, k, v = _graph_io_tensors( - graph, cudnn, q_aval, k_aval, v_aval, config - ) + info, io_dtype, q, k, v = _graph_io_tensors(graph, cudnn, q_aval, k_aval, v_aval, config) q_shape, k_shape, v_shape, o_shape = _logical_shapes(info) - input_bindings = list( - _qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize) - ) + input_bindings = list(_qkv_bindings(info, config.qkv_layout, jnp.dtype(q_aval.dtype).itemsize)) o = _tensor( graph, name="O", @@ -600,9 +580,7 @@ def build_fp8_bwd_graph( uid=_UID_SEQ_KV, ) options.update(seq_len_q=seq_q, seq_len_kv=seq_kv) - input_bindings.extend( - (GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10)) - ) + input_bindings.extend((GraphBinding(_UID_SEQ_Q, 9), GraphBinding(_UID_SEQ_KV, 10))) if _is_dropout(config): seed = _tensor( @@ -681,9 +659,7 @@ def build_fp8_bwd_graph( GraphBinding(_UID_DO_F16, 27), ) ) - padded = mxfp8_padded_sizes( - info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim - ) + padded = mxfp8_padded_sizes(info.q_max_seqlen, info.kv_max_seqlen, info.qk_dim, info.v_dim) tensors.update(_mx_bwd_scales(graph, cudnn, info, padded, input_bindings)) else: for name, uid, index in ( @@ -731,9 +707,9 @@ def build_fp8_bwd_graph( ("amax_dv", _UID_AMAX_DV, 8), ("amax_dp", _UID_AMAX_DP, 12), ): - op[name].set_output(True).set_uid(uid).set_data_type( - cudnn.data_type.FLOAT - ).set_dim((1, 1, 1, 1)).set_stride((1, 1, 1, 1)) + op[name].set_output(True).set_uid(uid).set_data_type(cudnn.data_type.FLOAT).set_dim( + (1, 1, 1, 1) + ).set_stride((1, 1, 1, 1)) output_bindings.append(GraphBinding(uid, 3, offset)) else: for amax in op["amax"]: @@ -951,9 +927,7 @@ def _validate_quantizer_modes(quantizers) -> str: if actual[role] != expected[role] ] if mismatches: - raise ValueError( - "Unsupported FP8 attention quantizer modes: " + ", ".join(mismatches) - ) + raise ValueError("Unsupported FP8 attention quantizer modes: " + ", ".join(mismatches)) return mode @@ -1097,18 +1071,14 @@ def _validate_fp8_support(qkv, quantizers, config, mode): window_size=tuple(int(value) for value in config.window_size), return_max_logit=False, cuda_graph=False, - deterministic=not bool( - int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")) - ), + deterministic=not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1"))), cudnn_version=cudnn_version, sm_arch=device_arch, ), scaling_mode=mode, ) if not support.supported: - raise ValueError( - f"Unsupported JAX FP8 attention configuration: {support.reason}." - ) + raise ValueError(f"Unsupported JAX FP8 attention configuration: {support.reason}.") def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): @@ -1118,15 +1088,11 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): if config.qkv_layout.is_thd(): raise NotImplementedError("FP8 attention does not support THD layouts in JAX.") if mode == "mxfp8" and not config.qkv_layout.is_separate(): - raise NotImplementedError( - "JAX MXFP8 attention currently requires separate BSHD Q/K/V." - ) + raise NotImplementedError("JAX MXFP8 attention currently requires separate BSHD Q/K/V.") if getattr(config.attn_bias_type, "name", "") != "NO_BIAS": raise NotImplementedError("FP8 attention does not support attention bias.") if getattr(config.softmax_type, "name", "") != "VANILLA_SOFTMAX": - raise NotImplementedError( - "JAX FP8 attention currently supports vanilla softmax only." - ) + raise NotImplementedError("JAX FP8 attention currently supports vanilla softmax only.") _validate_fp8_support(qkv, quantizers, config, mode) if mode == "mxfp8" and _is_padding(config): raise NotImplementedError("JAX MXFP8 attention does not support padding masks.") @@ -1169,9 +1135,7 @@ def fused_attn_fp8_fwd(qkv, sequence_descriptor, seed, quantizers, config): quantizers.s.update(amax[0:1]) if output_is_delayed: quantizers.o.update(amax[1:2]) - output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype( - qkv[0].dtype - ) + output = (raw_output.astype(qkv[0].dtype) * output_scale_inv).astype(qkv[0].dtype) else: output = raw_output return output, ( diff --git a/transformer_engine/jax/csrc/extensions/attention.cpp b/transformer_engine/jax/csrc/extensions/attention.cpp index 5d842a8d06c..da6989bcabb 100644 --- a/transformer_engine/jax/csrc/extensions/attention.cpp +++ b/transformer_engine/jax/csrc/extensions/attention.cpp @@ -202,9 +202,9 @@ Error_Type ExecuteCudnnGraph(cudaStream_t stream, Dictionary &attrs, NVTE_CHECK_CUDNN(cudnnSetStream(handle, stream)); int device_id = 0; NVTE_CHECK_CUDA(cudaGetDevice(&device_id)); - attention_cache_debug::Record( - get_attr_value(attrs, "attention_backend"), - get_attr_value(attrs, "attention_direction"), "execute", device_id); + attention_cache_debug::Record(get_attr_value(attrs, "attention_backend"), + get_attr_value(attrs, "attention_direction"), + "execute", device_id); auto status = graph->execute(handle, variant_pack, workspace); NVTE_CHECK(status.is_good(), "cuDNN frontend graph execution failed: ", status.get_message()); return ffi_with_cuda_error_check(); diff --git a/transformer_engine/jax/csrc/extensions/attention_cache_debug.h b/transformer_engine/jax/csrc/extensions/attention_cache_debug.h index 757592fd3ca..a0b71e34c85 100644 --- a/transformer_engine/jax/csrc/extensions/attention_cache_debug.h +++ b/transformer_engine/jax/csrc/extensions/attention_cache_debug.h @@ -27,8 +27,7 @@ namespace detail { constexpr size_t kSiteCount = 4; constexpr size_t kStageCount = 5; inline constexpr std::array kStageNames = { - "validate", "build_operation_graph", "create_execution_plans", "check_support", - "build_plans"}; + "validate", "build_operation_graph", "create_execution_plans", "check_support", "build_plans"}; inline int DebugLevel() { static const int level = [] { @@ -137,15 +136,14 @@ inline std::string CounterLine(size_t site, const char *event = nullptr, int dev const Counters &counters = AllCounters()[site]; const std::string device_name = all_devices ? "all" : std::to_string(device); char line[640]; - std::snprintf( - line, sizeof(line), - "[FUSED-ATTN-CACHE] %sdev=%-3s | %s %s %-12s | hit=%4" PRIu64 - ", miss=%4" PRIu64 ", create_graph=%4" PRIu64 ", cache_graph=%4" PRIu64 - ", build_plans=%4" PRIu64 ", execute=%4" PRIu64 "\n", - RankTag().c_str(), device_name.c_str(), Backend(site), Direction(site), - event == nullptr ? "" : event, - Load(counters.hit), Load(counters.miss), Load(counters.create_graph), - Load(counters.cache_graph), Load(counters.build_plans), Load(counters.execute)); + std::snprintf(line, sizeof(line), + "[FUSED-ATTN-CACHE] %sdev=%-3s | %s %s %-12s | hit=%4" PRIu64 ", miss=%4" PRIu64 + ", create_graph=%4" PRIu64 ", cache_graph=%4" PRIu64 ", build_plans=%4" PRIu64 + ", execute=%4" PRIu64 "\n", + RankTag().c_str(), device_name.c_str(), Backend(site), Direction(site), + event == nullptr ? "" : event, Load(counters.hit), Load(counters.miss), + Load(counters.create_graph), Load(counters.cache_graph), Load(counters.build_plans), + Load(counters.execute)); return line; } @@ -168,8 +166,7 @@ inline void PrintSummary() { const double milliseconds = static_cast(Load(timing.elapsed_ns)) / calls / 1e6; char line[320]; std::snprintf(line, sizeof(line), - "[FUSED-ATTN-CACHE] %s%s %-3s %-22s | calls=%" PRIu64 - " | time=%9.3f ms/call\n", + "[FUSED-ATTN-CACHE] %s%s %-3s %-22s | calls=%" PRIu64 " | time=%9.3f ms/call\n", RankTag().c_str(), Backend(site), Direction(site), kStageNames[stage], calls, milliseconds); output += line; @@ -228,10 +225,9 @@ inline void Record(std::string_view backend, std::string_view direction, std::st counter->fetch_add(1, std::memory_order_relaxed); if (!detail::TraceEnabled()) return; if ((event == "hit" || event == "miss") && !key.empty()) { - detail::Write("[FUSED-ATTN-CACHE] " + detail::RankTag() + "dev=" + - std::to_string(device) + " | " + std::string(backend) + " " + - std::string(direction) + " " + std::string(event) + " | " + std::string(key) + - "\n"); + detail::Write("[FUSED-ATTN-CACHE] " + detail::RankTag() + "dev=" + std::to_string(device) + + " | " + std::string(backend) + " " + std::string(direction) + " " + + std::string(event) + " | " + std::string(key) + "\n"); } else { std::string uppercase(event == "plans_built" ? "build_plans" : event); for (char &character : uppercase) { diff --git a/transformer_engine/jax/csrc/extensions/pybind.cpp b/transformer_engine/jax/csrc/extensions/pybind.cpp index 0ba8755207f..a5820f4b3b9 100644 --- a/transformer_engine/jax/csrc/extensions/pybind.cpp +++ b/transformer_engine/jax/csrc/extensions/pybind.cpp @@ -145,8 +145,7 @@ PYBIND11_MODULE(transformer_engine_jax, m) { attention_cache_debug::Record(backend, direction, event, device, key, elapsed_ns); }, pybind11::arg("backend"), pybind11::arg("direction"), pybind11::arg("event"), - pybind11::arg("device") = -1, pybind11::arg("key") = "", - pybind11::arg("elapsed_ns") = 0); + pybind11::arg("device") = -1, pybind11::arg("key") = "", pybind11::arg("elapsed_ns") = 0); m.def("get_device_compute_capability", &GetDeviceComputeCapability); m.def("get_num_compute_streams", &nvte_get_num_compute_streams); m.def("get_cublasLt_version", &cublasLtGetVersion); diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 45eeb7f5427..77cf1f703c5 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -1132,14 +1132,10 @@ def _build_fp8_fwd_graph( use_legacy_offsets = is_ragged_q or is_ragged_kv graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch graph_seqlen_q = ( - _max_ragged_tokens(q_data.shape[0]) - if is_ragged_q and use_token_buckets - else max_seqlen_q + _max_ragged_tokens(q_data.shape[0]) if is_ragged_q and use_token_buckets else max_seqlen_q ) graph_seqlen_kv = ( - _max_ragged_tokens(k_data.shape[0]) - if is_ragged_kv and use_token_buckets - else max_seqlen_kv + _max_ragged_tokens(k_data.shape[0]) if is_ragged_kv and use_token_buckets else max_seqlen_kv ) graph = make_graph(_fp8_cudnn_dtype(q), q.device, name="te_fp8_sdpa_fwd") tensors: Dict[str, Any] = { @@ -1532,9 +1528,7 @@ def _fp8_forward( rng_elts = (max_seqlen_q * max_seqlen_q + _FP8_THREADS_PER_CTA - 1) // _FP8_THREADS_PER_CTA rng_state = _reserve_philox_state(q.device, rng_gen, rng_elts) cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded - cu_seqlens_kv_padded = ( - cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded - ) + cu_seqlens_kv_padded = cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded kv_offsets_source = _kv_ragged_offsets_source( qkv_layout, cu_seqlens_q_padded, @@ -2120,14 +2114,10 @@ def _build_fp8_bwd_graph( use_legacy_offsets = is_ragged_q or is_ragged_kv graph_batch = _max_ragged_batch(batch) if use_legacy_offsets and use_token_buckets else batch graph_seqlen_q = ( - _max_ragged_tokens(q_data.shape[0]) - if is_ragged_q and use_token_buckets - else max_seqlen_q + _max_ragged_tokens(q_data.shape[0]) if is_ragged_q and use_token_buckets else max_seqlen_q ) graph_seqlen_kv = ( - _max_ragged_tokens(k_data.shape[0]) - if is_ragged_kv and use_token_buckets - else max_seqlen_kv + _max_ragged_tokens(k_data.shape[0]) if is_ragged_kv and use_token_buckets else max_seqlen_kv ) graph = make_graph(_fp8_cudnn_dtype(q), q.device, name="te_fp8_sdpa_bwd") tensors: Dict[str, Any] = { @@ -2646,9 +2636,7 @@ def _fp8_backward( d_softmax_offset = torch.empty_like(softmax_offset) if softmax_offset is not None else None hidden_amax = [torch.zeros(1, dtype=torch.float32, device=q.device) for _ in range(4)] cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded - cu_seqlens_kv_padded = ( - cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded - ) + cu_seqlens_kv_padded = cu_seqlens_kv if cu_seqlens_kv_padded is None else cu_seqlens_kv_padded kv_offsets_source = _kv_ragged_offsets_source( qkv_layout, cu_seqlens_q_padded, From a5248dc8b523bfa13fddc2fa7f24885d35e02bd3 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Sat, 12 Sep 2026 19:20:08 +0000 Subject: [PATCH 28/36] Fix FlexAttention execution mock signature Signed-off-by: Vladimir Cherepanov --- tests/pytorch/attention/test_flex_attention.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/pytorch/attention/test_flex_attention.py b/tests/pytorch/attention/test_flex_attention.py index 55b8b3b9aa6..3d5c9731c4d 100644 --- a/tests/pytorch/attention/test_flex_attention.py +++ b/tests/pytorch/attention/test_flex_attention.py @@ -446,8 +446,8 @@ class FakeEntry: score_mod_graph_tensors = {"softcap": object()} workspace_size = 1 - def fake_execute(graph, variant_pack, workspace_size, device): - del graph, variant_pack, workspace_size, device + def fake_execute(graph, variant_pack, workspace_size, device, cache_site): + del graph, variant_pack, workspace_size, device, cache_site q, k, v, _, _ = _score_mod_cache_cpu_inputs() q = q.requires_grad_() From a497b13446d07c3c443c2be333276f14c14ecbaa Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Sat, 12 Sep 2026 19:31:35 +0000 Subject: [PATCH 29/36] Fix FP8 attention checkpoint test opt-in Signed-off-by: Vladimir Cherepanov --- tests/pytorch/attention/test_attention.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index 0812aeecf55..8897f717568 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -41,6 +41,7 @@ fused_attn_fwd, ) from transformer_engine.pytorch.distributed import CudaRNGStatesTracker +from transformer_engine.pytorch._extra_state import UNSAFE_PICKLE_EXTRA_STATE_ENV from transformer_engine.pytorch.module.base import TransformerEngineBaseModule from transformer_engine.pytorch.utils import ( init_method_normal, @@ -2406,7 +2407,7 @@ def _run_transformer_layer( @pytest.mark.skipif(get_cudnn_version() < (9, 3, 0), reason="cuDNN 9.3.0+ is required.") @pytest.mark.parametrize("model", ["large"]) @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) -def test_dpa_fp8_extra_state(model, dtype): +def test_dpa_fp8_extra_state(model, dtype, monkeypatch): """Test DotProductAttention module in FP8 with checkpointing""" config = model_configs_fp8_extra_state[model] # Test backend availability @@ -2423,6 +2424,8 @@ def test_dpa_fp8_extra_state(model, dtype): pytest.skip("No attention backend available.") outputs = _run_dpa_fp8_extra_state(dtype, config, checkpoint=False) + # The checkpoints are generated locally by this test and are therefore trusted. + monkeypatch.setenv(UNSAFE_PICKLE_EXTRA_STATE_ENV, "1") outputs_checkpoint = _run_dpa_fp8_extra_state(dtype, config, checkpoint=True) outputs_checkpoint_v1_6 = _run_dpa_fp8_extra_state( dtype, config, mimic_v1_6=True, checkpoint=True From 848cd072989ed31877fb945fb0d04ae19c302ed7 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Sat, 12 Sep 2026 20:23:18 +0000 Subject: [PATCH 30/36] Avoid retaining cuDNN graph workspaces Signed-off-by: Vladimir Cherepanov --- tests/pytorch/attention/test_cudnn_graph.py | 49 +++++++++++++++++++ .../dot_product_attention/_cudnn_graph.py | 38 ++++---------- 2 files changed, 59 insertions(+), 28 deletions(-) create mode 100644 tests/pytorch/attention/test_cudnn_graph.py diff --git a/tests/pytorch/attention/test_cudnn_graph.py b/tests/pytorch/attention/test_cudnn_graph.py new file mode 100644 index 00000000000..ddfd4717a1b --- /dev/null +++ b/tests/pytorch/attention/test_cudnn_graph.py @@ -0,0 +1,49 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Tests for the PyTorch cuDNN graph runtime.""" + +import weakref + +import torch + +from transformer_engine.pytorch.attention.dot_product_attention import _cudnn_graph + + +def test_graph_entry_does_not_retain_execution_workspaces(monkeypatch): + """Cached graph entries should not own per-execution scratch allocations.""" + + class Workspace: + pass + + class Graph: + def __init__(self): + self.workspace_refs = [] + + def execute(self, variant_pack, workspace, handle): + assert variant_pack == {"q": "tensor"} + assert handle == "handle" + self.workspace_refs.append(weakref.ref(workspace)) + + allocated_workspace_refs = [] + + def allocate_workspace(size, *, dtype, device): + assert size == 123 + assert dtype == torch.uint8 + assert device == torch.device("cpu") + workspace = Workspace() + allocated_workspace_refs.append(weakref.ref(workspace)) + return workspace + + monkeypatch.setattr(_cudnn_graph.torch, "empty", allocate_workspace) + monkeypatch.setattr(_cudnn_graph, "current_stream_handle", lambda _device: "handle") + + graph = Graph() + entry = _cudnn_graph.GraphEntry(graph=graph, tensors={}, workspace_size=123) + entry.execute({"q": "tensor"}, torch.device("cpu")) + entry.execute({"q": "tensor"}, torch.device("cpu")) + + assert len(allocated_workspace_refs) == 2 + assert all(ref() is None for ref in allocated_workspace_refs) + assert all(ref() is None for ref in graph.workspace_refs) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py index 718493c9667..721e0ccbba7 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py @@ -7,7 +7,7 @@ from __future__ import annotations import threading -from dataclasses import dataclass, field +from dataclasses import dataclass from typing import Any, Dict, Hashable, Optional, Tuple import torch @@ -119,46 +119,28 @@ def finalize_graph(graph, *, cache_site: Tuple[str, str]) -> int: @dataclass class GraphEntry: - """Built graph plus named graph tensors and stream-local workspaces.""" + """Built graph plus named graph tensors.""" graph: Any tensors: Dict[str, Any] workspace_size: int cache_site: Optional[Tuple[str, str]] = None - _workspaces: Dict[int, torch.Tensor] = field(default_factory=dict, repr=False) - - def workspace(self, device: torch.device) -> torch.Tensor: - """Get a stable workspace for the current stream. - - A workspace may be reused by asynchronous launches on one stream, but - must not be shared by independent streams. CUDA graph capture also - keeps the allocation alive for CUDA graph replay. PyTorch's caching - allocator is capture-aware, so a stream-specific workspace may be - created on the first captured invocation after the graph itself has - been warmed and cached. - """ - - device = torch.device(device) - stream = torch.cuda.current_stream(device) - stream_key = int(stream.cuda_stream) - workspace = self._workspaces.get(stream_key) - if workspace is None: - workspace = torch.empty( - self.workspace_size, - dtype=torch.uint8, - device=device, - ) - self._workspaces[stream_key] = workspace - return workspace def execute(self, variant_pack: Dict[Any, Any], device: torch.device) -> None: """Execute the graph on PyTorch's current stream.""" if self.cache_site is not None: record_event(*self.cache_site, "execute", device=_device_key(device)[1]) + # Workspaces are execution scratch. Keeping them in graph cache entries + # retains one potentially large allocation for every cached configuration. + workspace = torch.empty( + self.workspace_size, + dtype=torch.uint8, + device=device, + ) self.graph.execute( variant_pack, - self.workspace(device), + workspace, handle=current_stream_handle(device), ) From a12ab7d53b98500c11e0a7ee3c0f9c0ad305dc12 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Sat, 12 Sep 2026 20:39:29 +0000 Subject: [PATCH 31/36] Fix paged attention table descriptors Signed-off-by: Vladimir Cherepanov --- tests/pytorch/attention/test_cudnn_graph.py | 26 ++++++++++++++++++- .../dot_product_attention/cudnn_attention.py | 26 +++++++++++++++++-- 2 files changed, 49 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/attention/test_cudnn_graph.py b/tests/pytorch/attention/test_cudnn_graph.py index ddfd4717a1b..f91507f5b37 100644 --- a/tests/pytorch/attention/test_cudnn_graph.py +++ b/tests/pytorch/attention/test_cudnn_graph.py @@ -8,7 +8,31 @@ import torch -from transformer_engine.pytorch.attention.dot_product_attention import _cudnn_graph +from transformer_engine.pytorch.attention.dot_product_attention import ( + _cudnn_graph, + cudnn_attention, +) + + +def test_page_table_uses_cudnn_logical_layout(): + """cuDNN paged attention requires a four-dimensional table descriptor.""" + + class Graph: + @staticmethod + def tensor(**kwargs): + return kwargs + + page_table = torch.empty_strided((3, 4), (6, 1), dtype=torch.int32) + graph_tensor = cudnn_attention._make_page_table_graph_tensor( + Graph(), page_table, batch=2, name="page_table_k" + ) + + assert graph_tensor == { + "name": "page_table_k", + "dim": (2, 1, 4, 1), + "stride": (6, 6, 1, 1), + "data_type": torch.int32, + } def test_graph_entry_does_not_retain_execution_workspaces(monkeypatch): diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 77cf1f703c5..62e5c0d2dd8 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -256,6 +256,24 @@ def _is_paged_layout(qkv_layout: str) -> bool: return qkv_layout.startswith("paged_kv_") +def _make_page_table_graph_tensor( + graph, page_table: torch.Tensor, *, batch: int, name: str +): + """Describe a physical ``[batch, pages]`` table in cuDNN's logical layout.""" + + if page_table.ndim != 2: + raise ValueError( + f"Paged attention expects a 2D page table, got shape {tuple(page_table.shape)}." + ) + batch_stride, page_stride = page_table.stride() + return graph.tensor( + name=name, + dim=(batch, 1, page_table.shape[1], 1), + stride=(batch_stride, batch_stride, page_stride, page_stride), + data_type=page_table.dtype, + ) + + def _tensor_metadata(tensor: Optional[torch.Tensor]) -> Optional[Tuple[Any, ...]]: if tensor is None: return None @@ -660,8 +678,12 @@ def _build_f16_fwd_graph( options["seq_len_kv"] = seq_kv_t if page_table_k is not None: - page_k_t = graph.tensor_like(page_table_k, name="page_table_k") - page_v_t = graph.tensor_like(page_table_v, name="page_table_v") + page_k_t = _make_page_table_graph_tensor( + graph, page_table_k, batch=graph_batch, name="page_table_k" + ) + page_v_t = _make_page_table_graph_tensor( + graph, page_table_v, batch=graph_batch, name="page_table_v" + ) tensors["page_table_k"] = page_k_t tensors["page_table_v"] = page_v_t options["paged_attention_k_table"] = page_k_t From d4e3915f92e34ee7dd3b78dcc40b04d0e0000696 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sat, 12 Sep 2026 20:41:07 +0000 Subject: [PATCH 32/36] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- .../attention/dot_product_attention/cudnn_attention.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index 62e5c0d2dd8..bf3ad39e270 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -256,9 +256,7 @@ def _is_paged_layout(qkv_layout: str) -> bool: return qkv_layout.startswith("paged_kv_") -def _make_page_table_graph_tensor( - graph, page_table: torch.Tensor, *, batch: int, name: str -): +def _make_page_table_graph_tensor(graph, page_table: torch.Tensor, *, batch: int, name: str): """Describe a physical ``[batch, pages]`` table in cuDNN's logical layout.""" if page_table.ndim != 2: From 8e9b0ac35272897c6060670060f7cafa1ea09d48 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Mon, 14 Sep 2026 19:36:38 +0000 Subject: [PATCH 33/36] Fix cuDNN graph execution device context Signed-off-by: Vladimir Cherepanov --- tests/pytorch/attention/test_cudnn_graph.py | 18 ++++++++++++++++++ .../dot_product_attention/_cudnn_graph.py | 11 ++++++----- 2 files changed, 24 insertions(+), 5 deletions(-) diff --git a/tests/pytorch/attention/test_cudnn_graph.py b/tests/pytorch/attention/test_cudnn_graph.py index f91507f5b37..599b6cc8ce7 100644 --- a/tests/pytorch/attention/test_cudnn_graph.py +++ b/tests/pytorch/attention/test_cudnn_graph.py @@ -38,6 +38,20 @@ def tensor(**kwargs): def test_graph_entry_does_not_retain_execution_workspaces(monkeypatch): """Cached graph entries should not own per-execution scratch allocations.""" + active_devices = [] + entered_devices = [] + + class DeviceGuard: + def __init__(self, device): + self.device = device + + def __enter__(self): + active_devices.append(self.device) + entered_devices.append(self.device) + + def __exit__(self, *_args): + active_devices.pop() + class Workspace: pass @@ -48,6 +62,7 @@ def __init__(self): def execute(self, variant_pack, workspace, handle): assert variant_pack == {"q": "tensor"} assert handle == "handle" + assert active_devices == [torch.device("cpu")] self.workspace_refs.append(weakref.ref(workspace)) allocated_workspace_refs = [] @@ -61,6 +76,7 @@ def allocate_workspace(size, *, dtype, device): return workspace monkeypatch.setattr(_cudnn_graph.torch, "empty", allocate_workspace) + monkeypatch.setattr(_cudnn_graph.torch.cuda, "device", DeviceGuard) monkeypatch.setattr(_cudnn_graph, "current_stream_handle", lambda _device: "handle") graph = Graph() @@ -68,6 +84,8 @@ def allocate_workspace(size, *, dtype, device): entry.execute({"q": "tensor"}, torch.device("cpu")) entry.execute({"q": "tensor"}, torch.device("cpu")) + assert entered_devices == [torch.device("cpu"), torch.device("cpu")] + assert not active_devices assert len(allocated_workspace_refs) == 2 assert all(ref() is None for ref in allocated_workspace_refs) assert all(ref() is None for ref in graph.workspace_refs) diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py index 721e0ccbba7..2c8d38f71d8 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py @@ -138,11 +138,12 @@ def execute(self, variant_pack: Dict[Any, Any], device: torch.device) -> None: dtype=torch.uint8, device=device, ) - self.graph.execute( - variant_pack, - workspace, - handle=current_stream_handle(device), - ) + with torch.cuda.device(device): + self.graph.execute( + variant_pack, + workspace, + handle=current_stream_handle(device), + ) def graph_cache() -> Dict[Hashable, GraphEntry]: From 1969670812f30e56c607fb15a4f3723daba8023b Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Mon, 14 Sep 2026 19:58:27 +0000 Subject: [PATCH 34/36] Fix attention Python lint findings Signed-off-by: Vladimir Cherepanov --- transformer_engine/common/attention/cudnn.py | 46 ++++++++----------- .../dot_product_attention/_cudnn_graph.py | 3 +- .../dot_product_attention/context_parallel.py | 2 +- .../dot_product_attention/cudnn_attention.py | 16 ++----- .../pytorch/attention/fused_mla_q_uproj.py | 4 +- 5 files changed, 29 insertions(+), 42 deletions(-) diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index 4cca6daf98b..23963b80a81 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -345,7 +345,7 @@ def check_f16_fused_attention_support( or blackwell_d256_bwd ): return _unsupported("head dimensions are not supported") - if ( + unsupported_hopper_bwd_dims = ( version >= 91100 and is_training and arch == 90 @@ -353,8 +353,11 @@ def check_f16_fused_attention_support( and dv >= 128 and (dqk, dv) != (192, 128) and dqk != dv - ): - return _unsupported("this Hopper backward head-dimension combination is unsupported") + ) + if unsupported_hopper_bwd_dims: + return _unsupported( + "this Hopper backward head-dimension combination is unsupported" + ) alibi_supported = ( bias == "alibi" @@ -376,6 +379,8 @@ def check_f16_fused_attention_support( standard_format = layout.qkv_format in ("sbhd", "bshd") basic_masks = mask in ("no_mask", "causal", "padding", "padding_causal") + aligned_bottom_right = sq % 64 == 0 and skv % 64 == 0 and sq <= skv + no_bias_no_dropout = bias == "no_bias" and dropout == 0.0 mask_ok = version < 8906 and mask == "causal" if version >= 8906 and standard_format and basic_masks: mask_ok = True @@ -393,37 +398,25 @@ def check_f16_fused_attention_support( version >= 90300 and standard_format and mask == "causal_bottom_right" - and sq % 64 == 0 - and skv % 64 == 0 - and sq <= skv - and bias == "no_bias" - and dropout == 0.0 + and aligned_bottom_right + and no_bias_no_dropout ): mask_ok = True + paged_mask_supported = mask in ("padding", "padding_causal") or ( + mask == "padding_causal_bottom_right" and aligned_bottom_right + ) if ( version >= 90500 and layout.layout_group == "paged_separate" - and ( - mask in ("padding", "padding_causal") - or ( - mask == "padding_causal_bottom_right" - and sq % 64 == 0 - and skv % 64 == 0 - and sq <= skv - ) - ) - and bias == "no_bias" - and dropout == 0.0 + and paged_mask_supported + and no_bias_no_dropout ): mask_ok = True if ( version >= 90600 and mask == "padding_causal_bottom_right" - and sq % 64 == 0 - and skv % 64 == 0 - and sq <= skv - and bias == "no_bias" - and dropout == 0.0 + and aligned_bottom_right + and no_bias_no_dropout ): mask_ok = True if version >= 90700: @@ -539,14 +532,15 @@ def check_f16_fused_attention_support( "cuDNN 9.14.0 does not support this non-causal sliding window", "This non-causal sliding-window configuration requires cuDNN > 9.14.0", ) - if ( + unsupported_cuda_graph_bwd = ( version <= 91500 and is_training and standard_format and skv % 128 != 0 and config.cuda_graph and mask not in ("padding", "padding_causal", "padding_causal_bottom_right") - ): + ) + if unsupported_cuda_graph_bwd: return _unsupported( "this backward CUDA-graph configuration requires cuDNN 9.15.1", "This backward CUDA-graph configuration requires cuDNN 9.15.1 or newer", diff --git a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py index 2c8d38f71d8..6d34b70042f 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/_cudnn_graph.py @@ -7,8 +7,9 @@ from __future__ import annotations import threading +from collections.abc import Hashable from dataclasses import dataclass -from typing import Any, Dict, Hashable, Optional, Tuple +from typing import Any, Dict, Optional, Tuple import torch diff --git a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py index 95c265fe0d1..e5ddebdac66 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py @@ -1121,7 +1121,7 @@ def cp_p2p_fwd_fused_attn( attn_bias = rest[0] if len(rest) > 0 else None if return_max_logit: - return out_per_step, softmax_lse_per_step, rng_states, attn_bias, *max_logit + return out_per_step, softmax_lse_per_step, rng_states, attn_bias, max_logit[0] return out_per_step, softmax_lse_per_step, rng_states, attn_bias, None diff --git a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py index bf3ad39e270..492cf4c1f91 100644 --- a/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py +++ b/transformer_engine/pytorch/attention/dot_product_attention/cudnn_attention.py @@ -92,6 +92,8 @@ class FusedAttnBackend(IntEnum): @classmethod def cast(cls, backend: Union["FusedAttnBackend", int, Any]) -> "FusedAttnBackend": + """Convert a legacy backend integer to a fused-attention backend.""" + if isinstance(backend, cls): return backend return cls(int(backend)) @@ -515,7 +517,6 @@ def _build_f16_fwd_graph( cu_seqlens_kv_padded: torch.Tensor, page_table_k: Optional[torch.Tensor], page_table_v: Optional[torch.Tensor], - rng_state: torch.Tensor, softmax_offset: Optional[torch.Tensor], attn_scale: float, dropout: float, @@ -899,7 +900,6 @@ def _f16_forward( cu_seqlens_kv_padded=cu_seqlens_kv_padded, page_table_k=page_table_k, page_table_v=page_table_v, - rng_state=rng_state, softmax_offset=softmax_offset, attn_scale=attn_scale, dropout=dropout, @@ -1112,9 +1112,6 @@ def _build_fp8_fwd_graph( k, v, output, - stats, - amax_s, - amax_o, s_quantizer, o_quantizer, qkv_layout, @@ -1590,9 +1587,6 @@ def _fp8_forward( k=k, v=v, output=output, - stats=stats, - amax_s=amax_s, - amax_o=amax_o, s_quantizer=s_quantizer, o_quantizer=o_quantizer, qkv_layout=qkv_layout, @@ -1799,7 +1793,6 @@ def _build_f16_bwd_graph( v: torch.Tensor, o: torch.Tensor, d_o: torch.Tensor, - stats: torch.Tensor, d_q: torch.Tensor, d_k: torch.Tensor, d_v: torch.Tensor, @@ -2083,7 +2076,6 @@ def _build_fp8_bwd_graph( o, d_o, d_o_f16, - stats, d_q, d_k, d_v, @@ -2701,7 +2693,6 @@ def _fp8_backward( o=o, d_o=d_o, d_o_f16=d_o_f16, - stats=stats, d_q=d_q, d_k=d_k, d_v=d_v, @@ -2954,7 +2945,7 @@ def fused_attn_bwd( d_bias = None if attn_bias_type == "post_scale_bias": # cuDNN does not support the [1,1,1,S] reduction form. - if not tuple(attn_bias.shape[:3]) == (1, 1, 1): + if tuple(attn_bias.shape[:3]) != (1, 1, 1): d_bias = torch.empty_like(attn_bias) d_softmax_offset = torch.empty_like(softmax_offset) if softmax_type != "vanilla" else None cu_seqlens_q_padded = cu_seqlens_q if cu_seqlens_q_padded is None else cu_seqlens_q_padded @@ -3008,7 +2999,6 @@ def fused_attn_bwd( v=v, o=o, d_o=d_o, - stats=stats, d_q=d_q, d_k=d_k, d_v=d_v, diff --git a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py index 7015ff12e96..2e55ddedb35 100644 --- a/transformer_engine/pytorch/attention/fused_mla_q_uproj.py +++ b/transformer_engine/pytorch/attention/fused_mla_q_uproj.py @@ -203,7 +203,9 @@ def run( from cuda.bindings import driver as cuda - stream = cuda.CUstream(torch.cuda.current_stream(x.device).cuda_stream) + stream = cuda.CUstream( # pylint: disable=c-extension-no-member + torch.cuda.current_stream(x.device).cuda_stream + ) wrapper = cls._kernel() if isinstance(w, QuantizedTensor): From c6f373e89ec22bb1bfaa6fe84b0c30f78c9fde51 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 14 Sep 2026 20:00:17 +0000 Subject: [PATCH 35/36] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- transformer_engine/common/attention/cudnn.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/transformer_engine/common/attention/cudnn.py b/transformer_engine/common/attention/cudnn.py index 23963b80a81..f78d01101b7 100644 --- a/transformer_engine/common/attention/cudnn.py +++ b/transformer_engine/common/attention/cudnn.py @@ -355,9 +355,7 @@ def check_f16_fused_attention_support( and dqk != dv ) if unsupported_hopper_bwd_dims: - return _unsupported( - "this Hopper backward head-dimension combination is unsupported" - ) + return _unsupported("this Hopper backward head-dimension combination is unsupported") alibi_supported = ( bias == "alibi" From 033c701bfe720b2f91d51ad5da524733662b8713 Mon Sep 17 00:00:00 2001 From: Vladimir Cherepanov Date: Mon, 14 Sep 2026 20:34:41 +0000 Subject: [PATCH 36/36] Fix JAX attention lint warnings Signed-off-by: Vladimir Cherepanov --- transformer_engine/jax/cpp_extensions/attention.py | 2 +- transformer_engine/jax/cpp_extensions/cudnn_attention.py | 6 ++++-- transformer_engine/jax/cpp_extensions/cudnn_graph.py | 4 ++-- transformer_engine/jax/cpp_extensions/fp8_attention.py | 4 ++-- 4 files changed, 9 insertions(+), 7 deletions(-) diff --git a/transformer_engine/jax/cpp_extensions/attention.py b/transformer_engine/jax/cpp_extensions/attention.py index 4a6aca414d4..c9745675aaf 100644 --- a/transformer_engine/jax/cpp_extensions/attention.py +++ b/transformer_engine/jax/cpp_extensions/attention.py @@ -959,7 +959,7 @@ def abstract( ) ( - batch_shape, + _batch_shape, q_max_seqlen, kv_max_seqlen, attn_heads, diff --git a/transformer_engine/jax/cpp_extensions/cudnn_attention.py b/transformer_engine/jax/cpp_extensions/cudnn_attention.py index 2eb9168161c..47be240dd8d 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_attention.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_attention.py @@ -265,7 +265,7 @@ def _graph_dimensions(info: _LayoutInfo, config): return graph_batch, graph_sq, graph_skv, use_ragged_stats, stats_shape, max_shape -def _tensor(graph, cudnn, *, name, dim, stride, dtype, uid): +def _tensor(graph, _cudnn, *, name, dim, stride, dtype, uid): return graph.tensor( name=name, dim=tuple(int(x) for x in dim), @@ -319,7 +319,7 @@ def _mask_options(cudnn, info: _LayoutInfo, config): return options -def _scalar_tensor(graph, cudnn, name: str, uid: int, dtype): +def _scalar_tensor(graph, _cudnn, name: str, uid: int, dtype): return graph.tensor( name=name, dim=(1, 1, 1, 1), @@ -639,6 +639,8 @@ def _build_bwd_graph( doutput_aval, config, ) -> AttentionGraphInfo: + del stats_aval, output_aval, doutput_aval + cudnn = import_cudnn() info = _layout_info(q_aval, k_aval, v_aval, config.qkv_layout) graph_batch, graph_sq, graph_skv, ragged_stats, stats_shape, max_shape = _graph_dimensions( diff --git a/transformer_engine/jax/cpp_extensions/cudnn_graph.py b/transformer_engine/jax/cpp_extensions/cudnn_graph.py index 15a2a01b31b..c13d9ed7296 100644 --- a/transformer_engine/jax/cpp_extensions/cudnn_graph.py +++ b/transformer_engine/jax/cpp_extensions/cudnn_graph.py @@ -198,9 +198,9 @@ def make_graph(cudnn, io_dtype): return make_cudnn_graph(cudnn, io_dtype) -def graph_hash(serialized_graph: bytes) -> tuple[int, int]: +def graph_hash(graph_data: bytes) -> tuple[int, int]: """Return two signed int64 values used as the C++ graph-cache key.""" - digest = hashlib.sha256(serialized_graph).digest() + digest = hashlib.sha256(graph_data).digest() return ( int.from_bytes(digest[0:8], byteorder="little", signed=True), int.from_bytes(digest[8:16], byteorder="little", signed=True), diff --git a/transformer_engine/jax/cpp_extensions/fp8_attention.py b/transformer_engine/jax/cpp_extensions/fp8_attention.py index c8f2ddff656..8fd5e79eb6d 100644 --- a/transformer_engine/jax/cpp_extensions/fp8_attention.py +++ b/transformer_engine/jax/cpp_extensions/fp8_attention.py @@ -966,7 +966,7 @@ def _quantized_operands(qkv, layout, quantizer, *, both=False): tensor = quantized[0] return (tensor, tensor, tensor), (_rowwise(tensor).data, empty, empty) if layout.is_kvpacked(): - q_tensor, kv_tensor = quantized + q_tensor, kv_tensor = quantized[0], quantized[1] return ( q_tensor, kv_tensor, @@ -1178,7 +1178,7 @@ def fused_attn_fp8_bwd(ctx, doutput, config): q_tensor, k_tensor, v_tensor = quantized mode = _scaling_mode(quantizers.qkv) input_dtype = _rowwise(q_tensor).dq_dtype - (do_tensor,) = _quantize_many((doutput,), quantizers.do, both=mode == "mxfp8") + do_tensor = _quantize_many((doutput,), quantizers.do, both=mode == "mxfp8")[0] q_data, k_data, v_data = ( _rowwise(q_tensor).data, _rowwise(k_tensor).data,