Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
51 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
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
17 changes: 13 additions & 4 deletions docs/envvars.rst
Original file line number Diff line number Diff line change
Expand Up @@ -178,10 +178,13 @@ 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 the order is ``FusedAttention > FlashAttention >
FrostAttention > UnfusedDotProductAttention``; FrostAttention only becomes eligible for
symmetric ``head_dim`` in (256, 512], which flash and fused attention do not serve, so the
backend it can displace is UnfusedDotProductAttention. 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 +216,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 FrostAttention backend (the cuDNN FROST CuTe-DSL SDPA kernels in cuDNN Frontend) for DotProductAttention. **This backend is experimental and subject to change**, including the possibility of being folded into FusedAttention; the underlying cuDNN FROST engines are themselves experimental. When set to ``0``, FrostAttention will not be used. From released components it is the only backend serving symmetric ``head_dim`` in (256, 512] together with context parallelism; without context parallelism UnfusedDotProductAttention also covers that range, and FrostAttention is preferred over it where both are eligible. 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 supports context parallelism with ``cp_comm_type`` of ``p2p``, ``all_gather``, ``a2a`` or ``a2a+p2p``, and sliding-window attention with ``all_gather`` or ``a2a`` (declined with ``p2p`` and ``a2a+p2p``, whose ring shards KV across steps). It declines FP8, ``thd`` layouts, dropout, attention bias, softcap, KV caching, ``max_logit``, 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
22 changes: 22 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,14 @@ 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":
# Leave NVTE_FLASH_ATTN and NVTE_FUSED_ATTN at 0: FROST is the only backend that serves
# head_dim > 256, so get_attention_backend selects it on its own.
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 +605,19 @@ def run_dpa_with_cp(
pad_between_seqs=pad_between_seqs,
fp8_output=fp8_mha,
)
if kernel_backend == "FrostAttention":
# Assert the backend actually used, not just the one requested. FROST is currently
# the only selectable backend for these configs -- flash and fused are 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,
)

assert _attention_backends[
"use_frost_attention"
], "expected FrostAttention to be selected, got %s" % (_attention_backends,)
if config.return_max_logit:
out_, max_logit_ = out_
if is_training:
Expand Down
103 changes: 103 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: symmetric 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,99 @@ 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", "a2a+p2p"])
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 needs four ranks rather than two -- an a2a subgroup crossed with a p2p subgroup -- and
exercises no new attention code: it dispatches to the same AttnFuncWithCPAndKVP2P as plain p2p,
with an a2a communication stage on either side of the ring. It is covered here so that claim is
measured rather than assumed.
"""
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

# a2a requires num_heads and num_gqa_groups divisible by the a2a subgroup size; every config
# here satisfies that, but assert rather than rely on it staying true.
if cp_comm_type == "a2a+p2p":
assert config.num_heads % 2 == 0 and config.num_gqa_groups % 2 == 0, (
f"cp_comm_type=a2a+p2p needs num_heads ({config.num_heads}) and num_gqa_groups"
f" ({config.num_gqa_groups}) divisible by the a2a subgroup size"
)

pool = cp_pool(4 if cp_comm_type == "a2a+p2p" else 2)

_submit(
pool,
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,
)


@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
192 changes: 192 additions & 0 deletions tests/pytorch/attention/test_flex_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -705,3 +705,195 @@ 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)


@pytest.mark.parametrize("direction", ["fwd", "bwd"])
def test_score_mod_graph_signatures_stay_aligned(direction):
"""The cache key, the builder and the getter are splatted from one positional tuple.

`_get_cudnn_score_mod_*_graph` passes the same `build_args` tuple to the cache key and to the
builder, so the three parameter lists have to stay in the same order. Nothing enforced that,
and the failure is quiet in the worst direction: a parameter inserted in one signature and not
another shifts the rest by one, and a shifted *cache key* is not a crash, it is two different
configurations sharing a cached graph.

