Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
101 commits
Select commit Hold shift + click to select a range
7765b76
feat(attention): cuDNN FROST kernel wrapper for head_dim in (256, 512]
nvegesna-netizen Sep 16, 2026
3319826
feat(attention): select and dispatch FrostAttention from DotProductAt…
nvegesna-netizen Sep 16, 2026
d60b3fd
feat(attention): context parallelism for FROST across p2p, all_gather…
nvegesna-netizen Sep 16, 2026
62ffe39
test(attention): CP coverage for FrostAttention at head_dim 512
nvegesna-netizen Sep 16, 2026
fe72e4a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Sep 16, 2026
a8957ae
test(attention): run the FrostAttention CP configs from pytest
nvegesna-netizen Sep 16, 2026
bdbb36a
fix(attention): gate FROST on the cuDNN Frontend version and key plan…
nvegesna-netizen Sep 16, 2026
e30ea42
fix(attention): repair the FROST plan-cache key arity and harden the …
nvegesna-netizen Sep 16, 2026
0957f71
fix(attention): reject a v that does not match k, and refine the vers…
nvegesna-netizen Sep 16, 2026
438b4da
fix(attention): update the mixed-THD backend unpack for the new retur…
nvegesna-netizen Sep 16, 2026
c0d0713
fix(attention): bind a cuDNN stream and close the remaining silent-wr…
nvegesna-netizen Sep 16, 2026
7504bff
fix(attention): JIT-compile FROST plans under the device their handle…
nvegesna-netizen Sep 16, 2026
85a0f51
test(attention): anchor FROST numerics to an fp32 reference, not to i…
nvegesna-netizen Sep 16, 2026
c5f9cd2
feat(attention): honour deterministic on the FROST path, and document…
nvegesna-netizen Sep 16, 2026
a97496e
docs(attention): scope the FROST exclusivity claim to context paralle…
nvegesna-netizen Sep 16, 2026
5a675c2
fix(attention): drop a duplicate deterministic parameter on the fused…
nvegesna-netizen Sep 16, 2026
e82981f
fix(attention): decline FROST when determinism is required
nvegesna-netizen Sep 16, 2026
3edba85
test(attention): make the FROST oracle float64, since an fp32 one is …
nvegesna-netizen Sep 16, 2026
064396e
feat(attention): express FROST masking as a diagonal band, adding sli…
nvegesna-netizen Sep 16, 2026
441dee4
fix(attention): carry the sliding window through a2a, and decline it …
nvegesna-netizen Sep 16, 2026
9770ca5
test(attention): cover the sliding window in backward, at its boundar…
nvegesna-netizen Sep 16, 2026
ad9dfdc
fix(attention): let the CP sliding-window asserts know FROST exists
nvegesna-netizen Sep 16, 2026
3eea2f9
docs(attention): correct the claimed cuDNN import-ordering hazard
nvegesna-netizen Sep 16, 2026
e062e8d
test(attention): apply the window for every mask type in the reference
nvegesna-netizen Sep 16, 2026
dd0033c
docs(attention): justify the p2p sliding-window decline from the ring…
nvegesna-netizen Sep 16, 2026
a82903b
fix(attention): bind the FROST flag on the ONNX path, decline what wa…
nvegesna-netizen Sep 16, 2026
591955d
fix(attention): read qkv_type from attention_params, not the rebound …
nvegesna-netizen Sep 16, 2026
6832a9b
test(attention): cover the ONNX-export branch on hardware that can ru…
nvegesna-netizen Sep 16, 2026
ca6c95a
feat(attention): allow FrostAttention with cp_comm_type=a2a+p2p
nvegesna-netizen Sep 17, 2026
b4cdcb4
fix(attention): handle a list-valued cp_group in FrostAttention.forward
nvegesna-netizen Sep 17, 2026
9a8b474
test(attention): cover fp16 in the backward and under context paralle…
nvegesna-netizen Sep 17, 2026
9790575
docs(attention): narrow the FrostAttention availability claim
nvegesna-netizen Sep 17, 2026
9a548fb
docs(attention): mark FrostAttention experimental and trim review com…
nvegesna-netizen Sep 21, 2026
fe34dcc
docs(attention): mark the FrostAttention backend experimental in envvars
nvegesna-netizen Sep 21, 2026
9cc5a40
Merge branch 'main' into nvegesna/te-frost-d512-cp
nvegesna-netizen Sep 21, 2026
c5d8825
Merge remote-tracking branch 'origin/main' into nvegesna/te-frost-d51…
nvegesna-netizen Sep 22, 2026
7bbccff
refactor(attention): extract the shared cuDNN pygraph plumbing
nvegesna-netizen Oct 1, 2026
42fa333
feat(attention): let the flex cuDNN graphs carry a diagonal-band mask
nvegesna-netizen Oct 1, 2026
084f528
fix(attention): stop preparing the FROST graphs twice
nvegesna-netizen Oct 1, 2026
1f43e88
fix(attention): keep the flex builder call shape, and frost's cudnn h…
nvegesna-netizen Oct 1, 2026
469bad9
fix(attention): enable the FROST engines whichever backend imports cu…
nvegesna-netizen Oct 1, 2026
626bde2
test(attention): cover the flex mask_spec path on CPU
nvegesna-netizen Oct 1, 2026
40b8a2a
fix(attention): say why a pinned cuDNN engine declined the graph
nvegesna-netizen Oct 1, 2026
812d475
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 1, 2026
c981a0d
test(attention): skip the new cuDNN-frontend tests when the package i…
nvegesna-netizen Oct 1, 2026
73a205b
fix(attention): stop flex graphs running on a FROST engine
nvegesna-netizen Oct 1, 2026
c46dffc
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 1, 2026
2cfd6ff
test(attention): skip the FROST switch test where it cannot detect an…
nvegesna-netizen Oct 1, 2026
636d7a6
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 1, 2026
8ae1663
Merge branch 'main' into nvegesna/te-frost-d512-cp
nvegesna-netizen Oct 4, 2026
69a253b
fix(attention): index the FROST p2p results by the alternating slot
nvegesna-netizen Oct 4, 2026
7a55819
refactor(attention): make FROST a FusedAttention sub-backend
nvegesna-netizen Oct 6, 2026
9eb866d
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 6, 2026
0d1c4bf
fix(attention): decline FROST where the diagonal anchor is ambiguous
nvegesna-netizen Oct 6, 2026
91a37f1
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 6, 2026
521c9a8
refactor(attention): make frost_attention.py standalone again
nvegesna-netizen Oct 6, 2026
76f59c7
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 6, 2026
94577ea
fix(attention): return max_logit by index from the CP p2p fused step
nvegesna-netizen Oct 6, 2026
2cb9db1
fix(attention): put the new CP grad slot at the right index
nvegesna-netizen Oct 6, 2026
d400eea
test(attention): count FROST as a comparable fused sub-backend
nvegesna-netizen Oct 6, 2026
0f2c781
Merge branch 'main' into nvegesna/te-frost-d512-cp
nvegesna-netizen Oct 6, 2026
dbc298e
test(attention): drop the comment on the sub-backend list
nvegesna-netizen Oct 6, 2026
aec8b4b
fix(attention): bar the cuDNN FROST engines from flex attention graphs
nvegesna-netizen Oct 6, 2026
5d506d8
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 6, 2026
df5ae6e
refactor(attention): trim the FROST comments to the repo's register
nvegesna-netizen Oct 6, 2026
597e593
refactor(attention): shorten the max_logit unpack comment
nvegesna-netizen Oct 7, 2026
6b1a6d8
feat(attention): serve asymmetric head_dim on the FROST sub-backend
nvegesna-netizen Oct 7, 2026
2fca15d
revert(attention): move the flex FROST engine bar out of this change
nvegesna-netizen Oct 7, 2026
d04ece9
docs(attention): drop the stale symmetric head_dim claim for FROST
nvegesna-netizen Oct 7, 2026
0c4601a
fix(attention): resolve the FROST diagonal anchor the way the dispatc…
nvegesna-netizen Oct 7, 2026
c1c3a86
refactor(attention): give the cuDNN pygraph state one owner
nvegesna-netizen Oct 7, 2026
2e6787d
refactor(attention): describe the FROST tensors instead of permuting …
nvegesna-netizen Oct 7, 2026
7bc01fb
refactor(attention): drop the FROST kernel API, keep the sub-backend one
nvegesna-netizen Oct 7, 2026
314a91d
refactor(attention): share the graph cache mechanism and the per-tens…
nvegesna-netizen Oct 7, 2026
2931c95
fix(attention): scope the uncacheable graph build to its device too
nvegesna-netizen Oct 7, 2026
e0f2690
fix(attention): bar the cuDNN FROST engines by default, not on request
nvegesna-netizen Oct 7, 2026
631d8be
docs(attention): bring the new docstrings into the repo's register
nvegesna-netizen Oct 7, 2026
6654a68
fix(attention): decline asymmetric head_dim again, measured on B200
nvegesna-netizen Oct 7, 2026
16915e9
fix(attention): close two guards the fused path was bypassing
nvegesna-netizen Oct 7, 2026
86609ba
refactor(attention): move the diagonal-band translation to cudnn_pygraph
nvegesna-netizen Oct 7, 2026
639cdad
Revert "refactor(attention): move the diagonal-band translation to cu…
nvegesna-netizen Oct 7, 2026
0d25aa3
refactor(attention): share the graph execution step
nvegesna-netizen Oct 8, 2026
8c6c36a
docs(attention): mark the FROST engine bar as a workaround with an ex…
nvegesna-netizen Oct 8, 2026
f01a0ee
docs(attention): say what actually verifies the FROST engine bar
nvegesna-netizen Oct 8, 2026
5019bf3
test(attention): let a lane require the flex cuDNN tests to actually run
nvegesna-netizen Oct 8, 2026
598db0d
test(attention): trim the flex required-guard comment
nvegesna-netizen Oct 8, 2026
8176c30
test(attention): trim the FROST guard comments to match the siblings
nvegesna-netizen Oct 8, 2026
a88556b
test(attention): cut the FROST suite's CI cost where it buys nothing
nvegesna-netizen Oct 8, 2026
c01f48b
test(attention): put the cp_comm_type selector test back
nvegesna-netizen Oct 8, 2026
adbc951
Revert "test(attention): put the cp_comm_type selector test back"
nvegesna-netizen Oct 8, 2026
6c66bca
test(attention): put the CP sequence length back, it saved nothing
nvegesna-netizen Oct 8, 2026
dbb9500
test(attention): trim the CP seqlen comment
nvegesna-netizen Oct 8, 2026
029e0d7
test(attention): drop the CP seqlen note
nvegesna-netizen Oct 8, 2026
f23a850
test(attention): point the single a2a+p2p case at the shortest model
nvegesna-netizen Oct 8, 2026
b66c078
refactor(attention): build the cuDNN SDPA graphs through cudnn_pygraph
nvegesna-netizen Oct 8, 2026
5c099bf
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 8, 2026
b163269
refactor(attention): drop the FROST layout checks cuDNN already makes
nvegesna-netizen Oct 8, 2026
7364535
fix(attention): restore the FROST rank check, which cuDNN cannot make
nvegesna-netizen Oct 8, 2026
8a680d5
docs(attention): cite the engine capability for the FROST symmetry rule
nvegesna-netizen Oct 8, 2026
80672e3
fix(attention): check the whole softmax_lse shape, not just its leadi…
nvegesna-netizen Oct 8, 2026
0effc6c
docs(attention): shorten the cudnn_pygraph module docstring
nvegesna-netizen Oct 8, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 12 additions & 4 deletions docs/envvars.rst
Original file line number Diff line number Diff line change
Expand Up @@ -178,10 +178,12 @@ Then it applies a performance-based preference order among the remaining eligibl
In PyTorch, the broad preference order is ``FlashAttention > FusedAttention >
UnfusedDotProductAttention`` on supported pre-Hopper GPUs such as Ampere/Ada, and
``FusedAttention > FlashAttention > UnfusedDotProductAttention`` on Hopper and newer GPUs,
including Blackwell. In JAX, Transformer Engine uses cuDNN fused attention when
``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise it falls back to the
JAX-native implementation. See :doc:`examples/attention/attention` for a longer
backend-selection overview.
including Blackwell. On Blackwell SM100/SM103, FusedAttention has an extra sub-backend, FROST,
which is selected only for symmetric ``head_dim`` in (256, 512] and only when the cuDNN
sub-backends decline; it does not change the order above. In JAX, Transformer Engine uses cuDNN
fused attention when ``NVTE_FUSED_ATTN=1`` and an eligible cuDNN kernel is available; otherwise
it falls back to the JAX-native implementation. See :doc:`examples/attention/attention` for a
longer backend-selection overview.

