diff --git a/docs/envvars.rst b/docs/envvars.rst
index b3765a06bd..d6ffa241ef 100644
--- a/docs/envvars.rst
+++ b/docs/envvars.rst
@@ -171,18 +171,18 @@ 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)
: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 or 1)
+ :Default: ``0``
+ :Description: Enable diagnostic logging for the FusedAttention graph cache (covers both the F16 and FP8 kernels, forward and backward). When set to ``1``, prints to stderr (prefixed ``[FUSED-ATTN-CACHE]``) a per-lookup ``HIT``/``MISS`` line with the full graph-cache key, a ``BUILD`` line whenever a new graph is constructed, an ``EXEC`` line whenever a graph is executed, a ``SUMMARY`` of graph builds vs. executions at process exit, and a breakdown of cuDNN graph-build timings. Useful for diagnosing redundant graph rebuilds or stale-cache reuse, and for profiling graph-build cost. Has negligible overhead when unset.
+
.. 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 989661b543..4ffa804401 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/jax/test_distributed_fused_attn.py b/tests/jax/test_distributed_fused_attn.py
index 2abd9824b6..6657962e93 100644
--- a/tests/jax/test_distributed_fused_attn.py
+++ b/tests/jax/test_distributed_fused_attn.py
@@ -81,25 +81,6 @@ def impl_test_self_attn(
is_training = True
batch, seqlen, num_head, hidden = data_shape
- if not is_fused_attn_kernel_available(
- is_training,
- dtype,
- dtype,
- QKVLayout.BS3HD,
- attn_bias_type,
- attn_mask_type,
- softmax_type,
- dropout_prob,
- num_head,
- num_head,
- seqlen,
- seqlen,
- hidden,
- hidden,
- None, # no window
- ):
- pytest.skip("No FusedAttn backend found")
-
col_ref = self.generate_collectives_count_ref(
mesh_shape,
mesh_axes,
@@ -233,25 +214,6 @@ def test_cross_attn(
batch, seqlen, num_head, hidden = data_shape
- if not is_fused_attn_kernel_available(
- is_training,
- dtype,
- dtype,
- QKVLayout.BSHD_BS2HD,
- attn_bias_type,
- attn_mask_type,
- softmax_type,
- dropout_prob,
- num_head,
- num_head,
- seqlen,
- seqlen,
- hidden,
- hidden,
- None, # no window
- ):
- pytest.skip("No FusedAttn backend found")
-
col_ref = self.generate_collectives_count_ref()
runner = FusedAttnRunner(
batch,
@@ -425,6 +387,7 @@ def impl_test_context_parallel_attn(
def check_has_backend_for_mask(mask_type):
return is_fused_attn_kernel_available(
is_training,
+ batch,
dtype,
dtype,
qkv_layout,
diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py
index 0d1db0b9e1..30f29c12d4 100644
--- a/tests/jax/test_fused_attn.py
+++ b/tests/jax/test_fused_attn.py
@@ -53,6 +53,9 @@
# Get determinism
_deterministic = not bool(int(os.getenv("NVTE_ALLOW_NONDETERMINISTIC_ALGO", "1")))
+# CI test level
+_TEST_LEVEL = os.getenv("NVTE_JAX_UNITTEST_LEVEL", "L0")
+
@pytest.fixture(autouse=True, scope="module")
def init():
@@ -469,6 +472,27 @@ def _get_max_segments_per_sequence(self):
return 1
def _check_configs(self):
+ # Trim SWA configs for L0 and L1 to reduce test time; need to trim more in future test refactoring.
+ if self.window_size is not None and (
+ self.dropout_prob != 0.0 or self.attn_bias_type is not AttnBiasType.NO_BIAS
+ ):
+ if _TEST_LEVEL == "L0" and (
+ self.softmax_type != AttnSoftmaxType.VANILLA_SOFTMAX
+ or self.dtype != jnp.bfloat16
+ or self.attn_bias_type is not AttnBiasType.POST_SCALE_BIAS
+ or self.attn_mask_type is not AttnMaskType.NO_MASK
+ ):
+ pytest.skip(
+ "Trimmed SWA+bias/dropout config: only vanilla-softmax + bf16 + post_scale_bias"
+ " + no-mask runs at L0"
+ )
+ if _TEST_LEVEL == "L1" and (
+ self.dtype != jnp.float16 or self.softmax_type != AttnSoftmaxType.LEARNABLE_SOFTMAX
+ ):
+ pytest.skip(
+ "Trimmed SWA+bias/dropout config: only float16 + learnable-softmax runs at L1"
+ )
+
# TODO(rewang): probably adds this in is_fused_attn_available
if self.qkv_layout.is_thd() and not self.attn_mask_type.is_padding():
pytest.skip("THD format requires padding masks.")
@@ -520,8 +544,21 @@ def _check_configs(self):
"is either BSHD_BSHD_BSHD or THD_THD_THD"
)
- self.backend = FusedAttnHelper(
+ bias_batch = bias_heads = bias_seqlen_q = bias_seqlen_kv = None
+ if self.attn_bias_type == AttnBiasType.POST_SCALE_BIAS:
+ if self.bias_shape == BiasShape._1HSS:
+ bias_batch, bias_heads = 1, self.num_heads_q
+ elif self.bias_shape == BiasShape._B1SS:
+ bias_batch, bias_heads = self.batch_size, 1
+ elif self.bias_shape == BiasShape._BHSS:
+ bias_batch, bias_heads = self.batch_size, self.num_heads_q
+ elif self.bias_shape == BiasShape._11SS:
+ bias_batch, bias_heads = 1, 1
+ bias_seqlen_q, bias_seqlen_kv = self.max_seqlen_q, self.max_seqlen_kv
+
+ self.backend, message = FusedAttnHelper(
self.is_training,
+ self.batch_size,
self.dtype,
self.dtype,
self.qkv_layout,
@@ -536,9 +573,14 @@ def _check_configs(self):
self.head_dim_qk,
self.head_dim_v,
(-1, -1) if self.window_size is None else self.window_size,
+ self.attn_mask_type.is_bottom_right(),
+ bias_batch=bias_batch,
+ bias_heads=bias_heads,
+ bias_seqlen_q=bias_seqlen_q,
+ bias_seqlen_kv=bias_seqlen_kv,
).get_fused_attn_backend()
if self.backend != NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen:
- pytest.skip("Unsupported inputs combination or device compute capability.")
+ pytest.skip(message)
if (
self.attn_bias_type == AttnBiasType.POST_SCALE_BIAS
diff --git a/tests/jax/test_fused_attn_score_mod.py b/tests/jax/test_fused_attn_score_mod.py
index b1f165f491..f854f4b16a 100644
--- a/tests/jax/test_fused_attn_score_mod.py
+++ b/tests/jax/test_fused_attn_score_mod.py
@@ -18,7 +18,7 @@
)
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 transformer_engine_jax import get_device_compute_capability, NVTE_Fused_Attn_Backend
from test_fused_attn import FusedAttnRunner, SeqDescFormat
@@ -397,9 +397,17 @@ def _identity_score_mod(_graph, score, _tensors):
def _install_fake_flax_fused_attn(monkeypatch, *, kernel_available=True):
captured = {}
- def fake_fused_attn_kernel_check(*args, **kwargs):
- captured.setdefault("kernel_checks", []).append((args, kwargs))
- return kernel_available
+ class FakeFusedAttnHelper:
+ def __init__(self, *args, **kwargs):
+ captured.setdefault("kernel_checks", []).append((args, kwargs))
+
+ def get_fused_attn_backend(self):
+ if kernel_available:
+ return NVTE_Fused_Attn_Backend.NVTE_F16_arbitrary_seqlen, ""
+ return (
+ NVTE_Fused_Attn_Backend.NVTE_No_Backend,
+ "fake FusedAttnHelper: no fused attention backend available for this configuration",
+ )
def fake_fused_attn(
qkv,
@@ -454,11 +462,7 @@ def fake_fused_attn(
)
return qkv[0]
- monkeypatch.setattr(
- flax_transformer,
- "is_fused_attn_kernel_available",
- fake_fused_attn_kernel_check,
- )
+ monkeypatch.setattr(flax_transformer, "FusedAttnHelper", FakeFusedAttnHelper)
monkeypatch.setattr(flax_transformer, "fused_attn", fake_fused_attn)
return captured
@@ -533,7 +537,7 @@ def test_dot_product_attention_plumbs_score_mod_to_fused_attn(monkeypatch):
assert captured["attn_bias_type"] is AttnBiasType.NO_BIAS
assert captured["qkv_layout"] is QKVLayout.BSHD_BSHD_BSHD
assert captured["softmax_type"] is AttnSoftmaxType.VANILLA_SOFTMAX
- assert captured["kernel_checks"][0][0][3] is QKVLayout.BSHD_BSHD_BSHD
+ assert captured["kernel_checks"][0][0][4] is QKVLayout.BSHD_BSHD_BSHD
def test_dot_product_attention_unpacks_packed_score_mod_to_separate_layout(monkeypatch):
@@ -557,7 +561,7 @@ def test_dot_product_attention_unpacks_packed_score_mod_to_separate_layout(monke
assert captured["qkv"][0].shape == (1, 8, 1, 16)
assert captured["qkv_layout"] is QKVLayout.BSHD_BSHD_BSHD
assert captured["score_mod"] is _identity_score_mod
- assert captured["kernel_checks"][0][0][3] is QKVLayout.BSHD_BSHD_BSHD
+ assert captured["kernel_checks"][0][0][4] is QKVLayout.BSHD_BSHD_BSHD
def test_multi_head_attention_plumbs_score_mod_to_dot_product_attention(monkeypatch):
diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py
index 7bcdd10daf..6c876ad7cc 100644
--- a/tests/pytorch/attention/test_attention.py
+++ b/tests/pytorch/attention/test_attention.py
@@ -1,6 +1,7 @@
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.
+import copy
import logging
import os
import sys
@@ -140,7 +141,7 @@ def test_dot_product_attention(
tols = dict(atol=1e-3, rtol=1e-3)
if dtype == torch.bfloat16:
tols = dict(atol=1.5e-2, rtol=1.5e-2)
- config = model_configs[model]
+ config = copy.deepcopy(model_configs[model])
is_mla = config.head_dim_qk != config.head_dim_v
is_mqa_gqa = config.num_heads != config.num_gqa_groups
if qkv_layout is None:
@@ -555,6 +556,9 @@ def test_dpa_softmax(dtype, model_configs, model):
@pytest.mark.parametrize("model", model_configs_softmax.keys())
def test_dpa_softmax_thd(dtype, model_configs, model):
"""Test DotProductAttention module with different softmax types"""
+ config = model_configs[model]
+ if "padding" not in config.attn_mask_type:
+ pytest.skip(f"Duplicate test to others with THD and padding mask.")
test_dot_product_attention(dtype, model_configs, model, True, "thd_thd_thd", False, False)
@@ -825,6 +829,9 @@ def test_dpa_bias_shapes(dtype, model_configs, model):
@pytest.mark.parametrize("qkv_layout", ["thd_thd_thd", "sbhd_sbhd_sbhd"])
def test_dpa_sliding_window(dtype, model_configs, model, qkv_layout):
"""Test DotProductAttention module with sliding window attention"""
+ config = model_configs[model]
+ if qkv_layout == "thd_thd_thd" and "padding" not in config.attn_mask_type:
+ pytest.skip(f"Duplicate test to others with THD and padding mask.")
test_dot_product_attention(dtype, model_configs, model, False, qkv_layout, True, False)
@@ -1867,12 +1874,24 @@ def test_dpa_fp8_extra_state(model, dtype):
config = model_configs_fp8_extra_state[model]
# Test backend availability
is_training = True
+ fp8_recipe = recipe.DelayedScaling(
+ margin=0,
+ fp8_format=recipe.Format.HYBRID,
+ amax_history_len=1,
+ amax_compute_algo="most_recent",
+ fp8_dpa=True,
+ )
+ fp8_meta = {}
+ fp8_meta["recipe"] = fp8_recipe
available_backends, _, fused_attn_backends = get_available_attention_backends(
config,
qkv_dtype=torch.float8_e4m3fn,
+ nominal_dtype=dtype,
qkv_layout="sb3hd",
is_training=is_training,
deterministic=_deterministic,
+ fp8=True,
+ fp8_meta=fp8_meta,
)
flash_attn_supported, fused_attn_supported, unfused_attn_supported = available_backends
if not fused_attn_supported and not flash_attn_supported:
@@ -2068,6 +2087,9 @@ def test_mha_fp8_vs_f16(
scaling_mode,
):
"""Test MultiHeadAttention module in FP8"""
+ if not is_training and fp8_dpa_bwd:
+ pytest.skip("fp8_dpa_bwd=True not applicable for inference")
+
os.environ["NVTE_FP8_DPA_BWD"] = "1" if fp8_dpa_bwd else "0"
config = model_configs_fp8_vs_f16[model]
@@ -2098,6 +2120,7 @@ def test_mha_fp8_vs_f16(
available_backends, _, _ = get_available_attention_backends(
config,
qkv_dtype=torch.float8_e4m3fn,
+ nominal_dtype=dtype,
qkv_layout=qkv_format.replace("hd", "h3d"),
fp8=True,
fp8_meta=fp8_meta,
@@ -2316,6 +2339,10 @@ def get_dummy_cuda_rng_tracker() -> CudaRNGStatesTracker:
def test_dpa_fp8_vs_f16(dtype, model, qkv_layout, fp8_dpa_bwd, is_training, scaling_mode):
"""Test DotProductAttention module in FP8"""
config = model_configs_fp8_vs_f16[model]
+ if config.num_heads != config.num_gqa_groups and "3" in qkv_layout:
+ pytest.skip("qkv_layout not applicable for MQA/GQA")
+ if not is_training and fp8_dpa_bwd:
+ pytest.skip("fp8_dpa_bwd=True not applicable for inference")
# TODO(cyang): think of another way to verify dropout results
# test cuDNN FP8 dropout
@@ -2355,6 +2382,7 @@ def test_dpa_fp8_vs_f16(dtype, model, qkv_layout, fp8_dpa_bwd, is_training, scal
available_backends, _, _ = get_available_attention_backends(
config,
qkv_dtype=torch.float8_e4m3fn,
+ nominal_dtype=dtype,
qkv_layout=qkv_layout,
fp8=True,
fp8_meta=fp8_meta,
@@ -2374,8 +2402,6 @@ def test_dpa_fp8_vs_f16(dtype, model, qkv_layout, fp8_dpa_bwd, is_training, scal
pytest.skip("No FP8 attention backend available.")
if not fused_attn_supported_f16:
pytest.skip("No reference backend available.")
- if config.num_heads != config.num_gqa_groups and "3" in qkv_layout:
- pytest.skip("qkv_layout not applicable for MQA/GQA")
if flash_attn_supported:
os.environ["NVTE_FLASH_ATTN"] = "1"
@@ -2664,10 +2690,22 @@ def test_custom_mha_fp8_vs_f16(dtype, model):
# Test backend availability
is_training = True
+ fp8_meta = {}
+ fp8_recipe = recipe.DelayedScaling(
+ margin=0,
+ fp8_format=recipe.Format.HYBRID,
+ amax_history_len=1,
+ amax_compute_algo="most_recent",
+ fp8_dpa=True,
+ )
+ fp8_meta["recipe"] = fp8_recipe
available_backends, _, fused_attn_backends = get_available_attention_backends(
config,
qkv_dtype=torch.float8_e4m3fn,
+ nominal_dtype=dtype,
qkv_layout="bs3hd",
+ fp8=True,
+ fp8_meta=fp8_meta,
is_training=is_training,
deterministic=_deterministic,
)
@@ -2745,6 +2783,7 @@ def _run_custom_mha_fp8(dtype, config, backend):
fp8_format=recipe.Format.HYBRID,
amax_history_len=1,
amax_compute_algo="most_recent",
+ fp8_dpa=True,
)
mha = Custom_MHA_FP8(config).to(dtype=dtype, device="cuda")
diff --git a/tests/pytorch/attention/test_attention_with_cp.py b/tests/pytorch/attention/test_attention_with_cp.py
index 681ee5e6e0..9ffdb865fa 100644
--- a/tests/pytorch/attention/test_attention_with_cp.py
+++ b/tests/pytorch/attention/test_attention_with_cp.py
@@ -378,6 +378,8 @@ def test_cp_with_flash_attention(cp_pool, dtype, model, qkv_format, cp_comm_type
config,
qkv_dtype=dtypes[dtype],
qkv_layout="_".join([qkv_format] * 3),
+ cp_size=num_gpus,
+ cp_size_a2a=2 if cp_comm_type == "a2a+p2p" else 1,
)
flash_attn_supported, *_ = available_backends
if not flash_attn_supported:
@@ -638,6 +640,8 @@ def test_cp_with_fused_attention(
fp8_meta=fp8_meta,
is_training=is_training,
deterministic=_deterministic,
+ cp_size=num_gpus,
+ cp_size_a2a=2 if cp_comm_type == "a2a+p2p" else 1,
)
_, fused_attn_supported, _ = available_backends
diff --git a/tests/pytorch/attention/test_kv_cache.py b/tests/pytorch/attention/test_kv_cache.py
index cdd98d2445..ad57d584f3 100644
--- a/tests/pytorch/attention/test_kv_cache.py
+++ b/tests/pytorch/attention/test_kv_cache.py
@@ -4,6 +4,7 @@
from collections import OrderedDict
from typing import List
+import copy
import os
import sys
import pathlib
@@ -472,8 +473,11 @@ def test_kv_cache(dtype, model, qkv_format, is_paged, backend, module, is_cuda_g
qkv_layout = qkv_format + "_" + "_".join([inference_params_qkv_format] * 2)
if is_paged:
qkv_layout = "paged_kv_" + qkv_layout
- available_backends, _, fused_attn_backends = get_available_attention_backends(
- config,
+ # probe inference configs only; reference configs are widely supported
+ probe_config = copy.deepcopy(config)
+ probe_config.attn_mask_type = "padding_causal"
+ available_backends, _, _ = get_available_attention_backends(
+ probe_config,
qkv_dtype=dtype,
qkv_layout=qkv_layout,
pad_between_seqs=False,
diff --git a/tests/pytorch/utils.py b/tests/pytorch/utils.py
index 21601d8cdd..cdab36b2c8 100644
--- a/tests/pytorch/utils.py
+++ b/tests/pytorch/utils.py
@@ -312,6 +312,10 @@ def __init__(
self.attn_type = "self" if (self.max_seqlen_q == self.max_seqlen_kv) else "cross"
self.bias_shape = bias_shape
self.window_size = check_set_window_size(self.attn_mask_type, window_size)
+ self.bottom_right_diagonal = self.attn_mask_type not in {
+ "causal",
+ "padding_causal",
+ }
self.context_parallel = context_parallel
self.cp_comm_type = cp_comm_type
self.return_max_logit = return_max_logit
@@ -336,6 +340,7 @@ def get_available_attention_backends(
config: ModelConfig,
qkv_dtype: torch.dtype,
qkv_layout: str,
+ nominal_dtype: Optional[torch.dtype] = None,
pad_between_seqs: bool = False,
deterministic: bool = False,
fp8: bool = False,
@@ -344,6 +349,8 @@ def get_available_attention_backends(
inference_params: Optional[InferenceParams] = None,
score_mod: bool = False,
score_mod_bprop: bool = False,
+ cp_size: int = 1,
+ cp_size_a2a: int = 1,
) -> Tuple[List, List]:
"""Check for all available attention backends that support a model configuration"""
@@ -358,9 +365,15 @@ def get_available_attention_backends(
if config.bias_shape == "bhss":
alibi_slopes_shape = [config.batch_size, config.num_heads]
- core_attention_bias_shape = (
- config.bias_shape if config.attn_bias_type == "post_scale_bias" else None
- )
+ core_attention_bias_shape = None
+ if config.attn_bias_type == "post_scale_bias":
+ b_dim, h_dim, sq_dim, skv_dim = config.bias_shape
+ core_attention_bias_shape = (
+ config.batch_size if b_dim == "b" else 1,
+ config.num_heads if h_dim == "h" else 1,
+ config.max_seqlen_q if sq_dim == "s" else 1,
+ config.max_seqlen_kv if skv_dim == "s" else 1,
+ )
core_attention_bias_requires_grad = False
# d=256 is supported by cuDNN 9.0+ for inference but not training
if (
@@ -369,7 +382,7 @@ def get_available_attention_backends(
and config.head_dim_v <= 128
):
# TODO(KshitijLakhani): Remove this guard when cuDNN starts support dbias calculation for bias shape 111s
- if core_attention_bias_shape != "111s":
+ if config.bias_shape != "111s":
core_attention_bias_requires_grad = True
fused_attn_backends = []
@@ -380,6 +393,7 @@ def get_available_attention_backends(
def test():
attention_params = AttentionParams(
qkv_dtype=qkv_dtype,
+ nominal_dtype=nominal_dtype,
qkv_layout=qkv_layout,
batch_size=config.batch_size,
num_heads=config.num_heads,
@@ -390,6 +404,7 @@ def test():
head_dim_v=config.head_dim_v,
attn_mask_type=config.attn_mask_type,
window_size=config.window_size,
+ bottom_right_diagonal=config.bottom_right_diagonal,
alibi_slopes_shape=alibi_slopes_shape,
core_attention_bias_type=config.attn_bias_type,
core_attention_bias_shape=core_attention_bias_shape,
@@ -398,6 +413,8 @@ def test():
attention_dropout=config.dropout_p,
context_parallel=config.context_parallel,
cp_comm_type=config.cp_comm_type,
+ cp_size=cp_size,
+ cp_size_a2a=cp_size_a2a,
deterministic=deterministic,
fp8=fp8,
fp8_meta=fp8_meta,
@@ -437,12 +454,10 @@ def test():
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)
+ _attention_backends["backend_selection_requires_update"] = True
+ available_backends, flash_attention_backend, fused_attention_backend = test()
+ if fused_attention_backend in (FusedAttnBackend[name] for name in backends.values()):
+ fused_attn_backends.append(fused_attention_backend)
return available_backends, flash_attention_backend, fused_attn_backends
diff --git a/transformer_engine/common/CMakeLists.txt b/transformer_engine/common/CMakeLists.txt
index 13880e0d61..742d37a2f8 100644
--- a/transformer_engine/common/CMakeLists.txt
+++ b/transformer_engine/common/CMakeLists.txt
@@ -181,6 +181,7 @@ list(APPEND transformer_engine_cpp_sources
cudnn_utils.cpp
transformer_engine.cpp
fused_attn/fused_attn.cpp
+ fused_attn/config_and_params.cpp
gemm/config.cpp
normalization/common.cpp
normalization/rtc_dispatch.cpp
diff --git a/transformer_engine/common/fused_attn/config_and_params.cpp b/transformer_engine/common/fused_attn/config_and_params.cpp
new file mode 100644
index 0000000000..ca4214dac3
--- /dev/null
+++ b/transformer_engine/common/fused_attn/config_and_params.cpp
@@ -0,0 +1,1160 @@
+/*************************************************************************
+ * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+ *
+ * See LICENSE for license information.
+ ************************************************************************/
+
+#include "config_and_params.h"
+
+#include
+#include
+
+#include
+
+#include "../common.h"
+#include "../util/cuda_runtime.h"
+
+namespace {
+
+void bool_to_uint8(bool in, void *out) {
+ *reinterpret_cast(out) = static_cast(in);
+}
+
+void uint8_to_bool(const void *in, bool &out) {
+ out = static_cast(*reinterpret_cast(in));
+}
+
+} // namespace
+
+namespace transformer_engine {
+
+namespace fused_attn {
+
+// Forward declarations
+size_t get_max_batch_size(size_t batch_size);
+size_t get_max_tokens(size_t num_tokens);
+
+void FusedAttnConfig::derive() {
+ const int64_t b = static_cast(batch_size);
+ const int64_t sq = static_cast(max_seqlen_q);
+ const int64_t skv = static_cast(max_seqlen_kv);
+
+ // convenience fields
+ q_format = nvte_get_q_format(qkv_layout);
+ kv_format = nvte_get_kv_format(qkv_layout);
+ const NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout);
+ is_paged_kv = (layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD);
+ is_ragged_q = (q_format == NVTE_QKV_Format::NVTE_THD);
+ is_ragged_kv = (kv_format == NVTE_QKV_Format::NVTE_THD);
+ is_padding = (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);
+ is_causal = (attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_MASK) ||
+ (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK);
+ is_causal_bottom_right =
+ (attn_mask_type == NVTE_Mask_Type::NVTE_CAUSAL_BOTTOM_RIGHT_MASK) ||
+ (attn_mask_type == NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK);
+
+ // bucket the THD (ragged) batch and token counts
+ const size_t tokens_q = num_tokens_q != 0 ? num_tokens_q : static_cast(b * sq);
+ const size_t tokens_kv = num_tokens_kv != 0 ? num_tokens_kv : static_cast(b * skv);
+ bucketed_batch_size =
+ (is_ragged_q || is_ragged_kv) ? fused_attn::get_max_batch_size(batch_size) : 0;
+ bucketed_num_tokens_q = is_ragged_q ? fused_attn::get_max_tokens(tokens_q) : 0;
+ bucketed_num_tokens_kv = is_ragged_kv ? fused_attn::get_max_tokens(tokens_kv) : 0;
+
+ // use of cu_seqlens vs actual_seqlens
+ const size_t cudnn_runtime_version = cudnnGetVersion();
+ const bool is_dropout = is_training && dropout != 0.0f;
+ uses_cu_seqlens_directly = CUDNN_FRONTEND_VERSION >= 12500 &&
+ (CUDNN_VERSION >= 92400 && cudnn_runtime_version >= 92400) &&
+ !is_dropout;
+
+ // paged KV dimensions
+ if (is_paged_kv) {
+ if (num_pages_k == 0) {
+ num_pages_k = static_cast(b);
+ }
+ if (num_pages_v == 0) {
+ num_pages_v = static_cast(b);
+ }
+ if (page_size_k == 0) {
+ page_size_k = static_cast(skv);
+ }
+ if (page_size_v == 0) {
+ page_size_v = static_cast(skv);
+ }
+ if (max_pages_per_seq_k == 0) {
+ max_pages_per_seq_k = 1;
+ }
+ if (max_pages_per_seq_v == 0) {
+ max_pages_per_seq_v = 1;
+ }
+ }
+}
+
+FusedAttnConfig FusedAttnConfig::make_cache_key() const {
+ FusedAttnConfig cache_cfg = *this;
+
+ // Key the device ID for multi-GPU single-process runs
+ cache_cfg.device_id = cuda::current_device();
+
+ // Normalize bottom_right_diagonal
+ const bool has_window = cache_cfg.window_size_left != -1 || cache_cfg.window_size_right != -1;
+ if (!cache_cfg.is_causal && !cache_cfg.is_causal_bottom_right && !has_window) {
+ cache_cfg.bottom_right_diagonal = false;
+ } else if (cache_cfg.is_causal_bottom_right &&
+ cache_cfg.max_seqlen_q == cache_cfg.max_seqlen_kv && !cache_cfg.is_padding) {
+ cache_cfg.bottom_right_diagonal = false;
+ }
+
+ // Bucket THD (ragged) batch and token counts
+ if (cache_cfg.is_ragged_q || cache_cfg.is_ragged_kv) {
+ const auto cudnn_runtime_version = cudnnGetVersion();
+ const int sm_arch_ = cuda::sm_arch(cuda::current_device());
+ if (cudnn_runtime_version >= 90600 && sm_arch_ != 120) {
+ if (cache_cfg.is_ragged_q) {
+ cache_cfg.max_seqlen_q = cache_cfg.bucketed_num_tokens_q;
+ }
+ if (cache_cfg.is_ragged_kv) {
+ cache_cfg.max_seqlen_kv = cache_cfg.bucketed_num_tokens_kv;
+ }
+ cache_cfg.num_tokens_q = 0;
+ cache_cfg.num_tokens_kv = 0;
+ const bool bucket_batch = !is_forward || !cache_cfg.uses_cu_seqlens_directly;
+ if (bucket_batch) {
+ cache_cfg.batch_size = cache_cfg.bucketed_batch_size;
+ }
+ }
+ }
+
+ // attn_scale is a pass-by-value graph input and different scales can share the same cached graph
+ cache_cfg.attn_scale = 1.0f;
+
+ // Restrict each direction's key to the fields its graph actually consumes, so
+ // no redundant graphs are built and no cache misses either
+ if (is_forward) {
+ cache_cfg.do_dtype = kNVTEBFloat16;
+ cache_cfg.dqkv_dtype = kNVTEBFloat16;
+ cache_cfg.do_format = NVTE_QKV_Format_NOT_SET;
+ cache_cfg.dqkv_layout = NVTE_QKV_Layout_NOT_SET;
+ cache_cfg.do_scale_inv_format = NVTE_QKV_Format_NOT_SET;
+ cache_cfg.deterministic = false;
+ } else {
+ cache_cfg.return_max_logit = false;
+ }
+
+ return cache_cfg;
+}
+
+FusedAttnConfig FusedAttnFwdParams::make_config() const {
+ const FusedAttnFwdParams ¶ms = *this;
+ FusedAttnConfig cfg{};
+ cfg.is_forward = true;
+ cfg.is_training = params.is_training;
+ cfg.deterministic = false;
+ cfg.cuda_graph = params.cuda_graph;
+ cfg.return_max_logit = params.return_max_logit;
+ cfg.qkv_layout = params.qkv_layout;
+ cfg.o_format = params.o_format;
+ cfg.qkv_scale_inv_format = params.qkv_scale_inv_format;
+ cfg.bias_type = params.bias_type;
+ cfg.attn_mask_type = params.attn_mask_type;
+ cfg.softmax_type = params.softmax_type;
+ cfg.attn_scale = params.attn_scale;
+ cfg.dropout = params.dropout;
+ cfg.max_seqlen_q = params.max_seqlen_q;
+ cfg.max_seqlen_kv = params.max_seqlen_kv;
+ cfg.window_size_left = params.window_size_left;
+ cfg.window_size_right = params.window_size_right;
+ cfg.bottom_right_diagonal = params.bottom_right_diagonal;
+
+ const Tensor *input_cu_seqlens_q = convertNVTETensorCheck(params.cu_seqlens_q);
+ const Tensor *input_cu_seqlens_kv = convertNVTETensorCheck(params.cu_seqlens_kv);
+ const Tensor *input_page_table_k = convertNVTETensorCheck(params.page_table_k);
+ const Tensor *input_page_table_v = convertNVTETensorCheck(params.page_table_v);
+ const Tensor *input_Q = convertNVTETensorCheck(params.Q);
+ const Tensor *input_K = convertNVTETensorCheck(params.K);
+ const Tensor *input_V = convertNVTETensorCheck(params.V);
+ const Tensor *input_Bias = convertNVTETensorCheck(params.Bias);
+ const Tensor *output_O = convertNVTETensorCheck(params.O);
+
+ const NVTE_QKV_Format q_format = nvte_get_q_format(params.qkv_layout);
+ const NVTE_QKV_Format kv_format = nvte_get_kv_format(params.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, num_pages_v = 0, page_size_k = 0, page_size_v = 0;
+ int64_t max_pages_per_seq_k = 0, 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];
+ }
+ const NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(params.qkv_layout);
+ if (layout_group == NVTE_QKV_Layout_Group::NVTE_Paged_KV_HD_HD_HD) {
+ const NVTE_QKV_Format paged_kv_format = nvte_get_kv_format(params.qkv_layout);
+ if (paged_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 (paged_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];
+ }
+ }
+
+ const NVTEDType Q_type = static_cast(input_Q->data.dtype);
+ const NVTEDType KV_type = static_cast(input_K->data.dtype);
+ NVTE_CHECK(Q_type == KV_type, "Q and KV must have the same data type.");
+
+ cfg.scaling_mode = input_Q->scaling_mode;
+ cfg.qkv_dtype = Q_type;
+ cfg.o_dtype = static_cast(output_O->data.dtype);
+ cfg.batch_size = b;
+ cfg.num_attn_heads = h_q;
+ cfg.num_gqa_groups = h_kv;
+ cfg.head_dim_qk = d_qk;
+ cfg.head_dim_v = d_v;
+ cfg.num_pages_k = static_cast(num_pages_k);
+ cfg.num_pages_v = static_cast(num_pages_v);
+ cfg.page_size_k = static_cast(page_size_k);
+ cfg.page_size_v = static_cast(page_size_v);
+ cfg.max_pages_per_seq_k = static_cast(max_pages_per_seq_k);
+ cfg.max_pages_per_seq_v = static_cast(max_pages_per_seq_v);
+ cfg.num_tokens_q = t_q;
+ cfg.num_tokens_kv = t_kv;
+
+ if ((params.bias_type != NVTE_NO_BIAS) && (params.bias_type != NVTE_ALIBI) &&
+ input_Bias->data.shape.size() >= 4) {
+ cfg.bias_batch_size = input_Bias->data.shape[0];
+ cfg.bias_num_heads = input_Bias->data.shape[1];
+ cfg.bias_seqlen_q = input_Bias->data.shape[2];
+ cfg.bias_seqlen_kv = input_Bias->data.shape[3];
+ }
+ return cfg;
+}
+
+FusedAttnConfig FusedAttnBwdParams::make_config() const {
+ const FusedAttnBwdParams ¶ms = *this;
+ FusedAttnConfig cfg{};
+ cfg.is_training = true;
+ cfg.deterministic = params.deterministic;
+ cfg.cuda_graph = params.cuda_graph;
+ cfg.return_max_logit = false;
+ cfg.qkv_layout = params.qkv_layout;
+ cfg.o_format = params.o_format;
+ cfg.do_format = params.do_format;
+ cfg.dqkv_layout = params.dqkv_layout;
+ cfg.qkv_scale_inv_format = params.qkv_scale_inv_format;
+ cfg.do_scale_inv_format = params.do_scale_inv_format;
+ cfg.bias_type = params.bias_type;
+ cfg.attn_mask_type = params.attn_mask_type;
+ cfg.softmax_type = params.softmax_type;
+ cfg.attn_scale = params.attn_scale;
+ cfg.dropout = params.dropout;
+ cfg.max_seqlen_q = params.max_seqlen_q;
+ cfg.max_seqlen_kv = params.max_seqlen_kv;
+ cfg.window_size_left = params.window_size_left;
+ cfg.window_size_right = params.window_size_right;
+ cfg.bottom_right_diagonal = params.bottom_right_diagonal;
+
+ const Tensor *input_cu_seqlens_q = convertNVTETensorCheck(params.cu_seqlens_q);
+ const Tensor *input_cu_seqlens_kv = convertNVTETensorCheck(params.cu_seqlens_kv);
+ const Tensor *input_Q = convertNVTETensorCheck(params.Q);
+ const Tensor *input_K = convertNVTETensorCheck(params.K);
+ const Tensor *input_V = convertNVTETensorCheck(params.V);
+ const Tensor *input_O = convertNVTETensorCheck(params.O);
+ const Tensor *input_dO = convertNVTETensorCheck(params.dO);
+ const Tensor *output_dQ = convertNVTETensorCheck(params.dQ);
+ const Tensor *output_dBias = convertNVTETensorCheck(params.dBias);
+
+ const NVTE_QKV_Format q_format = nvte_get_q_format(params.qkv_layout);
+ const NVTE_QKV_Format kv_format = nvte_get_kv_format(params.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;
+ }
+
+ const NVTEDType Q_type = static_cast(input_Q->data.dtype);
+ const NVTEDType KV_type = static_cast(input_K->data.dtype);
+ NVTE_CHECK(Q_type == KV_type, "Q and KV must have the same data type.");
+
+ cfg.scaling_mode = input_Q->scaling_mode;
+ cfg.qkv_dtype = Q_type;
+ cfg.o_dtype = static_cast(input_O->data.dtype);
+ cfg.do_dtype = static_cast(input_dO->data.dtype);
+ cfg.dqkv_dtype = static_cast(output_dQ->data.dtype);
+ cfg.batch_size = b;
+ cfg.num_attn_heads = h_q;
+ cfg.num_gqa_groups = h_kv;
+ cfg.head_dim_qk = d_qk;
+ cfg.head_dim_v = d_v;
+ cfg.num_tokens_q = t_q;
+ cfg.num_tokens_kv = t_kv;
+
+ if ((params.bias_type != NVTE_NO_BIAS) && (params.bias_type != NVTE_ALIBI) &&
+ output_dBias->data.shape.size() >= 4) {
+ cfg.bias_batch_size = output_dBias->data.shape[0];
+ cfg.bias_num_heads = output_dBias->data.shape[1];
+ cfg.bias_seqlen_q = output_dBias->data.shape[2];
+ cfg.bias_seqlen_kv = output_dBias->data.shape[3];
+ }
+ return cfg;
+}
+
+} // namespace fused_attn
+} // namespace transformer_engine
+
+NVTEFusedAttnConfig nvte_create_fused_attn_config() {
+ return new transformer_engine::fused_attn::FusedAttnConfig{};
+}
+
+void nvte_destroy_fused_attn_config(NVTEFusedAttnConfig config) {
+ delete transformer_engine::fused_attn::get_fused_attn_config_mutable(config);
+}
+
+void nvte_get_fused_attn_config_attribute(NVTEFusedAttnConfig config,
+ NVTEFusedAttnConfigAttribute attr, void *buf,
+ size_t size_in_bytes, size_t *size_written) {
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+
+ NVTE_CHECK(attr < kNVTEFusedAttnConfigNumAttributes, "Invalid NVTEFusedAttnConfigAttribute (got ",
+ static_cast(attr), ")");
+ const auto &attr_size = FusedAttnConfig::attr_sizes[attr];
+ if (size_written != nullptr) {
+ *size_written = attr_size;
+ }
+ if (buf == nullptr) {
+ return;
+ }
+ NVTE_CHECK(size_in_bytes >= attr_size,
+ "Buffer is too small for fused attention config attribute (attribute ",
+ static_cast(attr), " needs ", attr_size, " bytes, but buffer has ", size_in_bytes,
+ " bytes)");
+
+ const auto &cfg = *get_fused_attn_config(config);
+ switch (attr) {
+ case kNVTEFusedAttnConfigIsTraining:
+ bool_to_uint8(cfg.is_training, buf);
+ break;
+ case kNVTEFusedAttnConfigDeterministic:
+ bool_to_uint8(cfg.deterministic, buf);
+ break;
+ case kNVTEFusedAttnConfigCudaGraph:
+ bool_to_uint8(cfg.cuda_graph, buf);
+ break;
+ case kNVTEFusedAttnConfigReturnMaxLogit:
+ bool_to_uint8(cfg.return_max_logit, buf);
+ break;
+ case kNVTEFusedAttnConfigAttnMaskType:
+ std::memcpy(buf, &cfg.attn_mask_type, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasType:
+ std::memcpy(buf, &cfg.bias_type, attr_size);
+ break;
+ case kNVTEFusedAttnConfigWindowSizeLeft:
+ std::memcpy(buf, &cfg.window_size_left, attr_size);
+ break;
+ case kNVTEFusedAttnConfigWindowSizeRight:
+ std::memcpy(buf, &cfg.window_size_right, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBottomRightDiagonal:
+ bool_to_uint8(cfg.bottom_right_diagonal, buf);
+ break;
+ case kNVTEFusedAttnConfigSoftmaxType:
+ std::memcpy(buf, &cfg.softmax_type, attr_size);
+ break;
+ case kNVTEFusedAttnConfigScalingMode:
+ std::memcpy(buf, &cfg.scaling_mode, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDropout:
+ std::memcpy(buf, &cfg.dropout, attr_size);
+ break;
+ case kNVTEFusedAttnConfigAttnScale:
+ std::memcpy(buf, &cfg.attn_scale, attr_size);
+ break;
+ case kNVTEFusedAttnConfigQKVDtype:
+ std::memcpy(buf, &cfg.qkv_dtype, attr_size);
+ break;
+ case kNVTEFusedAttnConfigODtype:
+ std::memcpy(buf, &cfg.o_dtype, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDODtype:
+ std::memcpy(buf, &cfg.do_dtype, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDQKVDtype:
+ std::memcpy(buf, &cfg.dqkv_dtype, attr_size);
+ break;
+ case kNVTEFusedAttnConfigQKVLayout:
+ std::memcpy(buf, &cfg.qkv_layout, attr_size);
+ break;
+ case kNVTEFusedAttnConfigOFormat:
+ std::memcpy(buf, &cfg.o_format, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDOFormat:
+ std::memcpy(buf, &cfg.do_format, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDQKVLayout:
+ std::memcpy(buf, &cfg.dqkv_layout, attr_size);
+ break;
+ case kNVTEFusedAttnConfigQKVScaleInvFormat:
+ std::memcpy(buf, &cfg.qkv_scale_inv_format, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDOScaleInvFormat:
+ std::memcpy(buf, &cfg.do_scale_inv_format, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBatchSize:
+ std::memcpy(buf, &cfg.batch_size, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumAttnHeads:
+ std::memcpy(buf, &cfg.num_attn_heads, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumGQAGroups:
+ std::memcpy(buf, &cfg.num_gqa_groups, attr_size);
+ break;
+ case kNVTEFusedAttnConfigHeadDimQK:
+ std::memcpy(buf, &cfg.head_dim_qk, attr_size);
+ break;
+ case kNVTEFusedAttnConfigHeadDimV:
+ std::memcpy(buf, &cfg.head_dim_v, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxSeqlenQ:
+ std::memcpy(buf, &cfg.max_seqlen_q, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxSeqlenKV:
+ std::memcpy(buf, &cfg.max_seqlen_kv, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumTokensQ:
+ std::memcpy(buf, &cfg.num_tokens_q, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumTokensKV:
+ std::memcpy(buf, &cfg.num_tokens_kv, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumPagesK:
+ std::memcpy(buf, &cfg.num_pages_k, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumPagesV:
+ std::memcpy(buf, &cfg.num_pages_v, attr_size);
+ break;
+ case kNVTEFusedAttnConfigPageSizeK:
+ std::memcpy(buf, &cfg.page_size_k, attr_size);
+ break;
+ case kNVTEFusedAttnConfigPageSizeV:
+ std::memcpy(buf, &cfg.page_size_v, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxPagesPerSeqK:
+ std::memcpy(buf, &cfg.max_pages_per_seq_k, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxPagesPerSeqV:
+ std::memcpy(buf, &cfg.max_pages_per_seq_v, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasBatchSize:
+ std::memcpy(buf, &cfg.bias_batch_size, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasNumHeads:
+ std::memcpy(buf, &cfg.bias_num_heads, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasSeqlenQ:
+ std::memcpy(buf, &cfg.bias_seqlen_q, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasSeqlenKV:
+ std::memcpy(buf, &cfg.bias_seqlen_kv, attr_size);
+ break;
+ default:
+ NVTE_ERROR("Unsupported NVTEFusedAttnConfigAttribute (got ", static_cast(attr), ")");
+ }
+}
+
+void nvte_set_fused_attn_config_attribute(NVTEFusedAttnConfig config,
+ NVTEFusedAttnConfigAttribute attr, const void *buf,
+ size_t size_in_bytes) {
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+
+ NVTE_CHECK(attr < kNVTEFusedAttnConfigNumAttributes, "Invalid NVTEFusedAttnConfigAttribute (got ",
+ static_cast(attr), ")");
+ const auto &attr_size = FusedAttnConfig::attr_sizes[attr];
+ NVTE_CHECK(size_in_bytes >= attr_size,
+ "Buffer is too small for fused attention config attribute (attribute ",
+ static_cast(attr), " needs ", attr_size, " bytes, but buffer has ", size_in_bytes,
+ " bytes)");
+ NVTE_CHECK(buf != nullptr, "Invalid buffer (got NULL)");
+
+ auto &cfg = *get_fused_attn_config_mutable(config);
+ switch (attr) {
+ case kNVTEFusedAttnConfigIsTraining:
+ uint8_to_bool(buf, cfg.is_training);
+ break;
+ case kNVTEFusedAttnConfigDeterministic:
+ uint8_to_bool(buf, cfg.deterministic);
+ break;
+ case kNVTEFusedAttnConfigCudaGraph:
+ uint8_to_bool(buf, cfg.cuda_graph);
+ break;
+ case kNVTEFusedAttnConfigReturnMaxLogit:
+ uint8_to_bool(buf, cfg.return_max_logit);
+ break;
+ case kNVTEFusedAttnConfigAttnMaskType:
+ std::memcpy(&cfg.attn_mask_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasType:
+ std::memcpy(&cfg.bias_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigWindowSizeLeft:
+ std::memcpy(&cfg.window_size_left, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigWindowSizeRight:
+ std::memcpy(&cfg.window_size_right, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBottomRightDiagonal:
+ uint8_to_bool(buf, cfg.bottom_right_diagonal);
+ break;
+ case kNVTEFusedAttnConfigSoftmaxType:
+ std::memcpy(&cfg.softmax_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigScalingMode:
+ std::memcpy(&cfg.scaling_mode, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDropout:
+ std::memcpy(&cfg.dropout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigAttnScale:
+ std::memcpy(&cfg.attn_scale, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigQKVDtype:
+ std::memcpy(&cfg.qkv_dtype, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigODtype:
+ std::memcpy(&cfg.o_dtype, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDODtype:
+ std::memcpy(&cfg.do_dtype, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDQKVDtype:
+ std::memcpy(&cfg.dqkv_dtype, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigQKVLayout:
+ std::memcpy(&cfg.qkv_layout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigOFormat:
+ std::memcpy(&cfg.o_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDOFormat:
+ std::memcpy(&cfg.do_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDQKVLayout:
+ std::memcpy(&cfg.dqkv_layout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigQKVScaleInvFormat:
+ std::memcpy(&cfg.qkv_scale_inv_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigDOScaleInvFormat:
+ std::memcpy(&cfg.do_scale_inv_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBatchSize:
+ std::memcpy(&cfg.batch_size, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumAttnHeads:
+ std::memcpy(&cfg.num_attn_heads, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumGQAGroups:
+ std::memcpy(&cfg.num_gqa_groups, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigHeadDimQK:
+ std::memcpy(&cfg.head_dim_qk, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigHeadDimV:
+ std::memcpy(&cfg.head_dim_v, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxSeqlenQ:
+ std::memcpy(&cfg.max_seqlen_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxSeqlenKV:
+ std::memcpy(&cfg.max_seqlen_kv, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumTokensQ:
+ std::memcpy(&cfg.num_tokens_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumTokensKV:
+ std::memcpy(&cfg.num_tokens_kv, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumPagesK:
+ std::memcpy(&cfg.num_pages_k, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigNumPagesV:
+ std::memcpy(&cfg.num_pages_v, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigPageSizeK:
+ std::memcpy(&cfg.page_size_k, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigPageSizeV:
+ std::memcpy(&cfg.page_size_v, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxPagesPerSeqK:
+ std::memcpy(&cfg.max_pages_per_seq_k, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigMaxPagesPerSeqV:
+ std::memcpy(&cfg.max_pages_per_seq_v, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasBatchSize:
+ std::memcpy(&cfg.bias_batch_size, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasNumHeads:
+ std::memcpy(&cfg.bias_num_heads, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasSeqlenQ:
+ std::memcpy(&cfg.bias_seqlen_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnConfigBiasSeqlenKV:
+ std::memcpy(&cfg.bias_seqlen_kv, buf, attr_size);
+ break;
+ default:
+ NVTE_ERROR("Unsupported NVTEFusedAttnConfigAttribute (got ", static_cast(attr), ")");
+ }
+}
+
+NVTEFusedAttnFwdParams nvte_create_fused_attn_fwd_params() {
+ return new transformer_engine::fused_attn::FusedAttnFwdParams{};
+}
+
+void nvte_destroy_fused_attn_fwd_params(NVTEFusedAttnFwdParams params) {
+ delete transformer_engine::fused_attn::get_fused_attn_fwd_params_mutable(params);
+}
+
+void nvte_get_fused_attn_fwd_params_attribute(NVTEFusedAttnFwdParams params,
+ NVTEFusedAttnFwdParamsAttribute attr, void *buf,
+ size_t size_in_bytes, size_t *size_written) {
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+ NVTE_CHECK(attr < kNVTEFusedAttnFwdParamsNumAttributes,
+ "Invalid NVTEFusedAttnFwdParamsAttribute (got ", static_cast(attr), ")");
+ const auto &attr_size = FusedAttnFwdParams::attr_sizes[attr];
+ if (size_written != nullptr) {
+ *size_written = attr_size;
+ }
+ if (buf == nullptr) {
+ return;
+ }
+ NVTE_CHECK(size_in_bytes >= attr_size, "Buffer is too small for attribute (need ", attr_size,
+ ", got ", size_in_bytes, ")");
+ const auto &p = *get_fused_attn_fwd_params(params);
+ switch (attr) {
+ case kNVTEFusedAttnFwdParamsQ:
+ std::memcpy(buf, &p.Q, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsK:
+ std::memcpy(buf, &p.K, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsV:
+ std::memcpy(buf, &p.V, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsBias:
+ std::memcpy(buf, &p.Bias, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsSoftmaxOffset:
+ std::memcpy(buf, &p.SoftmaxOffset, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensQ:
+ std::memcpy(buf, &p.cu_seqlens_q, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensKV:
+ std::memcpy(buf, &p.cu_seqlens_kv, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensQPadded:
+ std::memcpy(buf, &p.cu_seqlens_q_padded, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensKVPadded:
+ std::memcpy(buf, &p.cu_seqlens_kv_padded, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsPageTableK:
+ std::memcpy(buf, &p.page_table_k, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsPageTableV:
+ std::memcpy(buf, &p.page_table_v, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsRngState:
+ std::memcpy(buf, &p.rng_state, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsS:
+ std::memcpy(buf, &p.S, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsO:
+ std::memcpy(buf, &p.O, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsAuxCtxTensors:
+ std::memcpy(buf, &p.Aux_CTX_Tensors, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsIsTraining:
+ bool_to_uint8(p.is_training, buf);
+ break;
+ case kNVTEFusedAttnFwdParamsCudaGraph:
+ bool_to_uint8(p.cuda_graph, buf);
+ break;
+ case kNVTEFusedAttnFwdParamsReturnMaxLogit:
+ bool_to_uint8(p.return_max_logit, buf);
+ break;
+ case kNVTEFusedAttnFwdParamsAttnMaskType:
+ std::memcpy(buf, &p.attn_mask_type, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsBiasType:
+ std::memcpy(buf, &p.bias_type, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsSoftmaxType:
+ std::memcpy(buf, &p.softmax_type, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsWindowSizeLeft:
+ std::memcpy(buf, &p.window_size_left, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsWindowSizeRight:
+ std::memcpy(buf, &p.window_size_right, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsBottomRightDiagonal:
+ bool_to_uint8(p.bottom_right_diagonal, buf);
+ break;
+ case kNVTEFusedAttnFwdParamsDropout:
+ std::memcpy(buf, &p.dropout, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsAttnScale:
+ std::memcpy(buf, &p.attn_scale, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsQKVLayout:
+ std::memcpy(buf, &p.qkv_layout, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsOFormat:
+ std::memcpy(buf, &p.o_format, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsQKVScaleInvFormat:
+ std::memcpy(buf, &p.qkv_scale_inv_format, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsMaxSeqlenQ:
+ std::memcpy(buf, &p.max_seqlen_q, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsMaxSeqlenKV:
+ std::memcpy(buf, &p.max_seqlen_kv, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsWorkspace:
+ std::memcpy(buf, &p.workspace, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsStream:
+ std::memcpy(buf, &p.stream, attr_size);
+ break;
+ default:
+ NVTE_ERROR("Unsupported NVTEFusedAttnFwdParamsAttribute (got ", static_cast(attr), ")");
+ }
+}
+
+void nvte_set_fused_attn_fwd_params_attribute(NVTEFusedAttnFwdParams params,
+ NVTEFusedAttnFwdParamsAttribute attr, const void *buf,
+ size_t size_in_bytes) {
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+ NVTE_CHECK(attr < kNVTEFusedAttnFwdParamsNumAttributes,
+ "Invalid NVTEFusedAttnFwdParamsAttribute (got ", static_cast(attr), ")");
+ const auto &attr_size = FusedAttnFwdParams::attr_sizes[attr];
+ NVTE_CHECK(buf != nullptr, "Input buffer must not be NULL.");
+ NVTE_CHECK(size_in_bytes >= attr_size, "Buffer is too small for attribute (need ", attr_size,
+ ", got ", size_in_bytes, ")");
+ auto &p = *get_fused_attn_fwd_params_mutable(params);
+ switch (attr) {
+ case kNVTEFusedAttnFwdParamsQ:
+ std::memcpy(&p.Q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsK:
+ std::memcpy(&p.K, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsV:
+ std::memcpy(&p.V, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsBias:
+ std::memcpy(&p.Bias, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsSoftmaxOffset:
+ std::memcpy(&p.SoftmaxOffset, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensQ:
+ std::memcpy(&p.cu_seqlens_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensKV:
+ std::memcpy(&p.cu_seqlens_kv, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensQPadded:
+ std::memcpy(&p.cu_seqlens_q_padded, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsCuSeqlensKVPadded:
+ std::memcpy(&p.cu_seqlens_kv_padded, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsPageTableK:
+ std::memcpy(&p.page_table_k, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsPageTableV:
+ std::memcpy(&p.page_table_v, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsRngState:
+ std::memcpy(&p.rng_state, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsS:
+ std::memcpy(&p.S, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsO:
+ std::memcpy(&p.O, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsAuxCtxTensors:
+ std::memcpy(&p.Aux_CTX_Tensors, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsIsTraining:
+ uint8_to_bool(buf, p.is_training);
+ break;
+ case kNVTEFusedAttnFwdParamsCudaGraph:
+ uint8_to_bool(buf, p.cuda_graph);
+ break;
+ case kNVTEFusedAttnFwdParamsReturnMaxLogit:
+ uint8_to_bool(buf, p.return_max_logit);
+ break;
+ case kNVTEFusedAttnFwdParamsAttnMaskType:
+ std::memcpy(&p.attn_mask_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsBiasType:
+ std::memcpy(&p.bias_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsSoftmaxType:
+ std::memcpy(&p.softmax_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsWindowSizeLeft:
+ std::memcpy(&p.window_size_left, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsWindowSizeRight:
+ std::memcpy(&p.window_size_right, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsBottomRightDiagonal:
+ uint8_to_bool(buf, p.bottom_right_diagonal);
+ break;
+ case kNVTEFusedAttnFwdParamsDropout:
+ std::memcpy(&p.dropout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsAttnScale:
+ std::memcpy(&p.attn_scale, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsQKVLayout:
+ std::memcpy(&p.qkv_layout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsOFormat:
+ std::memcpy(&p.o_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsQKVScaleInvFormat:
+ std::memcpy(&p.qkv_scale_inv_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsMaxSeqlenQ:
+ std::memcpy(&p.max_seqlen_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsMaxSeqlenKV:
+ std::memcpy(&p.max_seqlen_kv, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsWorkspace:
+ std::memcpy(&p.workspace, buf, attr_size);
+ break;
+ case kNVTEFusedAttnFwdParamsStream:
+ std::memcpy(&p.stream, buf, attr_size);
+ break;
+ default:
+ NVTE_ERROR("Unsupported NVTEFusedAttnFwdParamsAttribute (got ", static_cast(attr), ")");
+ }
+}
+
+NVTEFusedAttnBwdParams nvte_create_fused_attn_bwd_params() {
+ return new transformer_engine::fused_attn::FusedAttnBwdParams{};
+}
+
+void nvte_destroy_fused_attn_bwd_params(NVTEFusedAttnBwdParams params) {
+ delete transformer_engine::fused_attn::get_fused_attn_bwd_params_mutable(params);
+}
+
+void nvte_get_fused_attn_bwd_params_attribute(NVTEFusedAttnBwdParams params,
+ NVTEFusedAttnBwdParamsAttribute attr, void *buf,
+ size_t size_in_bytes, size_t *size_written) {
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+ NVTE_CHECK(attr < kNVTEFusedAttnBwdParamsNumAttributes,
+ "Invalid NVTEFusedAttnBwdParamsAttribute (got ", static_cast(attr), ")");
+ const auto &attr_size = FusedAttnBwdParams::attr_sizes[attr];
+ if (size_written != nullptr) {
+ *size_written = attr_size;
+ }
+ if (buf == nullptr) {
+ return;
+ }
+ NVTE_CHECK(size_in_bytes >= attr_size, "Buffer is too small for attribute (need ", attr_size,
+ ", got ", size_in_bytes, ")");
+ const auto &p = *get_fused_attn_bwd_params(params);
+ switch (attr) {
+ case kNVTEFusedAttnBwdParamsQ:
+ std::memcpy(buf, &p.Q, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsK:
+ std::memcpy(buf, &p.K, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsV:
+ std::memcpy(buf, &p.V, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsO:
+ std::memcpy(buf, &p.O, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDO:
+ std::memcpy(buf, &p.dO, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsS:
+ std::memcpy(buf, &p.S, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDP:
+ std::memcpy(buf, &p.dP, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsAuxCtxTensors:
+ std::memcpy(buf, &p.Aux_CTX_Tensors, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDQ:
+ std::memcpy(buf, &p.dQ, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDK:
+ std::memcpy(buf, &p.dK, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDV:
+ std::memcpy(buf, &p.dV, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDBias:
+ std::memcpy(buf, &p.dBias, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDSoftmaxOffset:
+ std::memcpy(buf, &p.dSoftmaxOffset, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensQ:
+ std::memcpy(buf, &p.cu_seqlens_q, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensKV:
+ std::memcpy(buf, &p.cu_seqlens_kv, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensQPadded:
+ std::memcpy(buf, &p.cu_seqlens_q_padded, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensKVPadded:
+ std::memcpy(buf, &p.cu_seqlens_kv_padded, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCudaGraph:
+ bool_to_uint8(p.cuda_graph, buf);
+ break;
+ case kNVTEFusedAttnBwdParamsDeterministic:
+ bool_to_uint8(p.deterministic, buf);
+ break;
+ case kNVTEFusedAttnBwdParamsAttnMaskType:
+ std::memcpy(buf, &p.attn_mask_type, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsBiasType:
+ std::memcpy(buf, &p.bias_type, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsSoftmaxType:
+ std::memcpy(buf, &p.softmax_type, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsWindowSizeLeft:
+ std::memcpy(buf, &p.window_size_left, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsWindowSizeRight:
+ std::memcpy(buf, &p.window_size_right, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsBottomRightDiagonal:
+ bool_to_uint8(p.bottom_right_diagonal, buf);
+ break;
+ case kNVTEFusedAttnBwdParamsDropout:
+ std::memcpy(buf, &p.dropout, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsAttnScale:
+ std::memcpy(buf, &p.attn_scale, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsQKVLayout:
+ std::memcpy(buf, &p.qkv_layout, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsOFormat:
+ std::memcpy(buf, &p.o_format, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDOFormat:
+ std::memcpy(buf, &p.do_format, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDQKVLayout:
+ std::memcpy(buf, &p.dqkv_layout, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsQKVScaleInvFormat:
+ std::memcpy(buf, &p.qkv_scale_inv_format, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDOScaleInvFormat:
+ std::memcpy(buf, &p.do_scale_inv_format, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsMaxSeqlenQ:
+ std::memcpy(buf, &p.max_seqlen_q, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsMaxSeqlenKV:
+ std::memcpy(buf, &p.max_seqlen_kv, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsWorkspace:
+ std::memcpy(buf, &p.workspace, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsStream:
+ std::memcpy(buf, &p.stream, attr_size);
+ break;
+ default:
+ NVTE_ERROR("Unsupported NVTEFusedAttnBwdParamsAttribute (got ", static_cast(attr), ")");
+ }
+}
+
+void nvte_set_fused_attn_bwd_params_attribute(NVTEFusedAttnBwdParams params,
+ NVTEFusedAttnBwdParamsAttribute attr, const void *buf,
+ size_t size_in_bytes) {
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+ NVTE_CHECK(attr < kNVTEFusedAttnBwdParamsNumAttributes,
+ "Invalid NVTEFusedAttnBwdParamsAttribute (got ", static_cast(attr), ")");
+ const auto &attr_size = FusedAttnBwdParams::attr_sizes[attr];
+ NVTE_CHECK(buf != nullptr, "Input buffer must not be NULL.");
+ NVTE_CHECK(size_in_bytes >= attr_size, "Buffer is too small for attribute (need ", attr_size,
+ ", got ", size_in_bytes, ")");
+ auto &p = *get_fused_attn_bwd_params_mutable(params);
+ switch (attr) {
+ case kNVTEFusedAttnBwdParamsQ:
+ std::memcpy(&p.Q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsK:
+ std::memcpy(&p.K, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsV:
+ std::memcpy(&p.V, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsO:
+ std::memcpy(&p.O, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDO:
+ std::memcpy(&p.dO, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsS:
+ std::memcpy(&p.S, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDP:
+ std::memcpy(&p.dP, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsAuxCtxTensors:
+ std::memcpy(&p.Aux_CTX_Tensors, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDQ:
+ std::memcpy(&p.dQ, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDK:
+ std::memcpy(&p.dK, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDV:
+ std::memcpy(&p.dV, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDBias:
+ std::memcpy(&p.dBias, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDSoftmaxOffset:
+ std::memcpy(&p.dSoftmaxOffset, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensQ:
+ std::memcpy(&p.cu_seqlens_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensKV:
+ std::memcpy(&p.cu_seqlens_kv, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensQPadded:
+ std::memcpy(&p.cu_seqlens_q_padded, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCuSeqlensKVPadded:
+ std::memcpy(&p.cu_seqlens_kv_padded, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsCudaGraph:
+ uint8_to_bool(buf, p.cuda_graph);
+ break;
+ case kNVTEFusedAttnBwdParamsDeterministic:
+ uint8_to_bool(buf, p.deterministic);
+ break;
+ case kNVTEFusedAttnBwdParamsAttnMaskType:
+ std::memcpy(&p.attn_mask_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsBiasType:
+ std::memcpy(&p.bias_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsSoftmaxType:
+ std::memcpy(&p.softmax_type, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsWindowSizeLeft:
+ std::memcpy(&p.window_size_left, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsWindowSizeRight:
+ std::memcpy(&p.window_size_right, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsBottomRightDiagonal:
+ uint8_to_bool(buf, p.bottom_right_diagonal);
+ break;
+ case kNVTEFusedAttnBwdParamsDropout:
+ std::memcpy(&p.dropout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsAttnScale:
+ std::memcpy(&p.attn_scale, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsQKVLayout:
+ std::memcpy(&p.qkv_layout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsOFormat:
+ std::memcpy(&p.o_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDOFormat:
+ std::memcpy(&p.do_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDQKVLayout:
+ std::memcpy(&p.dqkv_layout, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsQKVScaleInvFormat:
+ std::memcpy(&p.qkv_scale_inv_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsDOScaleInvFormat:
+ std::memcpy(&p.do_scale_inv_format, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsMaxSeqlenQ:
+ std::memcpy(&p.max_seqlen_q, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsMaxSeqlenKV:
+ std::memcpy(&p.max_seqlen_kv, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsWorkspace:
+ std::memcpy(&p.workspace, buf, attr_size);
+ break;
+ case kNVTEFusedAttnBwdParamsStream:
+ std::memcpy(&p.stream, buf, attr_size);
+ break;
+ default:
+ NVTE_ERROR("Unsupported NVTEFusedAttnBwdParamsAttribute (got ", static_cast(attr), ")");
+ }
+}
diff --git a/transformer_engine/common/fused_attn/config_and_params.h b/transformer_engine/common/fused_attn/config_and_params.h
new file mode 100644
index 0000000000..ebc5b3eb07
--- /dev/null
+++ b/transformer_engine/common/fused_attn/config_and_params.h
@@ -0,0 +1,385 @@
+/*************************************************************************
+ * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+ *
+ * See LICENSE for license information.
+ ************************************************************************/
+
+/*! \file config_and_params.h
+ * \brief Internal objects for fused-attention config and parameter handles.
+ */
+
+#ifndef TRANSFORMER_ENGINE_COMMON_FUSED_ATTN_CONFIG_AND_PARAMS_H_
+#define TRANSFORMER_ENGINE_COMMON_FUSED_ATTN_CONFIG_AND_PARAMS_H_
+
+#include
+
+#include "common/common.h"
+#include "transformer_engine/fused_attn.h"
+
+namespace transformer_engine {
+namespace fused_attn {
+
+struct FusedAttnConfig {
+ // basic attention settings
+ bool is_training = true;
+ bool deterministic = false;
+ bool cuda_graph = false;
+ bool return_max_logit = false;
+ NVTE_Mask_Type attn_mask_type = NVTE_NO_MASK;
+ NVTE_Bias_Type bias_type = NVTE_NO_BIAS;
+ int64_t window_size_left = -1;
+ int64_t window_size_right = -1;
+ bool bottom_right_diagonal = true;
+ NVTE_Softmax_Type softmax_type = NVTE_VANILLA_SOFTMAX;
+ NVTEScalingMode scaling_mode = NVTE_DELAYED_TENSOR_SCALING;
+ float dropout = 0.0f;
+ float attn_scale = 1.0f;
+
+ // tensor types
+ NVTEDType qkv_dtype = kNVTEBFloat16;
+ NVTEDType o_dtype = kNVTEBFloat16;
+ NVTEDType do_dtype = kNVTEBFloat16;
+ NVTEDType dqkv_dtype = kNVTEBFloat16;
+
+ // tensor layouts
+ NVTE_QKV_Layout qkv_layout = NVTE_QKV_Layout_NOT_SET;
+ NVTE_QKV_Format o_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Format do_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Layout dqkv_layout = NVTE_QKV_Layout_NOT_SET;
+ NVTE_QKV_Format qkv_scale_inv_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Format do_scale_inv_format = NVTE_QKV_Format_NOT_SET;
+
+ // tensor dimensions
+ size_t batch_size = 0;
+ size_t num_attn_heads = 0;
+ size_t num_gqa_groups = 0;
+ size_t head_dim_qk = 0;
+ size_t head_dim_v = 0;
+ size_t max_seqlen_q = 0;
+ size_t max_seqlen_kv = 0;
+ size_t num_tokens_q = 0;
+ size_t num_tokens_kv = 0;
+
+ // paged KV dimensions
+ size_t num_pages_k = 0;
+ size_t num_pages_v = 0;
+ size_t page_size_k = 0;
+ size_t page_size_v = 0;
+ size_t max_pages_per_seq_k = 0;
+ size_t max_pages_per_seq_v = 0;
+
+ // bias dimensions
+ size_t bias_batch_size = 0;
+ size_t bias_num_heads = 0;
+ size_t bias_seqlen_q = 0;
+ size_t bias_seqlen_kv = 0;
+
+ // device ID: not part of attribute serialization, but part of operator< and used to
+ // differentiate graphs built for different devices in multi-GPU single-process runs
+ int device_id = -1;
+
+ // Internal-only fields: never part of attribute serialization, operator<, or the graph cache key.
+ // Filled by derive() or set by caller (i.e. is_forward). Added for convinence purposes and do not
+ // represent any graph properties.
+
+ // Direction to build the cuDNN graph for; steers make_cache_key() normalization.
+ bool is_forward = false;
+ // THD batch/token counts; make_cache_key() folds these into batch_size/max_seqlen_*.
+ size_t bucketed_batch_size = 0;
+ size_t bucketed_num_tokens_q = 0;
+ size_t bucketed_num_tokens_kv = 0;
+ // Uses cu_seqlens or actual_seqlens.
+ bool uses_cu_seqlens_directly = false;
+ // Convinence fields to avoid recompute.
+ NVTE_QKV_Format q_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Format kv_format = NVTE_QKV_Format_NOT_SET;
+ bool is_ragged_q = false;
+ bool is_ragged_kv = false;
+ bool is_paged_kv = false;
+ bool is_padding = false;
+ bool is_causal = false;
+ bool is_causal_bottom_right = false;
+
+ static constexpr size_t attr_sizes[] = {
+ // basic attention settings
+ sizeof(uint8_t), // is_training
+ sizeof(uint8_t), // deterministic
+ sizeof(uint8_t), // cuda_graph
+ sizeof(uint8_t), // return_max_logit
+ sizeof(NVTE_Mask_Type), // attn_mask_type
+ sizeof(NVTE_Bias_Type), // bias_type
+ sizeof(int64_t), // window_size_left
+ sizeof(int64_t), // window_size_right
+ sizeof(uint8_t), // bottom_right_diagonal
+ sizeof(NVTE_Softmax_Type), // softmax_type
+ sizeof(NVTEScalingMode), // scaling_mode
+ sizeof(float), // dropout
+ sizeof(float), // attn_scale
+ // tensor types
+ sizeof(NVTEDType), // qkv_dtype
+ sizeof(NVTEDType), // o_dtype
+ sizeof(NVTEDType), // do_dtype
+ sizeof(NVTEDType), // dqkv_dtype
+ // tensor layouts
+ sizeof(NVTE_QKV_Layout), // qkv_layout
+ sizeof(NVTE_QKV_Format), // o_format
+ sizeof(NVTE_QKV_Format), // do_format
+ sizeof(NVTE_QKV_Layout), // dqkv_layout
+ sizeof(NVTE_QKV_Format), // qkv_scale_inv_format
+ sizeof(NVTE_QKV_Format), // do_scale_inv_format
+ // tensor dimensions
+ sizeof(size_t), // batch_size
+ sizeof(size_t), // num_attn_heads
+ sizeof(size_t), // num_gqa_groups
+ sizeof(size_t), // head_dim_qk
+ sizeof(size_t), // head_dim_v
+ sizeof(size_t), // max_seqlen_q
+ sizeof(size_t), // max_seqlen_kv
+ sizeof(size_t), // num_tokens_q
+ sizeof(size_t), // num_tokens_kv
+ // paged KV dimensions
+ sizeof(size_t), // num_pages_k
+ sizeof(size_t), // num_pages_v
+ sizeof(size_t), // page_size_k
+ sizeof(size_t), // page_size_v
+ sizeof(size_t), // max_pages_per_seq_k
+ sizeof(size_t), // max_pages_per_seq_v
+ // bias dimensions
+ sizeof(size_t), // bias_batch_size
+ sizeof(size_t), // bias_num_heads
+ sizeof(size_t), // bias_seqlen_q
+ sizeof(size_t), // bias_seqlen_kv
+ };
+
+ bool operator<(const FusedAttnConfig &rhs) const {
+ return std::tie(is_training, deterministic, cuda_graph, return_max_logit, attn_mask_type,
+ bias_type, window_size_left, window_size_right, bottom_right_diagonal,
+ softmax_type, scaling_mode, dropout, attn_scale, qkv_dtype, o_dtype, do_dtype,
+ dqkv_dtype, qkv_layout, o_format, do_format, dqkv_layout, qkv_scale_inv_format,
+ do_scale_inv_format, batch_size, num_attn_heads, num_gqa_groups, head_dim_qk,
+ head_dim_v, max_seqlen_q, max_seqlen_kv, num_tokens_q, num_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_batch_size, bias_num_heads, bias_seqlen_q,
+ bias_seqlen_kv, device_id) <
+ std::tie(rhs.is_training, rhs.deterministic, rhs.cuda_graph, rhs.return_max_logit,
+ rhs.attn_mask_type, rhs.bias_type, rhs.window_size_left, rhs.window_size_right,
+ rhs.bottom_right_diagonal, rhs.softmax_type, rhs.scaling_mode, rhs.dropout,
+ rhs.attn_scale, rhs.qkv_dtype, rhs.o_dtype, rhs.do_dtype, rhs.dqkv_dtype,
+ rhs.qkv_layout, rhs.o_format, rhs.do_format, rhs.dqkv_layout,
+ rhs.qkv_scale_inv_format, rhs.do_scale_inv_format, rhs.batch_size,
+ rhs.num_attn_heads, rhs.num_gqa_groups, rhs.head_dim_qk, rhs.head_dim_v,
+ rhs.max_seqlen_q, rhs.max_seqlen_kv, rhs.num_tokens_q, rhs.num_tokens_kv,
+ 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_batch_size,
+ rhs.bias_num_heads, rhs.bias_seqlen_q, rhs.bias_seqlen_kv, rhs.device_id);
+ }
+
+ // Derive fields such as bucketed batch_size or num_tokens for THD, based on input fields
+ // that have been set by the caller.
+ void derive();
+
+ // Return a normalized copy of this config to be used as a key for the cuDNN graph cache.
+ // It drops fields that are invariant (e.g. attn_scale) or irrelevant (e.g. dO/dQKV dtypes
+ // and `deterministic` for forward, and `return_max_logit` for backward) to the corresponding graph.
+ // This helps avoid redundant graph builds and cache misses.
+ FusedAttnConfig make_cache_key() const;
+};
+
+inline const FusedAttnConfig *get_fused_attn_config(NVTEFusedAttnConfig config) {
+ NVTE_CHECK(config != nullptr, "NVTEFusedAttnConfig must not be NULL.");
+ return reinterpret_cast(config);
+}
+
+inline FusedAttnConfig *get_fused_attn_config_mutable(NVTEFusedAttnConfig config) {
+ NVTE_CHECK(config != nullptr, "NVTEFusedAttnConfig must not be NULL.");
+ return reinterpret_cast(config);
+}
+
+struct FusedAttnFwdParams {
+ NVTETensor Q = nullptr;
+ NVTETensor K = nullptr;
+ NVTETensor V = nullptr;
+ NVTETensor Bias = nullptr;
+ NVTETensor SoftmaxOffset = nullptr;
+ NVTETensor cu_seqlens_q = nullptr;
+ NVTETensor cu_seqlens_kv = nullptr;
+ NVTETensor cu_seqlens_q_padded = nullptr;
+ NVTETensor cu_seqlens_kv_padded = nullptr;
+ NVTETensor page_table_k = nullptr;
+ NVTETensor page_table_v = nullptr;
+ NVTETensor rng_state = nullptr;
+ NVTETensor S = nullptr;
+ NVTETensor O = nullptr;
+ NVTETensorPack *Aux_CTX_Tensors = nullptr;
+ bool is_training = true;
+ bool cuda_graph = false;
+ bool return_max_logit = false;
+ NVTE_Mask_Type attn_mask_type = NVTE_NO_MASK;
+ NVTE_Bias_Type bias_type = NVTE_NO_BIAS;
+ NVTE_Softmax_Type softmax_type = NVTE_VANILLA_SOFTMAX;
+ int64_t window_size_left = -1;
+ int64_t window_size_right = -1;
+ bool bottom_right_diagonal = true;
+ float dropout = 0.0f;
+ float attn_scale = 1.0f;
+ NVTE_QKV_Layout qkv_layout = NVTE_QKV_Layout_NOT_SET;
+ NVTE_QKV_Format o_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Format qkv_scale_inv_format = NVTE_QKV_Format_NOT_SET;
+ size_t max_seqlen_q = 0;
+ size_t max_seqlen_kv = 0;
+ NVTETensor workspace = nullptr;
+ cudaStream_t stream = nullptr;
+
+ static constexpr size_t attr_sizes[] = {
+ sizeof(NVTETensor), // Q
+ sizeof(NVTETensor), // K
+ sizeof(NVTETensor), // V
+ sizeof(NVTETensor), // Bias
+ sizeof(NVTETensor), // SoftmaxOffset
+ sizeof(NVTETensor), // cu_seqlens_q
+ sizeof(NVTETensor), // cu_seqlens_kv
+ sizeof(NVTETensor), // cu_seqlens_q_padded
+ sizeof(NVTETensor), // cu_seqlens_kv_padded
+ sizeof(NVTETensor), // page_table_k
+ sizeof(NVTETensor), // page_table_v
+ sizeof(NVTETensor), // rng_state
+ sizeof(NVTETensor), // S
+ sizeof(NVTETensor), // O
+ sizeof(NVTETensorPack *), // Aux_CTX_Tensors
+ sizeof(uint8_t), // is_training
+ sizeof(uint8_t), // cuda_graph
+ sizeof(uint8_t), // return_max_logit
+ sizeof(NVTE_Mask_Type), // attn_mask_type
+ sizeof(NVTE_Bias_Type), // bias_type
+ sizeof(NVTE_Softmax_Type), // softmax_type
+ sizeof(int64_t), // window_size_left
+ sizeof(int64_t), // window_size_right
+ sizeof(uint8_t), // bottom_right_diagonal
+ sizeof(float), // dropout
+ sizeof(float), // attn_scale
+ sizeof(NVTE_QKV_Layout), // qkv_layout
+ sizeof(NVTE_QKV_Format), // o_format
+ sizeof(NVTE_QKV_Format), // qkv_scale_inv_format
+ sizeof(size_t), // max_seqlen_q
+ sizeof(size_t), // max_seqlen_kv
+ sizeof(NVTETensor), // workspace
+ sizeof(cudaStream_t), // stream
+ };
+
+ // Build a FusedAttnConfig from the scalar "knobs" carried here (e.g. attn_mask_type, bias_type)
+ // and the fields derived from the tensor handles (dtypes, dims, scaling mode, paged-KV and bias
+ // broadcast shapes). Returns the real execution config; call FusedAttnConfig::make_cache_key on
+ // it to obtain the normalized cuDNN graph-cache key.
+ FusedAttnConfig make_config() const;
+};
+
+inline const FusedAttnFwdParams *get_fused_attn_fwd_params(NVTEFusedAttnFwdParams params) {
+ NVTE_CHECK(params != nullptr, "NVTEFusedAttnFwdParams must not be NULL.");
+ return reinterpret_cast(params);
+}
+
+inline FusedAttnFwdParams *get_fused_attn_fwd_params_mutable(NVTEFusedAttnFwdParams params) {
+ NVTE_CHECK(params != nullptr, "NVTEFusedAttnFwdParams must not be NULL.");
+ return reinterpret_cast(params);
+}
+
+struct FusedAttnBwdParams {
+ NVTETensor Q = nullptr;
+ NVTETensor K = nullptr;
+ NVTETensor V = nullptr;
+ NVTETensor O = nullptr;
+ NVTETensor dO = nullptr;
+ NVTETensor S = nullptr;
+ NVTETensor dP = nullptr;
+ const NVTETensorPack *Aux_CTX_Tensors = nullptr;
+ NVTETensor dQ = nullptr;
+ NVTETensor dK = nullptr;
+ NVTETensor dV = nullptr;
+ NVTETensor dBias = nullptr;
+ NVTETensor dSoftmaxOffset = nullptr;
+ NVTETensor cu_seqlens_q = nullptr;
+ NVTETensor cu_seqlens_kv = nullptr;
+ NVTETensor cu_seqlens_q_padded = nullptr;
+ NVTETensor cu_seqlens_kv_padded = nullptr;
+ bool cuda_graph = false;
+ bool deterministic = false;
+ NVTE_Mask_Type attn_mask_type = NVTE_NO_MASK;
+ NVTE_Bias_Type bias_type = NVTE_NO_BIAS;
+ NVTE_Softmax_Type softmax_type = NVTE_VANILLA_SOFTMAX;
+ int64_t window_size_left = -1;
+ int64_t window_size_right = -1;
+ bool bottom_right_diagonal = true;
+ float dropout = 0.0f;
+ float attn_scale = 1.0f;
+ NVTE_QKV_Layout qkv_layout = NVTE_QKV_Layout_NOT_SET;
+ NVTE_QKV_Format o_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Format do_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Layout dqkv_layout = NVTE_QKV_Layout_NOT_SET;
+ NVTE_QKV_Format qkv_scale_inv_format = NVTE_QKV_Format_NOT_SET;
+ NVTE_QKV_Format do_scale_inv_format = NVTE_QKV_Format_NOT_SET;
+ size_t max_seqlen_q = 0;
+ size_t max_seqlen_kv = 0;
+ NVTETensor workspace = nullptr;
+ cudaStream_t stream = nullptr;
+
+ static constexpr size_t attr_sizes[] = {
+ sizeof(NVTETensor), // Q
+ sizeof(NVTETensor), // K
+ sizeof(NVTETensor), // V
+ sizeof(NVTETensor), // O
+ sizeof(NVTETensor), // dO
+ sizeof(NVTETensor), // S
+ sizeof(NVTETensor), // dP
+ sizeof(const NVTETensorPack *), // Aux_CTX_Tensors
+ sizeof(NVTETensor), // dQ
+ sizeof(NVTETensor), // dK
+ sizeof(NVTETensor), // dV
+ sizeof(NVTETensor), // dBias
+ sizeof(NVTETensor), // dSoftmaxOffset
+ sizeof(NVTETensor), // cu_seqlens_q
+ sizeof(NVTETensor), // cu_seqlens_kv
+ sizeof(NVTETensor), // cu_seqlens_q_padded
+ sizeof(NVTETensor), // cu_seqlens_kv_padded
+ sizeof(uint8_t), // cuda_graph
+ sizeof(uint8_t), // deterministic
+ sizeof(NVTE_Mask_Type), // attn_mask_type
+ sizeof(NVTE_Bias_Type), // bias_type
+ sizeof(NVTE_Softmax_Type), // softmax_type
+ sizeof(int64_t), // window_size_left
+ sizeof(int64_t), // window_size_right
+ sizeof(uint8_t), // bottom_right_diagonal
+ sizeof(float), // dropout
+ sizeof(float), // attn_scale
+ sizeof(NVTE_QKV_Layout), // qkv_layout
+ sizeof(NVTE_QKV_Format), // o_format
+ sizeof(NVTE_QKV_Format), // do_format
+ sizeof(NVTE_QKV_Layout), // dqkv_layout
+ sizeof(NVTE_QKV_Format), // qkv_scale_inv_format
+ sizeof(NVTE_QKV_Format), // do_scale_inv_format
+ sizeof(size_t), // max_seqlen_q
+ sizeof(size_t), // max_seqlen_kv
+ sizeof(NVTETensor), // workspace
+ sizeof(cudaStream_t), // stream
+ };
+
+ // Build a FusedAttnConfig from the scalar "knobs" carried here (e.g. attn_mask_type, bias_type)
+ // and the fields derived from the tensor handles (e.g. dtypes, dims, scaling mode and bias broadcast
+ // shape). Returns the real execution config; call FusedAttnConfig::make_cache_key on it to
+ // obtain the normalized cuDNN graph-cache key.
+ FusedAttnConfig make_config() const;
+};
+
+inline const FusedAttnBwdParams *get_fused_attn_bwd_params(NVTEFusedAttnBwdParams params) {
+ NVTE_CHECK(params != nullptr, "NVTEFusedAttnBwdParams must not be NULL.");
+ return reinterpret_cast(params);
+}
+
+inline FusedAttnBwdParams *get_fused_attn_bwd_params_mutable(NVTEFusedAttnBwdParams params) {
+ NVTE_CHECK(params != nullptr, "NVTEFusedAttnBwdParams must not be NULL.");
+ return reinterpret_cast(params);
+}
+
+} // namespace fused_attn
+} // namespace transformer_engine
+
+#endif // TRANSFORMER_ENGINE_COMMON_FUSED_ATTN_CONFIG_AND_PARAMS_H_
diff --git a/transformer_engine/common/fused_attn/fused_attn.cpp b/transformer_engine/common/fused_attn/fused_attn.cpp
index fc21771297..b3b9922abf 100644
--- a/transformer_engine/common/fused_attn/fused_attn.cpp
+++ b/transformer_engine/common/fused_attn/fused_attn.cpp
@@ -10,6 +10,7 @@
#include "../cudnn_utils.h"
#include "../util/cuda_runtime.h"
#include "../util/system.h"
+#include "config_and_params.h"
#include "fused_attn_f16_arbitrary_seqlen.h"
#include "fused_attn_fp8.h"
#include "utils.h"
@@ -225,308 +226,215 @@ NVTE_QKV_Format nvte_get_kv_format(NVTE_QKV_Layout qkv_layout) {
}
}
-// 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) {
+namespace {
+
+// The per-thread storage for the diagnostic string; it's re-used (cleared + re-populated)
+// on every call to nvte_get_fused_attn_backend_v2 on the same thread.
+thread_local std::string fused_attn_backend_message_buffer;
+
+// Stash `reason` in the thread-local buffer and, if the caller asked for a diagnostic,
+// publish a NUL-terminated pointer to it via `*message`. Safe to call with `message == nullptr`.
+void set_message(const char **message, std::string reason) {
+ if (message == nullptr) return;
+ fused_attn_backend_message_buffer = std::move(reason);
+ *message = fused_attn_backend_message_buffer.c_str();
+}
+
+} // namespace
+
+// select a backend for fused attention; the diagnostic message is based on the first failure, not cumulative.
+NVTE_Fused_Attn_Backend nvte_get_fused_attn_backend_v2(NVTEFusedAttnConfig config,
+ const char **message) {
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);
- NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(qkv_layout);
- auto cudnn_runtime_version = cudnnGetVersion();
+ using namespace transformer_engine::fused_attn;
+ const FusedAttnConfig &cfg = *get_fused_attn_config(config);
+ set_message(message, "");
- // For ragged offsets we only support 32-bit prior to cuDNN 9.5
- // Only used when THD format is requested.
+ cudnnHandle_t handle = cudnnExecutionPlanManager::Instance().GetHandle();
+ const NVTE_QKV_Format qkv_format = nvte_get_qkv_format(cfg.qkv_layout);
+ const NVTE_QKV_Layout_Group layout_group = nvte_get_qkv_layout_group(cfg.qkv_layout);
+ const auto cudnn_runtime_version = cudnnGetVersion();
+
+ // THD + 64-bit ragged offsets require cuDNN >= 9.5
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)) ||
- // 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}
- ((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))) &&
- !requires_64bit_ragged_offset &&
- // 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.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 ||
- (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;
+ (qkv_format == NVTE_THD &&
+ fused_attn::get_ragged_offset_dtype(layout_group, cfg.num_attn_heads, cfg.num_gqa_groups,
+ cfg.max_seqlen_q, cfg.max_seqlen_kv, cfg.head_dim_qk,
+ cfg.head_dim_v) == DType::kInt64);
+ if (requires_64bit_ragged_offset && cudnn_runtime_version < 90500) {
+ set_message(message,
+ "Configuration requires 64-bit ragged offsets, which require "
+ "cuDNN >= 9.5.");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+ }
+
+ // THD requires padding-style mask
+ if (qkv_format == NVTE_QKV_Format::NVTE_THD &&
+ cfg.attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_MASK &&
+ cfg.attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK &&
+ cfg.attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK) {
+ set_message(message,
+ "THD format requires PADDING / PADDING_CAUSAL / PADDING_CAUSAL_BOTTOM_RIGHT mask.");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+ }
+
+ // cuDNN does not support pre-scale bias
+ if (cfg.bias_type == NVTE_Bias_Type::NVTE_PRE_SCALE_BIAS) {
+ set_message(message, "Fused attention does not support pre-scale bias.");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+ }
+
+ const bool is_fp8 =
+ (cfg.qkv_dtype == NVTEDType::kNVTEFloat8E4M3 || cfg.qkv_dtype == NVTEDType::kNVTEFloat8E5M2);
+ const bool is_f16_or_bf16 =
+ (cfg.qkv_dtype == NVTEDType::kNVTEFloat16 || cfg.qkv_dtype == NVTEDType::kNVTEBFloat16);
+
+ if (is_fp8) {
+ if (cfg.return_max_logit) {
+ set_message(message, "FP8 fused attention does not support return_max_logit=True.");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
}
- if (flag_arb) {
- backend = NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen;
+ if (qkv_format != NVTE_QKV_Format::NVTE_BSHD && qkv_format != NVTE_QKV_Format::NVTE_SBHD &&
+ qkv_format != NVTE_QKV_Format::NVTE_BHSD) {
+ set_message(message, "FP8 fused attention supports BSHD/SBHD/BHSD formats, found " +
+ std::to_string(static_cast(qkv_format)) + ".");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
}
- 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;
+ std::string fwd_reason = is_supported_fp8_fwd(cfg, handle);
+ if (!fwd_reason.empty()) {
+ set_message(message, std::move(fwd_reason));
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
}
- 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 (cfg.is_training && !cfg.is_forward) {
+ std::string bwd_reason = is_supported_fp8_bwd(cfg, handle);
+ if (!bwd_reason.empty()) {
+ set_message(message, std::move(bwd_reason));
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+ }
}
- if ((cudnn_runtime_version <= 91500) && is_training &&
+ return NVTE_Fused_Attn_Backend::NVTE_FP8;
+ }
+
+ if (is_f16_or_bf16) {
+ if (cudnn_runtime_version <= 91500 && cfg.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;
+ (cfg.max_seqlen_kv % 128 != 0) && cfg.cuda_graph &&
+ cfg.attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_MASK &&
+ cfg.attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_MASK &&
+ cfg.attn_mask_type != NVTE_Mask_Type::NVTE_PADDING_CAUSAL_BOTTOM_RIGHT_MASK) {
+ set_message(message, "Known cuDNN <= 9.15 issue with CUDA graph. Please upgrade cuDNN.");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+ }
+ std::string fwd_reason = is_supported_f16_fwd(cfg, handle);
+ if (!fwd_reason.empty()) {
+ set_message(message, std::move(fwd_reason));
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
}
- 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;
- }
+ if (cfg.is_training && !cfg.is_forward) {
+ std::string bwd_reason = is_supported_f16_bwd(cfg, handle);
+ if (!bwd_reason.empty()) {
+ set_message(message, std::move(bwd_reason));
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
}
}
+ return NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen;
+ }
+
+ set_message(message, "Unsupported QKV dtype qkv_dtype=" + std::to_string(cfg.qkv_dtype) + " .");
+ return NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+}
+
+// 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) {
+ transformer_engine::fused_attn::FusedAttnConfig cfg{};
+ cfg.qkv_layout = qkv_layout;
+ cfg.bias_type = bias_type;
+ cfg.attn_mask_type = attn_mask_type;
+ cfg.softmax_type = softmax_type;
+ cfg.dropout = dropout;
+ cfg.max_seqlen_q = max_seqlen_q;
+ cfg.max_seqlen_kv = max_seqlen_kv;
+ cfg.window_size_left = window_size_left;
+ cfg.window_size_right = window_size_right;
+ cfg.cuda_graph = cuda_graph;
+ NVTE_CHECK(q_dtype == kv_dtype, "Q and KV must have the same data type.");
+ cfg.qkv_dtype = q_dtype;
+ cfg.o_dtype = q_dtype;
+ cfg.do_dtype = q_dtype;
+ cfg.dqkv_dtype = q_dtype;
+ cfg.num_attn_heads = num_attn_heads;
+ cfg.num_gqa_groups = num_gqa_groups;
+ cfg.head_dim_qk = head_dim_qk;
+ cfg.head_dim_v = head_dim_v;
+ cfg.is_training = is_training;
+ cfg.return_max_logit = return_max_logit;
+ cfg.deterministic = deterministic;
+ // fill in missing fields so it doesn't always return NVTE_No_Backend
+ cfg.batch_size = 1;
+ cfg.o_format = nvte_get_q_format(qkv_layout);
+ cfg.do_format = cfg.o_format;
+ cfg.dqkv_layout = qkv_layout;
+ if (bias_type == NVTE_Bias_Type::NVTE_POST_SCALE_BIAS) {
+ cfg.bias_batch_size = cfg.batch_size;
+ cfg.bias_num_heads = num_attn_heads;
+ cfg.bias_seqlen_q = max_seqlen_q;
+ cfg.bias_seqlen_kv = max_seqlen_kv;
+ }
+
+ return nvte_get_fused_attn_backend_v2(reinterpret_cast(&cfg),
+ /*message=*/nullptr);
+}
+
+// fused attention forward
+void nvte_fused_attn_fwd_v2(NVTEFusedAttnFwdParams params) {
+ NVTE_API_CALL(nvte_fused_attn_fwd_v2);
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+ const FusedAttnFwdParams &p = *get_fused_attn_fwd_params(params);
+ const Tensor *input_cu_seqlens_q = convertNVTETensorCheck(p.cu_seqlens_q);
+ const Tensor *input_cu_seqlens_kv = convertNVTETensorCheck(p.cu_seqlens_kv);
+ const Tensor *input_cu_seqlens_q_padded = convertNVTETensorCheck(p.cu_seqlens_q_padded);
+ const Tensor *input_cu_seqlens_kv_padded = convertNVTETensorCheck(p.cu_seqlens_kv_padded);
+ const Tensor *input_page_table_k = convertNVTETensorCheck(p.page_table_k);
+ const Tensor *input_page_table_v = convertNVTETensorCheck(p.page_table_v);
+ const Tensor *input_rng_state = convertNVTETensorCheck(p.rng_state);
+ const Tensor *input_Q = convertNVTETensorCheck(p.Q);
+ const Tensor *input_K = convertNVTETensorCheck(p.K);
+ const Tensor *input_V = convertNVTETensorCheck(p.V);
+ const Tensor *input_Bias = convertNVTETensorCheck(p.Bias);
+ const Tensor *input_SoftmaxOffset = convertNVTETensorCheck(p.SoftmaxOffset);
+ Tensor *input_output_S = convertNVTETensorCheck(p.S);
+ Tensor *output_O = convertNVTETensorCheck(p.O);
+ Tensor *wkspace = convertNVTETensor(p.workspace);
+
+ auto handle = cudnnExecutionPlanManager::Instance().GetHandle();
+ FusedAttnConfig cfg = p.make_config();
+ const char *fused_attn_reject_reason = nullptr;
+ NVTE_Fused_Attn_Backend fused_attention_backend = nvte_get_fused_attn_backend_v2(
+ reinterpret_cast(&cfg), &fused_attn_reject_reason);
+
+ if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_F16_arbitrary_seqlen) {
+ fused_attn_arbitrary_seqlen_fwd(cfg, input_Q, input_K, input_V, input_Bias, input_SoftmaxOffset,
+ output_O, p.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, p.stream, handle);
+ } else if (fused_attention_backend == NVTE_Fused_Attn_Backend::NVTE_FP8) {
+ fused_attn_fp8_fwd(cfg, input_Q, input_K, input_V, input_SoftmaxOffset, input_output_S,
+ output_O, p.Aux_CTX_Tensors, input_cu_seqlens_q, input_cu_seqlens_kv,
+ input_rng_state, wkspace, p.stream, handle);
} else {
- backend = NVTE_Fused_Attn_Backend::NVTE_No_Backend;
+ const char *reject_reason =
+ (fused_attn_reject_reason != nullptr && fused_attn_reject_reason[0] != '\0')
+ ? fused_attn_reject_reason
+ : "no cuDNN fused-attention backend supports the requested parameters";
+ NVTE_ERROR("Fused attention is not supported for this configuration: ", reject_reason);
}
- return backend;
}
// NVTE fused attention FWD with separate Q, K and V
@@ -545,98 +453,114 @@ void nvte_fused_attn_fwd(const NVTETensor Q, const NVTETensor K, const NVTETenso
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;
- }
+ transformer_engine::fused_attn::FusedAttnFwdParams p{};
+ p.Q = Q;
+ p.K = K;
+ p.V = V;
+ p.Bias = Bias;
+ p.SoftmaxOffset = SoftmaxOffset;
+ p.S = S;
+ p.O = O;
+ p.Aux_CTX_Tensors = Aux_CTX_Tensors;
+ p.cu_seqlens_q = cu_seqlens_q;
+ p.cu_seqlens_kv = cu_seqlens_kv;
+ p.cu_seqlens_q_padded = cu_seqlens_q_padded;
+ p.cu_seqlens_kv_padded = cu_seqlens_kv_padded;
+ p.page_table_k = page_table_k;
+ p.page_table_v = page_table_v;
+ p.rng_state = rng_state;
+ p.max_seqlen_q = max_seqlen_q;
+ p.max_seqlen_kv = max_seqlen_kv;
+ p.is_training = is_training;
+ p.return_max_logit = return_max_logit;
+ p.cuda_graph = cuda_graph;
+ p.attn_scale = attn_scale;
+ p.dropout = dropout;
+ p.qkv_layout = qkv_layout;
+ p.o_format = o_format;
+ p.qkv_scale_inv_format = qkv_scale_inv_format;
+ p.bias_type = bias_type;
+ p.attn_mask_type = attn_mask_type;
+ p.softmax_type = softmax_type;
+ p.window_size_left = window_size_left;
+ p.window_size_right = window_size_right;
+ p.bottom_right_diagonal = bottom_right_diagonal;
+ p.workspace = workspace;
+ p.stream = stream;
+ nvte_fused_attn_fwd_v2(reinterpret_cast(&p));
+}
- 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];
- }
- }
+// fused attention backward
+void nvte_fused_attn_bwd_v2(NVTEFusedAttnBwdParams params) {
+ NVTE_API_CALL(nvte_fused_attn_bwd_v2);
+ using namespace transformer_engine;
+ using namespace transformer_engine::fused_attn;
+ const FusedAttnBwdParams &p = *get_fused_attn_bwd_params(params);
+ const Tensor *input_cu_seqlens_q = convertNVTETensorCheck(p.cu_seqlens_q);
+ const Tensor *input_cu_seqlens_kv = convertNVTETensorCheck(p.cu_seqlens_kv);
+ const Tensor *input_cu_seqlens_q_padded = convertNVTETensorCheck(p.cu_seqlens_q_padded);
+ const Tensor *input_cu_seqlens_kv_padded = convertNVTETensorCheck(p.cu_seqlens_kv_padded);
+ const Tensor *input_Q = convertNVTETensorCheck(p.Q);
+ const Tensor *input_K = convertNVTETensorCheck(p.K);
+ const Tensor *input_V = convertNVTETensorCheck(p.V);
+ const Tensor *input_O = convertNVTETensorCheck(p.O);
+ const Tensor *input_dO = convertNVTETensorCheck(p.dO);
+ const Tensor *input_S = convertNVTETensorCheck(p.S);
+ Tensor *input_output_dP = convertNVTETensorCheck(p.dP);
+ Tensor *output_dQ = convertNVTETensorCheck(p.dQ);
+ Tensor *output_dK = convertNVTETensorCheck(p.dK);
+ Tensor *output_dV = convertNVTETensorCheck(p.dV);
+ Tensor *output_dBias = convertNVTETensorCheck(p.dBias);
+ Tensor *output_dSoftmaxOffset = convertNVTETensorCheck(p.dSoftmaxOffset);
+ Tensor *wkspace = convertNVTETensor(p.workspace);
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);
+ FusedAttnConfig cfg = p.make_config();
+ const char *fused_attn_reject_reason = nullptr;
+ NVTE_Fused_Attn_Backend fused_attention_backend = nvte_get_fused_attn_backend_v2(
+ reinterpret_cast(&cfg), &fused_attn_reject_reason);
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);
+ size_t i = 0;
+ Tensor *output_S = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ Tensor *input_rng_state = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ Tensor *input_Bias = nullptr, *input_SoftmaxOffset = nullptr;
+ if ((p.bias_type != NVTE_NO_BIAS) && (p.bias_type != NVTE_ALIBI)) {
+ input_Bias = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ }
+ if (p.softmax_type != NVTE_VANILLA_SOFTMAX) {
+ input_SoftmaxOffset = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ }
+ fused_attn_arbitrary_seqlen_bwd(
+ cfg, 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, p.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, 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_rng_state, wkspace, stream, handle);
+ size_t i = 0;
+ const Tensor *input_M = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ const Tensor *input_rng_state = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ const Tensor *input_SoftmaxOffset = nullptr;
+ if (p.softmax_type != NVTE_VANILLA_SOFTMAX) {
+ input_SoftmaxOffset = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ }
+ const Tensor *input_dO_f16 = nullptr;
+ if (input_dO->scaling_mode == NVTE_MXFP8_1D_SCALING) {
+ input_dO_f16 = convertNVTETensorCheck(p.Aux_CTX_Tensors->tensors[i++]);
+ }
+ fused_attn_fp8_bwd(cfg, 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_rng_state, wkspace, p.stream, handle);
} else {
- NVTE_ERROR("Invalid combination of data type and sequence length for fused attention. \n");
+ const char *reject_reason =
+ (fused_attn_reject_reason != nullptr && fused_attn_reject_reason[0] != '\0')
+ ? fused_attn_reject_reason
+ : "no cuDNN fused-attention backend supports the requested parameters";
+ NVTE_ERROR("Fused attention is not supported for this configuration: ", reject_reason);
}
}
+
// 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,
@@ -654,92 +578,45 @@ void nvte_fused_attn_bwd(const NVTETensor Q, const NVTETensor K, const NVTETenso
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, 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_rng_state, wkspace, stream, handle);
- } else {
- NVTE_ERROR("Invalid combination of data type and sequence length for fused attention. \n");
- }
+ transformer_engine::fused_attn::FusedAttnBwdParams p{};
+ p.Q = Q;
+ p.K = K;
+ p.V = V;
+ p.O = O;
+ p.dO = dO;
+ p.S = S;
+ p.dP = dP;
+ p.Aux_CTX_Tensors = Aux_CTX_Tensors;
+ p.dQ = dQ;
+ p.dK = dK;
+ p.dV = dV;
+ p.dBias = dBias;
+ p.dSoftmaxOffset = dSoftmaxOffset;
+ p.cu_seqlens_q = cu_seqlens_q;
+ p.cu_seqlens_kv = cu_seqlens_kv;
+ p.cu_seqlens_q_padded = cu_seqlens_q_padded;
+ p.cu_seqlens_kv_padded = cu_seqlens_kv_padded;
+ p.max_seqlen_q = max_seqlen_q;
+ p.max_seqlen_kv = max_seqlen_kv;
+ p.attn_scale = attn_scale;
+ p.dropout = dropout;
+ p.qkv_layout = qkv_layout;
+ p.o_format = o_format;
+ p.do_format = do_format;
+ p.dqkv_layout = dqkv_layout;
+ p.qkv_scale_inv_format = qkv_scale_inv_format;
+ p.do_scale_inv_format = do_scale_inv_format;
+ p.bias_type = bias_type;
+ p.attn_mask_type = attn_mask_type;
+ p.softmax_type = softmax_type;
+ p.window_size_left = window_size_left;
+ p.window_size_right = window_size_right;
+ p.bottom_right_diagonal = bottom_right_diagonal;
+ p.deterministic = deterministic;
+ p.cuda_graph = cuda_graph;
+ p.workspace = workspace;
+ p.stream = stream;
+ nvte_fused_attn_bwd_v2(reinterpret_cast(&p));
}
uint32_t nvte_get_runtime_num_segments(NVTETensor cu_seqlen, NVTETensor workspace, size_t len,
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
index bf34758a35..b7c7a349af 100644
--- a/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu
+++ b/transformer_engine/common/fused_attn/fused_attn_f16_arbitrary_seqlen.cu
@@ -10,6 +10,7 @@
#include
#include