No GPU: this reads signatures only.
"""
import inspect

names = [
getattr(flex_attention, "_cudnn_score_mod_%s_cache_key" % direction),
getattr(flex_attention, "_build_cudnn_score_mod_%s_graph" % direction),
getattr(flex_attention, "_get_cudnn_score_mod_%s_graph" % direction),
]
signatures = [list(inspect.signature(fn).parameters) for fn in names]
reference = signatures[0]
for fn, params in zip(names[1:], signatures[1:]):
assert params == reference, (
"%s takes %s but _cudnn_score_mod_%s_cache_key takes %s; these are splatted from one"
" positional tuple and must stay in the same order"
% (fn.__name__, params, direction, reference)
)


@pytest.mark.parametrize(
"mask_spec,expected",
[
# Causal: top-left aligned, right bound pinned to the diagonal, no left bound.
(("causal", (-1, 0)), {"diagonal_alignment": "TOP_LEFT", "diagonal_band_right_bound": 0}),
# Bottom-right causal, which is what KV trimming produces whenever SKV > SQ.
(
("causal_bottom_right", (-1, 0)),
{"diagonal_alignment": "BOTTOM_RIGHT", "diagonal_band_right_bound": 0},
),
# Sliding window. cuDNN's left bound counts the diagonal itself and TE's window_size does
# not, so 511 must arrive as 512. Getting this wrong drops one token of context per layer
# and no shape-level test would notice.
(
("causal", (511, 0)),
{
"diagonal_alignment": "TOP_LEFT",
"diagonal_band_right_bound": 0,
"diagonal_band_left_bound": 512,
},
),
# No mask at all: no alignment, no bounds.
(("no_mask", (-1, -1)), {}),
],
)
def test_mask_spec_translates_to_a_diagonal_band(mask_spec, expected):
"""A mask_spec must become the cuDNN band kwargs, and never a score_mod.

No GPU: this builds no graph, it checks the kwargs the graph would be given. The frontend is
an optional dependency, so skip rather than fail where it is absent.
"""
try:
cudnn = flex_attention._import_cudnn_frontend()
except ImportError:
pytest.skip("cuDNN frontend Python package is required for the diagonal-band kwargs.")
got = flex_attention._mask_or_score_mod_kwargs(mask_spec, None)

assert "score_mod" not in got and "use_causal_mask" not in got
for key, want in expected.items():
if key == "diagonal_alignment":
assert got[key] == getattr(cudnn.diagonal_alignment, want)
else:
assert got[key] == want
assert set(got) == set(expected)


def test_mask_spec_and_score_mod_cannot_be_combined():
"""cuDNN's backward refuses the pair, so flex must refuse it before building the graph."""
with pytest.raises(ValueError, match="cannot be combined"):
flex_attention._mask_or_score_mod_kwargs(("causal", (-1, 0)), lambda *a, **k: None)


def test_no_mask_spec_still_takes_the_score_mod_path():
"""The default path must be byte-identical to what it was before mask_spec existed."""
sentinel = object()
assert flex_attention._mask_or_score_mod_kwargs(None, sentinel) == {
"use_causal_mask": False,
"score_mod": sentinel,
}


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 a FrostAttention call elsewhere 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.

No GPU: this checks the instruction is passed, not what cuDNN does with it.
"""
from transformer_engine.pytorch.attention.dot_product_attention import cudnn_pygraph

seen = {}

def fake_finalize(graph, **kwargs):
seen.update(kwargs)
return 4096, None

original = cudnn_pygraph.finalize_plans
cudnn_pygraph.finalize_plans = fake_finalize
try:
assert flex_attention._finalize_cudnn_graph(object()) == 4096
finally:
cudnn_pygraph.finalize_plans = original

excluded = seen.get("exclude_plan_tokens")
assert excluded, "flex did not ask cuDNN to exclude any engine"
assert "sdpa_fwd_prefill_sm100" in excluded and "sdpa_bwd_sm100" in excluded, excluded


@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 violated: with the engines on, an unpinned build selected
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 nearly everywhere: where the FROST engines are absent or
# decline on arch, both runs get a backend plan and agree no matter what flex does. The
# engines themselves are found lazily at planning time, so the switch works whenever it is
# set, but they still have to exist.
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 for this 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()
Comment thread
greptile-apps[bot] marked this conversation as resolved.
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