.. envvar:: NVTE_FLASH_ATTN

Expand Down Expand Up @@ -213,6 +215,12 @@ backend-selection overview.
:Default: ``1``
:Description: Enable or disable FusedAttention backend (cuDNN-based) for DotProductAttention. When set to ``0``, FusedAttention will not be used.

.. envvar:: NVTE_FROST_ATTN

:Type: ``int`` (0 or 1)
:Default: ``1``
:Description: Enable or disable the FROST sub-backend of FusedAttention for DotProductAttention. FROST wraps the cuDNN FROST CuTe-DSL SDPA kernels through the cuDNN Frontend python API, rather than the C++ fused-attention path the other sub-backends use. **It is experimental and subject to change**, as the underlying cuDNN FROST engines are. When set to ``0``, FROST will not be used. It is selected only where the cuDNN sub-backends decline and is the only released backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism. It is limited to SM100/SM103 with BF16/FP16 inputs, a ``head_dim`` that is a multiple of 8, and ``nvidia-cudnn-frontend>=1.29.0`` and ``nvidia-cutlass-dsl>=4.7.0`` installed. It declines FP8, ``thd`` layouts, dropout, attention bias, KV caching, ``max_logit``, CUDA graph capture, and deterministic execution, the last because cuDNN offers no deterministic backward for these kernels.

.. envvar:: NVTE_UNFUSED_ATTN

:Type: ``int`` (0 or 1)
Expand Down
1 change: 1 addition & 0 deletions qa/L0_pytorch_unittest/test.sh
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hybrid_quantizat
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_identity_quantizer.xml $TE_PATH/tests/pytorch/test_identity_quantizer.py || test_fail "test_identity_quantizer.py"
NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_attention.xml $TE_PATH/tests/pytorch/attention/test_attention.py || test_fail "test_attention.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_flex_attention.xml $TE_PATH/tests/pytorch/attention/test_flex_attention.py || test_fail "test_flex_attention.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_frost_attention.xml $TE_PATH/tests/pytorch/attention/test_frost_attention.py || test_fail "test_frost_attention.py"
NVTE_GDN_TEST_REQUIRED=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gdn_attention.xml $TE_PATH/tests/pytorch/attention/test_gdn_attention.py || test_fail "test_gdn_attention.py"
NVTE_GDN2_TEST_REQUIRED=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gdn2_attention.xml $TE_PATH/tests/pytorch/attention/test_gdn2_attention.py || test_fail "test_gdn2_attention.py"
NVTE_GDP_TEST_REQUIRED=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gdp_attention.xml $TE_PATH/tests/pytorch/attention/test_gdp_attention.py || test_fail "test_gdp_attention.py"
Expand Down
26 changes: 26 additions & 0 deletions tests/pytorch/attention/run_attention_with_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from transformer_engine.pytorch import DType
from test_attention_with_cp import (
model_configs_flash_attn,
model_configs_frost_attn,
model_configs_fused_attn,
)
from transformer_engine.pytorch import (
Expand Down Expand Up @@ -276,6 +277,15 @@ def run_dpa_with_cp(
config = copy.deepcopy(model_configs_fused_attn[model])
else:
assert False, f"{model=} is not a known FusedAttention CP config!"
if kernel_backend == "FrostAttention":
# FROST is a sub-backend of FusedAttention, so NVTE_FUSED_ATTN has to stay on. Flash is
# left off; nothing else serves head_dim > 256, so the selector reaches FROST on its own.
os.environ["NVTE_FUSED_ATTN"] = "1"
os.environ["NVTE_FROST_ATTN"] = "1"
if model in model_configs_frost_attn:
config = copy.deepcopy(model_configs_frost_attn[model])
else:
assert False, f"{model=} is not a known FrostAttention CP config!"
assert config.attn_mask_type in [
"causal",
"no_mask",
Expand Down Expand Up @@ -596,6 +606,22 @@ def run_dpa_with_cp(
pad_between_seqs=pad_between_seqs,
fp8_output=fp8_mha,
)
if kernel_backend == "FrostAttention":
# Assert the sub-backend actually used, not just the one requested. FROST is
# currently the only selectable backend for these configs -- flash is env-gated off
# and CP disables unfused -- so a silent substitution is impossible today and this
# would pass by construction. It is here so it stops passing if that stops being
# true, rather than quietly testing some other kernel.
from transformer_engine.pytorch.attention.dot_product_attention.dot_product_attention import ( # pylint: disable=import-outside-toplevel
_attention_backends,
)

# pylint: disable-next=import-outside-toplevel
from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend

assert (
_attention_backends["fused_attention_backend"] == FusedAttnBackend.FROST
), "expected the FROST sub-backend to be selected, got %s" % (_attention_backends,)
if config.return_max_logit:
out_, max_logit_ = out_
if is_training:
Expand Down
122 changes: 122 additions & 0 deletions tests/pytorch/attention/test_attention_with_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -459,6 +459,16 @@ def test_cp_with_flash_attention_softcap(cp_pool, cp_comm_type):
)


# cuDNN FROST: head_dim in (256, 512] on SM100/SM103, the range no other backend
# serves together with context parallelism. Shapes are Gemma-4 global layers, which is what
# motivated the backend. seqlen must stay divisible by cp_size * 2 for causal load balancing.
model_configs_frost_attn = {
# test: ModelConfig(b, sq, hq, dqk)
"cp_hd512_0": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="causal"),
"cp_hd512_1": ModelConfig(2, 4096, 8, 512, num_gqa_groups=4, attn_mask_type="no_mask"),
"cp_hd512_2": ModelConfig(2, 2048, 8, 512, num_gqa_groups=8, attn_mask_type="causal"),
}
Comment thread
greptile-apps[bot] marked this conversation as resolved.

model_configs_fused_attn = {
# test: ModelConfig(b, sq, hq, dqk)
"cp_1_0": ModelConfig(2, 4096, 12, 128, attn_mask_type="causal", return_max_logit=True), # MHA
Expand Down Expand Up @@ -747,6 +757,118 @@ def test_cp_with_fused_attention(
)


def _frost_availability():
"""Why FrostAttention cannot run here, or None if it can.

The backend needs cuDNN Frontend >= 1.29.0 and, less obviously,
nvidia-cutlass-dsl >= 4.7.0: cudnn-frontend only declares >= 4.6.2, and below the FROST floor
every FROST engine silently declines and ordinary backend plans are returned with no error.
Reporting the reason as a skip keeps that distinguishable from a real failure.
"""
if get_device_compute_capability() not in ((10, 0), (10, 3)):
return "FrostAttention requires SM100/SM103 (the cuDNN d512 backward is Blackwell-only)."
from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import (
is_frost_attention_available,
)

ok, reason = is_frost_attention_available()
return None if ok else reason


@pytest.mark.parametrize("model", model_configs_frost_attn.keys())
@pytest.mark.parametrize("qkv_format", ["bshd", "sbhd"])
@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"])
def test_cp_with_frost_attention(cp_pool, model, qkv_format, cp_comm_type):
"""Context parallelism at head_dim 512, which no other backend serves.

thd is excluded because the backend declines it: it needs varlen support that is not
implemented. a2a+p2p is covered separately below, since it needs four ranks.
"""
reason = _frost_availability()
if reason is not None:
pytest.skip(reason)

config = model_configs_frost_attn[model]
config.context_parallel = True
config.cp_comm_type = cp_comm_type

_submit(
cp_pool(2),
dtype="bf16",
model=model,
qkv_format=qkv_format,
kernel_backend="FrostAttention",
cp_comm_type=cp_comm_type,
is_training=True,
log_level=pytest_logging_level,
)


def test_cp_with_frost_attention_a2a_p2p(cp_pool):
"""One case, because a2a+p2p composes rather than adding a path.

It dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p, with an a2a communication stage
on either side of the ring, and it needs four ranks rather than two. One case says the
composition works; six would pay four-rank prices to re-cover p2p.
"""
reason = _frost_availability()
if reason is not None:
pytest.skip(reason)

# The shortest model, since this is the only four-rank case and so the only one that would
# take all of a2a+p2p down with it if it timed out.
config = model_configs_frost_attn["cp_hd512_2"]
config.context_parallel = True
config.cp_comm_type = "a2a+p2p"
# The a2a stage shards heads across its subgroup, so both counts must divide by it.
assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, (
f"a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups"
f" ({config.num_gqa_groups}) divisible by the a2a subgroup size"
)

_submit(
cp_pool(4),
dtype="bf16",
model="cp_hd512_2",
qkv_format="bshd",
kernel_backend="FrostAttention",
cp_comm_type="a2a+p2p",
is_training=True,
log_level=pytest_logging_level,
)


@pytest.mark.parametrize("cp_comm_type", ["p2p", "all_gather", "a2a"])
def test_cp_with_frost_attention_fp16(cp_pool, cp_comm_type):
"""One fp16 arm per comm type, since the matrix above is bf16 throughout.

The backend serves BF16 and FP16, but every context-parallel configuration was covered in bf16
only. fp16 has a far narrower exponent range, and the ring correction exponentiates a difference
of log-sum-exp values across steps, so a range problem would surface here rather than in the
non-CP numerics. One model and one layout keeps the cost to three cases rather than doubling
the matrix; a2a+p2p is omitted because it would need a second four-rank pool for a dtype that
exercises no additional code path.
"""
reason = _frost_availability()
if reason is not None:
pytest.skip(reason)

config = model_configs_frost_attn["cp_hd512_0"]
config.context_parallel = True
config.cp_comm_type = cp_comm_type

_submit(
cp_pool(2),
dtype="fp16",
model="cp_hd512_0",
qkv_format="bshd",
kernel_backend="FrostAttention",
cp_comm_type=cp_comm_type,
is_training=True,
log_level=pytest_logging_level,
)


@pytest.mark.skipif(
get_device_compute_capability() < (9, 0), reason="FusedAttention THD requires sm90+."
)
Expand Down
138 changes: 138 additions & 0 deletions tests/pytorch/attention/test_flex_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from transformer_engine.pytorch import DotProductAttention, is_bf16_available
from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends
import transformer_engine.pytorch.attention.dot_product_attention.flex_attention as flex_attention
from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph
from transformer_engine.pytorch.utils import get_device_compute_capability

_current_file = pathlib.Path(__file__).resolve()
Expand All @@ -22,6 +23,23 @@
get_available_attention_backends,
)


def _flex_availability():
"""Why the cuDNN score_mod path cannot run here, or None if it can."""
if not torch.cuda.is_available():
return "no CUDA device"
try:
flex_attention._import_cudnn_frontend()
except ImportError as exc:
return "nvidia-cudnn-frontend not importable: %s" % exc
return None


_SKIP = _flex_availability()
# Mirrors NVTE_GDN_TEST_REQUIRED: the tests that reach cuDNN skip silently without the frontend.
if os.getenv("NVTE_FLEX_TEST_REQUIRED", "0") == "1" and _SKIP is not None:
raise RuntimeError("NVTE_FLEX_TEST_REQUIRED=1, but flex attention is unavailable: %s" % _SKIP)

param_types = [torch.float16]
if torch.cuda.is_available() and is_bf16_available():
param_types.append(torch.bfloat16)
Expand Down Expand Up @@ -705,3 +723,123 @@ def test_dot_product_attention_score_mod(dtype, qkv_format, score_mod_case, scal
torch.testing.assert_close(q.grad, q_ref.grad, **tols)
torch.testing.assert_close(k.grad, k_ref.grad, **tols)
torch.testing.assert_close(v.grad, v_ref.grad, **tols)


def test_flex_bars_the_frost_engines():
"""flex must tell cuDNN not to use a FROST engine, not merely decline to ask for them.

The switch that offers those engines is process-wide, so another caller in the process, or a
user setting CUDNN_FRONTEND_ENABLE_FROST_ENGINES, puts them ahead of the backend engines for
these graphs too. They accept a score_mod graph, pass check_support, build, and then compute
without the callback.

The bar is the shared helper's default rather than something flex asks for, so this checks
flex's own path reaches it: a caller that forgot to opt in is the failure being prevented.

No GPU: this checks the instruction is passed, not what cuDNN does with it.
"""
barred = []

class FakeGraph:
"""Records the engines flex bars, and stops at the first call it cannot serve."""

def validate(self):
pass

def build_operation_graph(self):
pass

def create_execution_plans(self, _heuristics):
pass

def deselect_engines(self, names):
barred.extend(names)

def check_support(self):
pass

def build_plans(self, _policy):
pass

def get_workspace_size(self):
return 4096

try:
flex_attention._import_cudnn_frontend()
except ImportError:
pytest.skip("cuDNN frontend Python package is required for score_mod attention.")

assert flex_attention._finalize_cudnn_graph(FakeGraph()) == 4096
assert barred, "flex did not ask cuDNN to exclude any engine"
assert set(cudnn_pygraph.FROST_PLAN_TOKENS) <= set(barred), barred


@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required.")
def test_frost_switch_does_not_change_what_flex_computes():
"""Enabling the FROST engines must not change flex's output.

This is the property the silent drop violates: with the engines on, an unpinned build selects
a FROST plan at every head dim measured on B200, and that plan returns plain attention with
the score_mod discarded. Comparing flex against itself across the switch needs no reference
and no knowledge of which plan ran; if the two differ, a different kernel answered.
"""
try:
flex_attention._import_cudnn_frontend()
except ImportError:
pytest.skip("cuDNN frontend Python package is required for score_mod attention.")
# Without this the test is vacuous: if the engines are absent, decline on arch, or sit below
# their version floors, both runs get a backend plan and agree no matter what flex does. The
# availability probe covers all three, where an arch check alone would miss the version ones.
from transformer_engine.pytorch.attention.dot_product_attention.frost_attention import (
is_frost_attention_available,
)

frost_ok, frost_reason = is_frost_attention_available()
if not frost_ok:
pytest.skip("the FROST engines must be reachable to test anything: %s" % frost_reason)

env = "CUDNN_FRONTEND_ENABLE_FROST_ENGINES"
saved = os.environ.get(env)
torch.manual_seed(0)
b, h, s, d = 2, 4, 512, 64
dtype = torch.bfloat16 if is_bf16_available() else torch.float16
q, k, v = (torch.randn(b, s, h, d, device="cuda", dtype=dtype) for _ in range(3))

def bias_score_mod(score_mod_graph, score_tensor, _tensors):
"""score += (row - col). Self-contained, and large enough that dropping it is obvious."""
cudnn = flex_attention._import_cudnn_frontend()
row = score_mod_graph.gen_index(input=score_tensor, axis=2)
row.set_data_type(cudnn.data_type.INT32)
col = score_mod_graph.gen_index(input=score_tensor, axis=3)
col.set_data_type(cudnn.data_type.INT32)
bias = score_mod_graph.sub(a=row, b=col, compute_data_type=cudnn.data_type.FLOAT)
bias.set_data_type(cudnn.data_type.FLOAT)
return score_mod_graph.add(a=score_tensor, b=bias, compute_data_type=cudnn.data_type.FLOAT)

def run():
flex_attention._cudnn_score_mod_graph_cache.clear()
return flex_attention.FusedAttentionWithScoreModFunc.apply(
False, q, k, v, "bshd", "bshd", d**-0.5, bias_score_mod, None, None, None, False
)

try:
os.environ.pop(env, None)
without = run()
os.environ[env] = "1"
with_engines = run()
finally:
flex_attention._cudnn_score_mod_graph_cache.clear()
if saved is None:
os.environ.pop(env, None)
else:
os.environ[env] = saved

torch.testing.assert_close(
with_engines,
without,
msg=lambda m: (
"flex computed something different with the FROST engines enabled, which means a"
" FROST plan answered and dropped the score_mod:\n"
+ m
),
)
Loading
Loading