From 99d22075ecdf16f84a61fc99c9f08030c9cde378 Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Wed, 19 Aug 2026 18:06:09 -0400 Subject: [PATCH 1/7] Update to new QoLA/CK-JIT/AITER --- 3rdparty/QoLA | 2 +- 3rdparty/ck_jit | 2 +- transformer_engine/common/ck_fused_attn/CMakeLists.txt | 2 +- transformer_engine/common/ck_fused_attn/qola_manifest.toml | 5 +++-- .../common/ck_fused_attn/src/ck_fused_attn_fwd.cpp | 4 ++-- 5 files changed, 8 insertions(+), 7 deletions(-) diff --git a/3rdparty/QoLA b/3rdparty/QoLA index 239be99d59..a1f9245e42 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit 239be99d5906a0c7f7a202b293abb72dace5f283 +Subproject commit a1f9245e4263a172d56e1ec490b3729c0f790b6c diff --git a/3rdparty/ck_jit b/3rdparty/ck_jit index 882083cb84..3cf034d344 160000 --- a/3rdparty/ck_jit +++ b/3rdparty/ck_jit @@ -1 +1 @@ -Subproject commit 882083cb84e7af35cceec403c05d3d1594eaaf1f +Subproject commit 3cf034d3445d7faf5de852840c4419ab720483cf diff --git a/transformer_engine/common/ck_fused_attn/CMakeLists.txt b/transformer_engine/common/ck_fused_attn/CMakeLists.txt index 964ff6513c..2e4c02658d 100644 --- a/transformer_engine/common/ck_fused_attn/CMakeLists.txt +++ b/transformer_engine/common/ck_fused_attn/CMakeLists.txt @@ -222,7 +222,7 @@ add_library(ck_fused_attn SHARED ${ck_fused_attn_SOURCES}) set(CK_FUSED_ATTN_COMPILE_OPTIONS) list(APPEND CK_FUSED_ATTN_COMPILE_OPTIONS -DCK_TILE_FLOAT_TO_BFLOAT16_DEFAULT=${CK_FUSED_ATTN_FLOAT_TO_BFLOAT16_DEFAULT} - -DENABLE_CK=1 -DFAV_NATIVE_ON=1) + -DENABLE_CK=1 -DFA_WITH_NATIVE_SPLITKV=1) # Public QoLA headers ship alongside the .so libs in ${__AITER_MHA_PATH}/../include # (emitted by qola.cli build, or copied from the QoLA build dir above for the diff --git a/transformer_engine/common/ck_fused_attn/qola_manifest.toml b/transformer_engine/common/ck_fused_attn/qola_manifest.toml index f2dc42a1d0..a689234445 100644 --- a/transformer_engine/common/ck_fused_attn/qola_manifest.toml +++ b/transformer_engine/common/ck_fused_attn/qola_manifest.toml @@ -1,7 +1,7 @@ [qola] -aiter_commit = "7a24cd87525834fa9aeaa021e6ee80ab7de028f9" # pinned AITER submodule commit +aiter_commit = "a50a62aaa292427e00190dd7ba58e39a6db4e61e" # pinned AITER submodule commit namespace = "te" -rocm_versions = ["7.2"] +rocm_versions = ["7.14"] [build] architectures = ["gfx950", "gfx942"] @@ -12,6 +12,7 @@ mode = "cpp_itfs" receipt = 700 drop_srcs = ["mha_fwd_split.cu", "mha_fwd_batch_prefill.cu"] drop_directions = ["fwd_splitkv", "batch_prefill"] +flags_extra_cc = ["'-DFA_WITH_NATIVE_SPLITKV=1'"] [[modules]] name = "libmha_bwd" diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp index e73619127b..e6700147bf 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp @@ -267,7 +267,7 @@ hipError_t ck_attn_fwd(const CKAttnFwdArgs& args, hipStream_t stream){ } int ck_attn_fwd_num_splits(const CKAttnFwdArgs& args){ -#if FAV_NATIVE_ON +#if FA_WITH_NATIVE_SPLITKV aiter::mha_fwd_args fmha_args = build_fwd_fmha_args(args); return QOLA_NS(mha_fwd_calculate_num_splits)(fmha_args); #else @@ -276,7 +276,7 @@ int ck_attn_fwd_num_splits(const CKAttnFwdArgs& args){ } size_t ck_attn_fwd_workspace_size(const CKAttnFwdArgs& args){ -#if FAV_NATIVE_ON +#if FA_WITH_NATIVE_SPLITKV aiter::mha_fwd_args fmha_args = build_fwd_fmha_args(args); return QOLA_NS(mha_fwd_workspace_size)(fmha_args); #else From 5cced7ae5b8f9a312008d962c84f10e53294d2c2 Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Sat, 22 Aug 2026 21:03:52 -0400 Subject: [PATCH 2/7] Update CI prebuild CK blob list --- ci/ck_jit_prebuild.txt | 78 ++++++++++++++++++++++++++++++++++++------ 1 file changed, 67 insertions(+), 11 deletions(-) diff --git a/ci/ck_jit_prebuild.txt b/ci/ck_jit_prebuild.txt index 5c06af1cb4..13e927fd51 100644 --- a/ci/ck_jit_prebuild.txt +++ b/ci/ck_jit_prebuild.txt @@ -2,10 +2,12 @@ fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_npad_deterministic_gfx9 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_npad_deterministic_gfx950 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_npad_ndeterministic_gfx9 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_npad_ndeterministic_gfx950 +fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_pd_deterministic_gfx950 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_pd_ndeterministic_gfx9 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_pd_ndeterministic_gfx950 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_ps_ndeterministic_gfx9 fmha_bwd_convert_dq_d128_bf16_b64_batch_o2_ps_ndeterministic_gfx950 +fmha_bwd_convert_dq_d128_bf16_b64_group_o2_psd_deterministic_gfx950 fmha_bwd_convert_dq_d128_bf16_b64_group_o2_ps_deterministic_gfx9 fmha_bwd_convert_dq_d128_bf16_b64_group_o2_ps_deterministic_gfx950 fmha_bwd_convert_dq_d128_bf16_b64_group_o2_ps_ndeterministic_gfx9 @@ -29,6 +31,8 @@ fmha_bwd_convert_dq_d256_bf16_b64_batch_o2_psd_ndeterministic_gfx9 fmha_bwd_convert_dq_d256_bf16_b64_batch_o2_psd_ndeterministic_gfx950 fmha_bwd_convert_dq_d256_bf16_b64_batch_o2_ps_ndeterministic_gfx9 fmha_bwd_convert_dq_d256_bf16_b64_batch_o2_ps_ndeterministic_gfx950 +fmha_bwd_convert_dq_d256_bf16_b64_group_o2_psd_ndeterministic_gfx9 +fmha_bwd_convert_dq_d256_bf16_b64_group_o2_psd_ndeterministic_gfx950 fmha_bwd_convert_dq_d256_bf16_b64_group_o2_ps_ndeterministic_gfx9 fmha_bwd_convert_dq_d256_bf16_b64_group_o2_ps_ndeterministic_gfx950 fmha_bwd_convert_dq_d256_fp16_b64_batch_o2_pd_ndeterministic_gfx9 @@ -49,21 +53,33 @@ fmha_bwd_convert_dq_d64_bf16_b64_batch_o2_npad_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_bf16_b64_batch_o2_pd_deterministic_gfx950 fmha_bwd_convert_dq_d64_bf16_b64_batch_o2_pd_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_bf16_b64_batch_o2_psd_ndeterministic_gfx950 +fmha_bwd_convert_dq_d64_bf16_b64_batch_o2_ps_ndeterministic_gfx9 +fmha_bwd_convert_dq_d64_bf16_b64_batch_o2_ps_ndeterministic_gfx950 +fmha_bwd_convert_dq_d64_bf16_b64_group_o2_psd_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_bf16_b64_group_o2_ps_ndeterministic_gfx9 fmha_bwd_convert_dq_d64_bf16_b64_group_o2_ps_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_fp16_b64_batch_o2_npad_ndeterministic_gfx9 fmha_bwd_convert_dq_d64_fp16_b64_batch_o2_npad_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_fp16_b64_batch_o2_pd_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_fp16_b64_batch_o2_psd_ndeterministic_gfx950 +fmha_bwd_convert_dq_d64_fp16_b64_batch_o2_ps_ndeterministic_gfx9 +fmha_bwd_convert_dq_d64_fp16_b64_batch_o2_ps_ndeterministic_gfx950 fmha_bwd_convert_dq_d64_fp16_b64_group_o2_ps_ndeterministic_gfx9 fmha_bwd_convert_dq_d64_fp16_b64_group_o2_ps_ndeterministic_gfx950 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_alibi_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_alibi_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_mask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_mask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_nmask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_nmask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_mask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_nmask_dropout_wg16_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 @@ -81,11 +97,18 @@ fmha_bwd_d128_bf16_batch_b16x16x128x16x128x16x16x128x128_r1x1x1_r1x1x1_r1x1x1_w1 fmha_bwd_d128_bf16_batch_b16x16x128x16x128x16x16x128x128_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq16_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_alibi_ndbias_mask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_alibi_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_mask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_mask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_nmask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_dbias_nmask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_mask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_mask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_nmask_dropout_wg16_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_bias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_bias_dbias_mask_ndropout_deterministic_trload_gfx950 @@ -102,15 +125,19 @@ fmha_bwd_d128_bf16_batch_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w fmha_bwd_d128_bf16_batch_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_batch_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pdv8_nbias_ndbias_mask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pdv8_nbias_ndbias_nmask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_bf16_group_b16x16x128x16x128x16x16x128x128_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq16_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_trload_gfx950 +fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_group_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_mask_ndropout_deterministic_trload_gfx950 @@ -119,6 +146,7 @@ fmha_bwd_d128_bf16_group_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w fmha_bwd_d128_bf16_group_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_bf16_group_b32x128x128x32x128x32x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_fp16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_deterministic_ntrload_gfx9 +fmha_bwd_d128_fp16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d128_fp16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_ntrload_gfx9 fmha_bwd_d128_fp16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d128_fp16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 @@ -137,6 +165,7 @@ fmha_bwd_d128_fp16_batch_b16x128x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w fmha_bwd_d128_fp16_batch_b16x16x128x16x128x16x16x128x128_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq16_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_fp16_batch_b16x16x128x16x128x16x16x128x128_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq16_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_fp16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_deterministic_trload_gfx950 +fmha_bwd_d128_fp16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_fp16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_deterministic_trload_gfx950 fmha_bwd_d128_fp16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d128_fp16_batch_b16x192x128x16x128x16x32x128x128_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_bias_dbias_mask_dropout_wg16_ndeterministic_trload_gfx950 @@ -185,6 +214,10 @@ fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w1 fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx950 fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx950 +fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx950 +fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 +fmha_bwd_d256_bf16_group_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx950 fmha_bwd_d256_fp16_batch_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d256_fp16_batch_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx950 fmha_bwd_d256_fp16_batch_b16x64x256x16x256x16x32x256x256_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pd8dv8_nbias_ndbias_mask_ndropout_ndeterministic_ntrload_gfx9 @@ -248,6 +281,7 @@ fmha_bwd_d64_bf16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16 fmha_bwd_d64_bf16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pdv8_nbias_ndbias_nmask_dropout_wg16_ndeterministic_ntrload_gfx950 fmha_bwd_d64_bf16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pdv8_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d64_bf16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_pdv8_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx950 +fmha_bwd_d64_bf16_batch_b32x16x64x32x64x32x16x64x64_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq32_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d64_bf16_batch_b32x16x64x32x64x32x16x64x64_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq32_pd8dv8_nbias_ndbias_nmask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d64_bf16_batch_b32x256x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_alibi_ndbias_mask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d64_bf16_batch_b32x256x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_bias_dbias_mask_dropout_wg16_ndeterministic_trload_gfx950 @@ -316,6 +350,7 @@ fmha_bwd_d64_fp16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16 fmha_bwd_d64_fp16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_dropout_wg16_ndeterministic_ntrload_gfx9 fmha_bwd_d64_fp16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx9 fmha_bwd_d64_fp16_batch_b32x128x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x16_o1_maxq0_npad_nbias_ndbias_nmask_ndropout_ndeterministic_ntrload_gfx950 +fmha_bwd_d64_fp16_batch_b32x16x64x32x64x32x16x64x64_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq32_npad_nbias_ndbias_nmask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d64_fp16_batch_b32x16x64x32x64x32x16x64x64_r1x1x1_r1x1x1_r1x1x1_w16x16x32_w16x16x16_o2_maxq32_pd8dv8_nbias_ndbias_nmask_dropout_wg16_ndeterministic_trload_gfx950 fmha_bwd_d64_fp16_batch_b32x256x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_alibi_ndbias_mask_ndropout_ndeterministic_trload_gfx950 fmha_bwd_d64_fp16_batch_b32x256x64x32x64x32x32x64x64_r1x4x1_r4x1x1_r1x4x1_w16x16x32_w16x16x32_o1_maxq0_npad_bias_dbias_mask_dropout_wg16_ndeterministic_trload_gfx950 @@ -374,6 +409,8 @@ fmha_bwd_dot_do_o_d256_bf16_b64_batch_o2_psdv_gfx9 fmha_bwd_dot_do_o_d256_bf16_b64_batch_o2_psdv_gfx950 fmha_bwd_dot_do_o_d256_bf16_b64_batch_o2_ps_gfx9 fmha_bwd_dot_do_o_d256_bf16_b64_batch_o2_ps_gfx950 +fmha_bwd_dot_do_o_d256_bf16_b64_group_o2_psdv_gfx9 +fmha_bwd_dot_do_o_d256_bf16_b64_group_o2_psdv_gfx950 fmha_bwd_dot_do_o_d256_bf16_b64_group_o2_ps_gfx9 fmha_bwd_dot_do_o_d256_bf16_b64_group_o2_ps_gfx950 fmha_bwd_dot_do_o_d256_fp16_b64_batch_o2_pdv_gfx9 @@ -394,6 +431,8 @@ fmha_bwd_dot_do_o_d64_bf16_b64_batch_o2_npad_gfx950 fmha_bwd_dot_do_o_d64_bf16_b64_batch_o2_pdv_gfx9 fmha_bwd_dot_do_o_d64_bf16_b64_batch_o2_pdv_gfx950 fmha_bwd_dot_do_o_d64_bf16_b64_batch_o2_psdv_gfx950 +fmha_bwd_dot_do_o_d64_bf16_b64_batch_o2_ps_gfx9 +fmha_bwd_dot_do_o_d64_bf16_b64_batch_o2_ps_gfx950 fmha_bwd_dot_do_o_d64_bf16_b64_group_o2_psdv_gfx9 fmha_bwd_dot_do_o_d64_bf16_b64_group_o2_psdv_gfx950 fmha_bwd_dot_do_o_d64_bf16_b64_group_o2_ps_gfx9 @@ -402,16 +441,26 @@ fmha_bwd_dot_do_o_d64_fp16_b64_batch_o2_npad_gfx9 fmha_bwd_dot_do_o_d64_fp16_b64_batch_o2_npad_gfx950 fmha_bwd_dot_do_o_d64_fp16_b64_batch_o2_pdv_gfx950 fmha_bwd_dot_do_o_d64_fp16_b64_batch_o2_psdv_gfx950 +fmha_bwd_dot_do_o_d64_fp16_b64_batch_o2_ps_gfx9 +fmha_bwd_dot_do_o_d64_fp16_b64_batch_o2_ps_gfx950 fmha_bwd_dot_do_o_d64_fp16_b64_group_o2_ps_gfx9 fmha_bwd_dot_do_o_d64_fp16_b64_group_o2_ps_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_alibi_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_alibi_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_alibi_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_alibi_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_bias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 @@ -427,8 +476,12 @@ fmha_fwd_d128_bf16_batch_b16x32x64x128x32x128_r1x1x1_r1x1x1_w16x16x32_w16x16x32_ fmha_fwd_d128_bf16_batch_b16x32x64x128x32x128_r1x1x1_r1x1x1_w16x16x32_w16x16x32_qr_async_trload_vr_pddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_trload_nsink_gfx950 fmha_fwd_d128_bf16_batch_b64x128x32x128x32x128_r4x1x1_r4x1x1_w16x16x32_w16x16x16_qr_async_vr_psddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_batch_b64x128x32x128x32x128_r4x1x1_r4x1x1_w16x16x32_w16x16x16_qr_async_vr_psddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 +fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d128_bf16_group_b128x128x32x128x32x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d128_bf16_group_b128x64x32x128x16x128_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_trload_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_trload_nsink_gfx950 @@ -493,27 +546,27 @@ fmha_fwd_d192_fp16_batch_b128x128x32x128x32x192_r4x1x1_r4x1x1_w32x32x16_w32x32x1 fmha_fwd_d192_fp16_batch_b128x128x32x192x32x192_r4x1x1_r4x1x1_w32x32x16_w32x32x16_o1_qr_async_vr_psddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d192_fp16_batch_b128x128x32x192x32x192_r4x1x1_r4x1x1_w32x32x16_w32x32x16_o1_qr_async_vr_psddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_group_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_group_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_bf16_group_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_bf16_group_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_group_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_bf16_group_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_fp16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_fp16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_fp16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_fp16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_fp16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_fp16_batch_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_fp16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_fp16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_fp16_batch_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d256_fp16_group_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 -fmha_fwd_d256_fp16_group_b128x128x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d256_fp16_group_b128x64x32x256x32x256_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_pssk_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d32_bf16_batch_b128x64x16x32x32x32_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d32_bf16_batch_b128x64x16x32x32x32_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d32_bf16_batch_b128x64x16x32x32x32_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psddv_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 @@ -561,6 +614,7 @@ fmha_fwd_d64_bf16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr fmha_fwd_d64_bf16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_bias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d64_bf16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_bias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d64_bf16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_psskddv_nlogits_bias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d64_bf16_batch_b16x32x64x64x32x64_r1x1x1_r1x1x1_w16x16x32_w16x16x32_qr_async_trload_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_trload_nsink_gfx950 fmha_fwd_d64_bf16_group_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d64_bf16_group_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d64_bf16_group_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 @@ -587,6 +641,8 @@ fmha_fwd_d64_fp16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr fmha_fwd_d64_fp16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d64_fp16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d64_fp16_batch_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_vr_npad_nlogits_bias_nmask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx950 +fmha_fwd_d64_fp16_batch_b16x32x64x64x32x64_r1x1x1_r1x1x1_w16x16x32_w16x16x32_qr_async_trload_vr_npad_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_trload_nsink_gfx950 +fmha_fwd_d64_fp16_batch_b16x32x64x64x32x64_r1x1x1_r1x1x1_w16x16x32_w16x16x32_qr_async_trload_vr_npad_nlogits_nbias_nmask_lse_ndropout_nskip_nqscale_trload_nsink_gfx950 fmha_fwd_d64_fp16_group_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx9 fmha_fwd_d64_fp16_group_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_dropout_nskip_nqscale_ntrload_nsink_gfx950 fmha_fwd_d64_fp16_group_b128x64x32x64x32x64_r4x1x1_r4x1x1_w32x32x16_w32x32x16_qr_async_vr_psskddv_nlogits_nbias_mask_lse_ndropout_nskip_nqscale_ntrload_nsink_gfx9 From 4f651ad04ab06a3cc995b3288309ccbb61b3275a Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Fri, 28 Aug 2026 13:36:39 -0400 Subject: [PATCH 3/7] Skip known failing test config --- tests/jax/test_fused_attn.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index bf5e31d39b..36dde0d492 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1897,6 +1897,11 @@ def test_backward( swa, seq_desc_format, ): + if (is_hip_extension() and qkv_layout == QKVLayout.THD_THD_THD and + attn_mask_type == AttnMaskType.PADDING_MASK and + attn_bias_type == AttnBiasType.NO_BIAS): + # Issue #350 + pytest.skip("This confis is skipped due to known numerical failure.") """ Test backward with parameterized configs """ From 1e97485c43598719ca5f660159f21f9b0483bdd8 Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Mon, 31 Aug 2026 10:23:54 -0400 Subject: [PATCH 4/7] Update QoLA with CK patches fix, restore skip test, add custom QoLA support --- 3rdparty/QoLA | 2 +- tests/jax/test_fused_attn.py | 5 ----- transformer_engine/common/ck_fused_attn/CMakeLists.txt | 4 ++++ 3 files changed, 5 insertions(+), 6 deletions(-) diff --git a/3rdparty/QoLA b/3rdparty/QoLA index a1f9245e42..8dd830d7d6 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit a1f9245e4263a172d56e1ec490b3729c0f790b6c +Subproject commit 8dd830d7d63180b6a197b4ab760a5e0a07f17e8b diff --git a/tests/jax/test_fused_attn.py b/tests/jax/test_fused_attn.py index 36dde0d492..bf5e31d39b 100644 --- a/tests/jax/test_fused_attn.py +++ b/tests/jax/test_fused_attn.py @@ -1897,11 +1897,6 @@ def test_backward( swa, seq_desc_format, ): - if (is_hip_extension() and qkv_layout == QKVLayout.THD_THD_THD and - attn_mask_type == AttnMaskType.PADDING_MASK and - attn_bias_type == AttnBiasType.NO_BIAS): - # Issue #350 - pytest.skip("This confis is skipped due to known numerical failure.") """ Test backward with parameterized configs """ diff --git a/transformer_engine/common/ck_fused_attn/CMakeLists.txt b/transformer_engine/common/ck_fused_attn/CMakeLists.txt index 5e30c6da5a..e6ee23dfa0 100644 --- a/transformer_engine/common/ck_fused_attn/CMakeLists.txt +++ b/transformer_engine/common/ck_fused_attn/CMakeLists.txt @@ -24,6 +24,10 @@ if(NOT _gfx1250_idx EQUAL -1) endif() set(__QOLA_DIR "${CMAKE_CURRENT_LIST_DIR}/../../../3rdparty/QoLA") +if (DEFINED ENV{NVTE_QOLA_DIR} AND NOT $ENV{NVTE_QOLA_DIR} STREQUAL "") + set(__QOLA_DIR $ENV{NVTE_QOLA_DIR}) + message(STATUS "Using QoLA from NVTE_QOLA_DIR=${__QOLA_DIR}.") +endif() set(__AITER_SOURCE_DIR "${CMAKE_CURRENT_BINARY_DIR}/qola/third_party/aiter") if (DEFINED ENV{NVTE_AITER_SOURCE_DIR} AND NOT $ENV{NVTE_AITER_SOURCE_DIR} STREQUAL "") set(__AITER_SOURCE_DIR $ENV{NVTE_AITER_SOURCE_DIR}) From 12ad9e7b867ca4327addcf29ba91ae83fff1ca0b Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Wed, 2 Sep 2026 16:01:08 -0400 Subject: [PATCH 5/7] Update CK-JIT and QoLA to main branch-es commits --- 3rdparty/QoLA | 2 +- 3rdparty/ck_jit | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/3rdparty/QoLA b/3rdparty/QoLA index 8dd830d7d6..100b315ab2 160000 --- a/3rdparty/QoLA +++ b/3rdparty/QoLA @@ -1 +1 @@ -Subproject commit 8dd830d7d63180b6a197b4ab760a5e0a07f17e8b +Subproject commit 100b315ab23e26673519ecbdd7c9d8564b9a367e diff --git a/3rdparty/ck_jit b/3rdparty/ck_jit index 3cf034d344..e0b303100e 160000 --- a/3rdparty/ck_jit +++ b/3rdparty/ck_jit @@ -1 +1 @@ -Subproject commit 3cf034d3445d7faf5de852840c4419ab720483cf +Subproject commit e0b303100e7deda4c117acabaf49e825243f8748 From 9048129aa891d0b40f202af00bca9653f4e4e841 Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Thu, 3 Sep 2026 12:02:51 -0400 Subject: [PATCH 6/7] Gfx1250 changes on top of updated AITER --- tests/pytorch/attention/test_attention.py | 8 +- .../attention/test_attention_gfx1250.py | 382 ++++++++++++++++++ .../common/ck_fused_attn/CMakeLists.txt | 18 +- .../common/ck_fused_attn/qola_manifest.toml | 2 +- .../ck_fused_attn/src/ck_fused_attn_fwd.cpp | 110 ++++- .../common/fused_attn_rocm/fused_attn.cpp | 6 - .../fused_attn_rocm/fused_attn_aotriton.cpp | 5 + 7 files changed, 503 insertions(+), 28 deletions(-) create mode 100644 tests/pytorch/attention/test_attention_gfx1250.py diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index bcd2c48cf9..1bd7ea0272 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -310,7 +310,9 @@ def test_dot_product_attention( # Skip if only unfused backend is supported # Double-count the CK backend since we want to compare V2/V3 kernels - has_ck_backend = IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends + # Issue #16948 CK V2 is disabled for gfx1250 + has_ck_backend = ( IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends and + get_device_compute_capability() != (12, 5) ) if not has_ck_backend and ( len(fused_attn_backends) + flash_attn_supported + unfused_attn_supported ) < 2: @@ -1667,7 +1669,9 @@ def test_transformer_layer( # Skip if only unfused backend is supported # Double-count the CK backend since we want to compare V2/V3 kernels - has_ck_backend = IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends + # Issue #16948 CK V2 is disabled for gfx1250 + has_ck_backend = ( IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends and + get_device_compute_capability() != (12, 5) ) if not has_ck_backend and ( len(fused_attn_backends) + flash_attn_supported + unfused_attn_supported ) < 2: diff --git a/tests/pytorch/attention/test_attention_gfx1250.py b/tests/pytorch/attention/test_attention_gfx1250.py new file mode 100644 index 0000000000..17787b7c64 --- /dev/null +++ b/tests/pytorch/attention/test_attention_gfx1250.py @@ -0,0 +1,382 @@ +# Copyright (c) 2026, Advanced Micro Devices, Inc. All rights reserved. +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +""" +Gfx1250-targeted attention tests. + +Configs mirror the smoke tests in: + 3rdparty/aiter/op_tests/cpp/mha/smoke_test_fwd_v3_gfx1250.sh (fwd V3) + +FWD V3 notes (fmha_fwd_gfx1250_batched / fmha_fwd_with_sink_asm): + - Both D64 and D128 require a non-null sink_addr (fixed kernarg layout). + TE supplies a static [256] fp32 buffer initialized to -1e30f so that + exp(-1e30f) ≈ 0.0f adds no effective weight for non-fully-masked rows. + - D64 (ENABLE_SINK=1): kernel reads and uses the sink values as a logit + floor. Top-left causal; sq ≤ sk (rectangular) safe because even with + sink≈0 the real attention weights dominate. + - D128 (ENABLE_SINK=0): kernel ignores the sink values entirely. + Causal (top-left or bottom-right); sq == sk only — rectangular shapes + risk NaN on fully-masked KV tiles because the sink floor is disabled. + - No SWA (window_size_left must be -1). + +Each test forces CK V3 and compares against a pure-PyTorch scaled dot product +attention reference (torch.nn.functional.scaled_dot_product_attention) computed +in float32 for numerical stability. +""" +import os +import sys +import pathlib + +import pytest +import torch +import torch.nn.functional as F +from torch.utils.cpp_extension import IS_HIP_EXTENSION + +from transformer_engine.pytorch import DotProductAttention, get_device_compute_capability +from transformer_engine.pytorch.cpp_extensions.fused_attn import FusedAttnBackend +from transformer_engine.pytorch.attention.dot_product_attention import _attention_backends +from transformer_engine.pytorch.distributed import CudaRNGStatesTracker + +_current_file = pathlib.Path(__file__).resolve() +sys.path.append(str(_current_file.parent.parent)) +from utils import ( + reset_rng_states, + ModelConfig, + get_available_attention_backends, +) + +# Whole file is ROCm + gfx1250 only. +pytestmark = [ + pytest.mark.skipif(not IS_HIP_EXTENSION, reason="ROCm TE specific test."), + pytest.mark.skipif( + get_device_compute_capability() != (12, 5), + reason="gfx1250 (compute capability 12.5) required.", + ), +] + +_SEED = 1234 + + +def _make_rng_tracker(): + tracker = CudaRNGStatesTracker() + tracker.add("model-parallel-rng", _SEED) + return tracker + + +def _run_dpa( + config: ModelConfig, + backend_env: dict, + dtype: torch.dtype, + is_training: bool, + q: torch.Tensor = None, + k: torch.Tensor = None, + v: torch.Tensor = None, +): + """Run one forward (and optional backward) pass, return (out, grads). + + If q/k/v are provided they are used directly (requires_grad set appropriately); + otherwise fresh random tensors are created. + """ + reset_rng_states() + + for key in [ + "NVTE_FLASH_ATTN", "NVTE_FUSED_ATTN", "NVTE_UNFUSED_ATTN", + "NVTE_FUSED_ATTN_CK", "NVTE_FUSED_ATTN_AOTRITON", + "NVTE_CK_USES_FWD_V3", "NVTE_CK_USES_BWD_V3", + ]: + os.environ.pop(key, None) + for key, val in backend_env.items(): + os.environ[key] = str(val) + _attention_backends["backend_selection_requires_update"] = True + + device = "cuda" + b = config.batch_size + sq = config.max_seqlen_q + sk = config.max_seqlen_kv + hq = config.num_heads + hk = config.num_gqa_groups + dqk = config.head_dim_qk + + if q is None: + q = torch.randn(b, sq, hq, dqk, dtype=dtype, device=device) + k = torch.randn(b, sk, hk, dqk, dtype=dtype, device=device) + v = torch.randn(b, sk, hk, dqk, dtype=dtype, device=device) + + q = q.detach().requires_grad_(is_training) + k = k.detach().requires_grad_(is_training) + v = v.detach().requires_grad_(is_training) + + block = DotProductAttention( + num_attention_heads=hq, + kv_channels=dqk, + num_gqa_groups=hk, + attention_dropout=0.0, + qkv_format="bshd", + attn_mask_type=config.attn_mask_type, + tp_size=1, + tp_group=None, + get_rng_state_tracker=_make_rng_tracker, + ).to(dtype=dtype, device=device) + if not is_training: + block.eval() + + out = block( + q, k, v, + qkv_format="bshd", + attn_mask_type=config.attn_mask_type, + window_size=config.window_size, + max_seqlen_q=sq, + max_seqlen_kv=sk, + ) + + grads = None + if is_training: + out.sum().backward() + grads = (q.grad.clone(), k.grad.clone(), v.grad.clone()) + + return out.detach(), grads + + +def _pytorch_ref( + config: ModelConfig, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + is_training: bool, +) -> tuple: + """Pure-PyTorch SDPA reference computed in float32. + + Inputs are BSHD; GQA heads are expanded before calling SDPA. + Returns (out_bf16, (dq, dk, dv)) where dq/dk/dv are None when not training. + The reference runs in float32 for numerical stability so that the comparison + is against a correct high-precision result rather than a TE-specific backend + that may itself have issues on gfx1250. + """ + b, sq, hq, d = q.shape + hk = k.shape[2] + + # Upcast to float32 for reference stability. + qf = q.float() + kf = k.float() + vf = v.float() + + if is_training: + qf = qf.detach().requires_grad_(True) + kf = kf.detach().requires_grad_(True) + vf = vf.detach().requires_grad_(True) + + # SDPA expects [B, H, S, D]. + qf_t = qf.permute(0, 2, 1, 3) # [b, hq, sq, d] + kf_t = kf.permute(0, 2, 1, 3) # [b, hk, sk, d] + vf_t = vf.permute(0, 2, 1, 3) + + # Expand KV heads for GQA so SDPA sees [b, hq, sk, d]. + # Use a separate leaf for the expanded tensor so grad flows back to hk heads. + gqa = hk != hq + if gqa: + kf_t_exp = kf_t.repeat_interleave(hq // hk, dim=1).detach().requires_grad_(is_training) + vf_t_exp = vf_t.repeat_interleave(hq // hk, dim=1).detach().requires_grad_(is_training) + else: + kf_t_exp = kf_t + vf_t_exp = vf_t + + attn_mask_type = config.attn_mask_type + window_size = config.window_size + sk_len = kf_t.shape[2] + + # Use explicit mask for bottom-right causal or SWA; otherwise let SDPA handle it. + has_swa = window_size is not None and window_size not in ((-1, -1), (-1, 0)) + needs_explicit_mask = attn_mask_type == "causal_bottom_right" or has_swa + + if needs_explicit_mask: + # Start with all-attend, then apply causal + SWA constraints. + rows = torch.arange(sq, device=q.device).unsqueeze(1) # [sq, 1] + cols = torch.arange(sk_len, device=q.device).unsqueeze(0) # [1, sk] + mask = torch.ones(sq, sk_len, dtype=torch.bool, device=q.device) + if attn_mask_type in ("causal", "causal_bottom_right"): + offset = sk_len - sq if attn_mask_type == "causal_bottom_right" else 0 + mask = mask & (cols <= rows + offset) + if has_swa: + left, right = window_size + lo = (rows - left).clamp(min=0) if left >= 0 else torch.zeros_like(rows) + hi = (rows + right) if right >= 0 else torch.full_like(rows, sk_len - 1) + mask = mask & (cols >= lo) & (cols <= hi) + float_mask = torch.zeros(sq, sk_len, dtype=torch.float32, device=q.device) + float_mask[~mask] = float("-inf") + out_t = F.scaled_dot_product_attention(qf_t, kf_t_exp, vf_t_exp, attn_mask=float_mask) + elif attn_mask_type == "no_mask": + out_t = F.scaled_dot_product_attention(qf_t, kf_t_exp, vf_t_exp, is_causal=False) + else: + # causal top-left: PyTorch is_causal=True is exactly this + out_t = F.scaled_dot_product_attention(qf_t, kf_t_exp, vf_t_exp, is_causal=True) + + out_bf16 = out_t.permute(0, 2, 1, 3).to(dtype=q.dtype).detach() # [b, sq, hq, d] + + dq = dk = dv = None + if is_training: + out_t.sum().backward() + dq = qf.grad.to(dtype=q.dtype).detach() # [b, sq, hq, d] + if gqa: + # kf_t_exp.grad: [b, hq, sk, d] — reduce over the hq/hk groups back to hk heads + group = hq // hk + dk_exp = kf_t_exp.grad.view(b, hk, group, sk_len, d).sum(dim=2) # [b, hk, sk, d] + dv_exp = vf_t_exp.grad.view(b, hk, group, sk_len, d).sum(dim=2) + dk = dk_exp.permute(0, 2, 1, 3).to(dtype=k.dtype).detach() # [b, sk, hk, d] + dv = dv_exp.permute(0, 2, 1, 3).to(dtype=v.dtype).detach() + else: + dk = kf.grad.to(dtype=k.dtype).detach() # [b, sk, hk, d] + dv = vf.grad.to(dtype=v.dtype).detach() + + return out_bf16, (dq, dk, dv) + + +def _compare(config: ModelConfig, dtype: torch.dtype = torch.bfloat16, is_training: bool = True): + """Run CK V3 and compare against a float32 PyTorch SDPA reference.""" + tols = dict(atol=2e-2, rtol=2e-2) + + _, _, fused_backends = get_available_attention_backends( + config, + qkv_dtype=dtype, + qkv_layout="bshd_bshd_bshd", + is_training=is_training, + ) + if FusedAttnBackend["CK"] not in fused_backends: + pytest.skip("CK backend not available for this config") + + reset_rng_states() + for _evar in [ + "NVTE_FLASH_ATTN", "NVTE_FUSED_ATTN", "NVTE_UNFUSED_ATTN", + "NVTE_FUSED_ATTN_CK", "NVTE_FUSED_ATTN_AOTRITON", + "NVTE_CK_USES_FWD_V3", "NVTE_CK_USES_BWD_V3", + ]: + os.environ.pop(_evar, None) + + device = "cuda" + b = config.batch_size + sq = config.max_seqlen_q + sk = config.max_seqlen_kv + hq = config.num_heads + hk = config.num_gqa_groups + dqk = config.head_dim_qk + + q = torch.randn(b, sq, hq, dqk, dtype=dtype, device=device) + k = torch.randn(b, sk, hk, dqk, dtype=dtype, device=device) + v = torch.randn(b, sk, hk, dqk, dtype=dtype, device=device) + + ref_out, ref_grads = _pytorch_ref(config, q, k, v, is_training) + + ck_v3_env = { + "NVTE_FUSED_ATTN": "1", + "NVTE_FLASH_ATTN": "0", + "NVTE_UNFUSED_ATTN": "0", + "NVTE_FUSED_ATTN_CK": "1", + "NVTE_FUSED_ATTN_AOTRITON": "0", + "NVTE_CK_USES_FWD_V3": "1", + "NVTE_CK_USES_BWD_V3": "1", + } + ck_out, ck_grads = _run_dpa(config, ck_v3_env, dtype, is_training, q=q, k=k, v=v) + + # TE DotProductAttention returns [b, sq, hq*d] for bshd format; reshape to [b, sq, hq, d]. + ck_out = ck_out.view(b, sq, hq, dqk) + + torch.testing.assert_close(ck_out, ref_out, **tols) + if is_training and ref_grads is not None: + dq_ref, dk_ref, dv_ref = ref_grads + dq_ck, dk_ck, dv_ck = ck_grads + # grads from _run_dpa are already [b, s, h, d] since q/k/v were passed as BSHD leaves + torch.testing.assert_close(dq_ck, dq_ref, **tols) + torch.testing.assert_close(dk_ck, dk_ref, **tols) + torch.testing.assert_close(dv_ck, dv_ref, **tols) + + +# --------------------------------------------------------------------------- +# FWD V3 D64 configs (smoke_test_fwd_sink.sh — run_d64) +# +# Kernel: fmha_fwd_with_sink_asm, ENABLE_SINK=1, bottom-right causal. +# Sink floor prevents div-by-zero, so sq < sk (rectangular) is safe. +# h=8, h_k∈{1,2,4}, b∈{1,2}. FWD-only (no backward for this path). +# +# NOTE: the smoke test labels some cases "mha" but actually uses h_k=1 +# (GQA-1 / MQA). The TE test mirrors that: "mha" suffix = num_gqa_groups=1. +# The square cases (sq==sk) additionally exercise num_gqa_groups=8 (true MHA). +# --------------------------------------------------------------------------- + +_fwd_v3_d64: dict = {} +# The gfx1250 ASM kernel implements bottom-right causal (mask_type 1 and 2 both +# map to is_causal=1 inside fmha_fwd_with_sink_asm, which follows bottom-right +# causal semantics). Square shapes (sq==sk) are equivalent for both variants; +# rectangular shapes (sq < sk) must use "causal_bottom_right" to match. +for _s in (512, 1024, 2048): + # h_k=1 matches smoke test "mha" label; also add true MHA (h_k=h=8) for square. + _fwd_v3_d64[f"d64_sq{_s}_sk{_s}_b1_mqa"] = ModelConfig( + 1, _s, 8, 64, max_seqlen_kv=_s, num_gqa_groups=1, attn_mask_type="causal_bottom_right") + _fwd_v3_d64[f"d64_sq{_s}_sk{_s}_b1_mha"] = ModelConfig( + 1, _s, 8, 64, max_seqlen_kv=_s, num_gqa_groups=8, attn_mask_type="causal_bottom_right") + _fwd_v3_d64[f"d64_sq{_s}_sk{_s}_b2_gqa2"] = ModelConfig( + 2, _s, 8, 64, max_seqlen_kv=_s, num_gqa_groups=2, attn_mask_type="causal_bottom_right") +for _sq in (128, 256, 512): + for _sk in (512, 2048): + for _b in (1, 2): + for _hq, _hk in ((8, 1), (8, 2), (4, 4)): + _key = f"d64_sq{_sq}_sk{_sk}_b{_b}_h{_hq}k{_hk}" + _fwd_v3_d64[_key] = ModelConfig( + _b, _sq, _hq, 64, max_seqlen_kv=_sk, num_gqa_groups=_hk, + attn_mask_type="causal_bottom_right") +# Rectangular tail configs: smoke test uses h_k=1 for these non-standard shapes. +for _sq in (130, 300): + _fwd_v3_d64[f"d64_sq{_sq}_sk2048_b1_mha"] = ModelConfig( + 1, _sq, 8, 64, max_seqlen_kv=2048, num_gqa_groups=1, attn_mask_type="causal_bottom_right") +for _sk in (768, 2300): + _fwd_v3_d64[f"d64_sq128_sk{_sk}_b1_mha"] = ModelConfig( + 1, 128, 8, 64, max_seqlen_kv=_sk, num_gqa_groups=1, attn_mask_type="causal_bottom_right") + + +@pytest.mark.parametrize("model", sorted(_fwd_v3_d64.keys())) +def test_gfx1250_fwd_v3_d64(model): + """FWD V3 D64 correctness — run_d64 configs from smoke_test_fwd_v3_gfx1250.sh.""" + _compare(_fwd_v3_d64[model], dtype=torch.bfloat16, is_training=False) + + +# --------------------------------------------------------------------------- +# FWD V3 D128 configs (smoke_test_fwd_sink.sh — run_d128) +# +# Kernel: fmha_fwd_with_sink_asm, ENABLE_SINK=0 (sink_ptr ignored/nullptr). +# D128 uses bottom-right causal (kernel is_causal=1, same path as D64). +# h=8, h_k∈{1,2,4}, b∈{1,2}. FWD-only. +# +# NOTE: smoke test "mha" cases use h_k=1 (MQA); true MHA (h_k=8) is added +# for square shapes only. +# --------------------------------------------------------------------------- + +_fwd_v3_d128: dict = {} +for _s in (512, 1024, 2048): + # h_k=1 matches smoke test; also add true MHA for square shapes. + _fwd_v3_d128[f"d128_sq{_s}_sk{_s}_b1_mqa"] = ModelConfig( + 1, _s, 8, 128, max_seqlen_kv=_s, num_gqa_groups=1, + attn_mask_type="causal_bottom_right") + _fwd_v3_d128[f"d128_sq{_s}_sk{_s}_b1_mha"] = ModelConfig( + 1, _s, 8, 128, max_seqlen_kv=_s, num_gqa_groups=8, + attn_mask_type="causal_bottom_right") + _fwd_v3_d128[f"d128_sq{_s}_sk{_s}_b2_gqa2"] = ModelConfig( + 2, _s, 8, 128, max_seqlen_kv=_s, num_gqa_groups=2, + attn_mask_type="causal_bottom_right") +# Rectangular configs from smoke test run_d128 (bottom-right causal, sq < sk safe). +for _sq in (128, 256): + for _hq, _hk in ((8, 1), (8, 2), (4, 4)): + _fwd_v3_d128[f"d128_sq{_sq}_sk2048_h{_hq}k{_hk}"] = ModelConfig( + 1, _sq, _hq, 128, max_seqlen_kv=2048, num_gqa_groups=_hk, + attn_mask_type="causal_bottom_right") +# Unaligned tail configs: h_k=1 matching smoke test. +_fwd_v3_d128["d128_sq130_sk2048_b1_mha"] = ModelConfig( + 1, 130, 8, 128, max_seqlen_kv=2048, num_gqa_groups=1, + attn_mask_type="causal_bottom_right") +_fwd_v3_d128["d128_sq128_sk2300_b1_mha"] = ModelConfig( + 1, 128, 8, 128, max_seqlen_kv=2300, num_gqa_groups=1, + attn_mask_type="causal_bottom_right") + + +@pytest.mark.parametrize("model", sorted(_fwd_v3_d128.keys())) +def test_gfx1250_fwd_v3_d128(model): + """FWD V3 D128 correctness — run_d128 configs from smoke_test_fwd_v3_gfx1250.sh.""" + _compare(_fwd_v3_d128[model], dtype=torch.bfloat16, is_training=False) diff --git a/transformer_engine/common/ck_fused_attn/CMakeLists.txt b/transformer_engine/common/ck_fused_attn/CMakeLists.txt index e6ee23dfa0..539849db22 100644 --- a/transformer_engine/common/ck_fused_attn/CMakeLists.txt +++ b/transformer_engine/common/ck_fused_attn/CMakeLists.txt @@ -8,21 +8,6 @@ project(ck_fused_attn LANGUAGES HIP CXX) set(AITER_MHA_INSTALL_DIR "${CMAKE_INSTALL_PREFIX}/transformer_engine/lib") -#Corresponding runtime check is in nvte_get_fused_attn_backend() -list(FIND CMAKE_HIP_ARCHITECTURES "gfx1250" _gfx1250_idx) -if(NOT _gfx1250_idx EQUAL -1) - message(WARNING - "Removing unsupported gfx1250 from CMAKE_HIP_ARCHITECTURES for ck_fused_attn build.") - list(REMOVE_ITEM CMAKE_HIP_ARCHITECTURES "gfx1250") - list(LENGTH CMAKE_HIP_ARCHITECTURES _hip_arch_count) - if(_hip_arch_count EQUAL 0) - message(FATAL_ERROR - "No supported architectures remain for the ck_fused_attn build. " - "Re-run the build with FUSED_ATTN_CK backend disabled.") - endif() - set(GPU_TARGETS ${CMAKE_HIP_ARCHITECTURES}) -endif() - set(__QOLA_DIR "${CMAKE_CURRENT_LIST_DIR}/../../../3rdparty/QoLA") if (DEFINED ENV{NVTE_QOLA_DIR} AND NOT $ENV{NVTE_QOLA_DIR} STREQUAL "") set(__QOLA_DIR $ENV{NVTE_QOLA_DIR}) @@ -228,7 +213,7 @@ set(CK_FUSED_ATTN_COMPILE_OPTIONS) #pragma in CK headers which is not supported by the compiler used to build TE. list(APPEND CK_FUSED_ATTN_COMPILE_OPTIONS -DCK_TILE_FLOAT_TO_BFLOAT16_DEFAULT=${CK_FUSED_ATTN_FLOAT_TO_BFLOAT16_DEFAULT} - -DENABLE_CK=1 -DFA_WITH_NATIVE_SPLITKV=1 -Wno-unused-value -Wno-unknown-warning-option) + -DENABLE_CK=1 -DFA_WITH_NATIVE_SPLITKV=1 -DFA_WITH_SINK=1 -Wno-unused-value -Wno-unknown-warning-option) # Public QoLA headers ship alongside the .so libs in ${__AITER_MHA_PATH}/../include # (emitted by qola.cli build, or copied from the QoLA build dir above for the @@ -249,6 +234,7 @@ target_link_libraries(ck_fused_attn PUBLIC ${ck_fused_attn_LINKER_LIBS}) target_compile_options(ck_fused_attn PRIVATE ${CK_FUSED_ATTN_COMPILE_OPTIONS}) set_target_properties(ck_fused_attn PROPERTIES INSTALL_RPATH "$ORIGIN") +# In CK_JIT mode the AITER libs are emitted directly into AITER_MHA_INSTALL_DIR if (NOT "${__AITER_MHA_PATH}" STREQUAL "${AITER_MHA_INSTALL_DIR}") install(FILES ${__AITER_MHA_PATH}/te_libmha_fwd.so ${__AITER_MHA_PATH}/te_libmha_bwd.so DESTINATION ${AITER_MHA_INSTALL_DIR}) endif() diff --git a/transformer_engine/common/ck_fused_attn/qola_manifest.toml b/transformer_engine/common/ck_fused_attn/qola_manifest.toml index a689234445..091fd2fbdb 100644 --- a/transformer_engine/common/ck_fused_attn/qola_manifest.toml +++ b/transformer_engine/common/ck_fused_attn/qola_manifest.toml @@ -12,7 +12,7 @@ mode = "cpp_itfs" receipt = 700 drop_srcs = ["mha_fwd_split.cu", "mha_fwd_batch_prefill.cu"] drop_directions = ["fwd_splitkv", "batch_prefill"] -flags_extra_cc = ["'-DFA_WITH_NATIVE_SPLITKV=1'"] +flags_extra_cc = ["'-DFA_WITH_NATIVE_SPLITKV=1'", "'-DFA_WITH_SINK=1'"] [[modules]] name = "libmha_bwd" diff --git a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp index e6700147bf..f270cabf56 100644 --- a/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp +++ b/transformer_engine/common/ck_fused_attn/src/ck_fused_attn_fwd.cpp @@ -6,14 +6,103 @@ #include #include +#include #include +#include #include +#include +#include #include "ck_fused_attn/ck_fused_attn.hpp" #include "qola_mha_fwd.h" #include "ck_fused_attn_utils.hpp" namespace ck_fused_attn{ +#if FA_WITH_SINK +namespace { +// D64 gfx1250 fmha_fwd_with_sink_asm (ENABLE_SINK=1): requires non-null sink_ptr +// of shape [nhead] fp32 in "AITER post-scale domain". The kernel adds +// exp(sink_val[h]) to every row's softmax denominator. We initialize to +// -1e30f so expf(-1e30f)=0.0f in fp32 — zero contribution, matching the +// UnfusedDotProductAttention reference which has no sink term. +// D128 (ENABLE_SINK=0): dispatch guard rejects sink_ptr!=nullptr; leave null. +// +// The buffer is read-only to the kernel (an input logit vector, never written +// back), so a single immutable buffer can be shared safely by any number of +// concurrent kernel launches / host threads. +// +// We keep one buffer *per device*, allocated lazily and kept for the process +// lifetime. +constexpr int kSinkBufMaxHeads = 256; +static std::mutex s_sink_mutex; +static std::unordered_map s_sink_bufs; // device id -> device buffer + +// Resolve the device a kernel launched on `stream` will execute on. +int device_for_stream(hipStream_t stream){ + hipDevice_t hip_dev = 0; + if(stream != nullptr) { + if (hipStreamGetDevice(stream, &hip_dev) == hipSuccess){ + return static_cast(hip_dev); + } + } + // Default/null stream, or query unsupported/failed: use the current device. + int dev = 0; + if(hipGetDevice(&dev) != hipSuccess){ + throw std::runtime_error( + "ck_fused_attn fwd: hipGetDevice failed while resolving gfx1250 sink buffer device."); + } + return dev; +} + +const void* get_gfx1250_sink_buf(int dev){ + std::lock_guard lock(s_sink_mutex); + + auto it = s_sink_bufs.find(dev); + if(it != s_sink_bufs.end()){ + return it->second; + } + + // Allocate on the intended device regardless of the current device, then + // restore the previous current device so we don't perturb the caller. + int prev_dev = 0; + bool switched = false; + if(hipGetDevice(&prev_dev) == hipSuccess && prev_dev != dev){ + if(hipSetDevice(dev) == hipSuccess){ + switched = true; + } + } + + float* buf = nullptr; + hipError_t err = hipMalloc(&buf, kSinkBufMaxHeads * sizeof(float)); + if(err == hipSuccess && buf != nullptr){ + std::vector fill(kSinkBufMaxHeads, -1e30f); + err = hipMemcpy(buf, fill.data(), + kSinkBufMaxHeads * sizeof(float), hipMemcpyHostToDevice); + if(err != hipSuccess){ + hipFree(buf); + buf = nullptr; + } + } else { + buf = nullptr; + } + + if(switched){ + hipSetDevice(prev_dev); + } + + if(buf == nullptr){ + throw std::runtime_error( + std::string("ck_fused_attn fwd: failed to allocate/initialize gfx1250 sink " + "buffer on device ") + std::to_string(dev) + + " (" + hipGetErrorString(err) + ")."); + } + + s_sink_bufs.emplace(dev, buf); + return buf; +} +} // namespace +#endif //FA_WITH_SINK + // print the fmha traits and fmha_args when calling ck apis void log_fwd_config(const char* func_name, bool has_dropout, const aiter::mha_fwd_args& fmha_args, std::ostream* log_file){ @@ -90,6 +179,10 @@ void log_fwd_config(const char* func_name, bool has_dropout, const aiter::mha_fw log_value(log_file, "dropout_seed_ptr", std::get<0>(std::get>(fmha_args.drop_seed_offset))); log_value(log_file, "dropout_offset_ptr", std::get<1>(std::get>(fmha_args.drop_seed_offset))); + + log_value(log_file, "num_splits", fmha_args.num_splits); + log_value(log_file, "splitkv_workspace_ptr", fmha_args.splitkv_workspace_ptr); + log_value(log_file, "sink_ptr", fmha_args.sink_ptr); } void dump_fwd_timings(const char* dump_path, float average_runtime){ @@ -103,7 +196,7 @@ void dump_fwd_timings(const char* dump_path, float average_runtime){ // never disagree with the launch. v3_api_check is left false here; callers flip it. // The stream-dependent max_seqlen override (NVTE_CK_RUNTIME_MAX_SEQLEN) is applied // by ck_attn_fwd after this returns; it does not affect v3 kernel selection. -aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ +aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args, hipStream_t stream = nullptr){ bias_enum bias_type = bias_enum::no_bias; BiasShape bias_shape = BiasShape::k11SS; @@ -212,6 +305,11 @@ aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ fmha_args.num_splits = args.num_splits; fmha_args.splitkv_workspace_ptr = args.splitkv_workspace_ptr; +#if FA_WITH_SINK + if(args.h <= kSinkBufMaxHeads && QOLA_NS(mha_fwd_with_sink_supported)(fmha_args)) { + fmha_args.sink_ptr = get_gfx1250_sink_buf(device_for_stream(stream)); + } +#endif return fmha_args; } @@ -219,11 +317,18 @@ aiter::mha_fwd_args build_fwd_fmha_args(const CKAttnFwdArgs& args){ // launching any kernel. Builds the same args as ck_attn_fwd and relies on AITER's // v3_api_check dry-run (returns 1 when v3 is available, -1 otherwise). bool ck_attn_fwd_uses_v3(const CKAttnFwdArgs& args){ +#if 1 + //Split-KV and Sink dispatchers do not honor v3_api_check + // so we cannot rely on the AITER backend to tell us whether v3 is available. + throw std::runtime_error( + "FWD path V3 API checking is not robust on the backend."); +#else aiter::mha_fwd_args fmha_args = build_fwd_fmha_args(args); fmha_args.v3_api_check = true; // No kernel is launched in check mode, so the stream/log flags are irrelevant. ck_tile::stream_config stream_config{nullptr, false, false}; return QOLA_NS(mha_fwd)(fmha_args, stream_config) == 1; +#endif } hipError_t ck_attn_fwd(const CKAttnFwdArgs& args, hipStream_t stream){ @@ -240,7 +345,7 @@ hipError_t ck_attn_fwd(const CKAttnFwdArgs& args, hipStream_t stream){ // print kernel name on verbose mode ck_tile::stream_config stream_config{stream, dump_path!=nullptr, get_ck_log_stream() != nullptr}; - aiter::mha_fwd_args fmha_args = build_fwd_fmha_args(args); + aiter::mha_fwd_args fmha_args = build_fwd_fmha_args(args, stream); if(const char* env_p = std::getenv("NVTE_CK_RUNTIME_MAX_SEQLEN")){ if(args.is_group_mode() && std::string(env_p) == "1"){ @@ -285,4 +390,3 @@ size_t ck_attn_fwd_workspace_size(const CKAttnFwdArgs& args){ } }//namespace ck_fused_attn - diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp index 29b0806988..d32bf18152 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn.cpp @@ -282,12 +282,6 @@ NVTE_Fused_Attn_Backend nvte_get_fused_attn_backend( int64_t window_size_right, bool return_max_logit, bool /*cuda_graph*/, bool deterministic) { using namespace transformer_engine; - //gfx1250 is disabled in ck_fused_attn/CMakeLists.txt and is not supported by curretnt aotriton - const int gpu_arch = cuda::sm_arch(cuda::current_device()); - if (gpu_arch == 125) { - return NVTE_Fused_Attn_Backend::NVTE_No_Backend; - } - // TODO: Add return_max_logit support if (return_max_logit) return NVTE_Fused_Attn_Backend::NVTE_No_Backend; diff --git a/transformer_engine/common/fused_attn_rocm/fused_attn_aotriton.cpp b/transformer_engine/common/fused_attn_rocm/fused_attn_aotriton.cpp index 2928520fc4..caf08565eb 100644 --- a/transformer_engine/common/fused_attn_rocm/fused_attn_aotriton.cpp +++ b/transformer_engine/common/fused_attn_rocm/fused_attn_aotriton.cpp @@ -54,6 +54,11 @@ bool is_aotriton_backend_supported( int64_t window_size_right) { #ifdef USE_FUSED_ATTN_AOTRITON + // AOTriton has no gfx1250 support. + if(cuda::sm_arch(cuda::current_device()) == 125){ + return false; + } + //TODO: release after AOTriton support support Multi-latent attention if(head_dim_qk != head_dim_v){ return false; From fe1cf36d714cf33dac2ac87e6291c49995efbc41 Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Thu, 3 Sep 2026 12:10:03 -0400 Subject: [PATCH 7/7] Enable CK tests. Enable MXFP4, MXFP8 and NVFP4 in core, still disabled for FW integration --- examples/jax/encoder/common.py | 2 +- tests/cpp/operator/CMakeLists.txt | 11 +-- .../cpp/operator/test_cast_mxfp4_transpose.cu | 2 + tests/cpp/operator/test_cast_mxfp8.cu | 4 + .../cpp/operator/test_cast_nvfp4_transpose.cu | 2 + tests/cpp/operator/test_cublaslt_gemm.cu | 89 ++++++++++--------- tests/pytorch/attention/test_attention.py | 8 +- tests/pytorch/test_grouped_linear.py | 5 +- .../common/cast/mxfp4/quantize_mxfp4.cuh | 6 +- transformer_engine/pytorch/quantization.py | 28 +++--- 10 files changed, 88 insertions(+), 69 deletions(-) diff --git a/examples/jax/encoder/common.py b/examples/jax/encoder/common.py index 4a8efb517a..17c624aa09 100644 --- a/examples/jax/encoder/common.py +++ b/examples/jax/encoder/common.py @@ -56,7 +56,7 @@ def is_nvfp4_supported(): gpu_arch = get_device_compute_capability(0) if is_hip_extension(): # only GFX12.5 machines support nvfp4 - return False #TODO add gfx1250 (gpu_arch == 125) when ready + return gpu_arch >= 120 return gpu_arch >= 100 diff --git a/tests/cpp/operator/CMakeLists.txt b/tests/cpp/operator/CMakeLists.txt index 2b1ce6ef5e..b9e872e5cf 100644 --- a/tests/cpp/operator/CMakeLists.txt +++ b/tests/cpp/operator/CMakeLists.txt @@ -50,21 +50,18 @@ if(USE_ROCM) test_dequantize_nvfp4.cu test_cublaslt_gemm.cu test_ck_grouped_gemm.cu + test_ck_grouped_mxfp8.cu test_cast_mxfp4_transpose.cu test_multi_quantize_mxfp8.cu) - list(FIND CMAKE_HIP_ARCHITECTURES "gfx1250" _gfx1250_idx) - if(NOT _gfx1250_idx EQUAL -1) - list(APPEND test_cuda_sources test_ck_grouped_mxfp8.cu) - target_include_directories(test_operator BEFORE PRIVATE - ${CMAKE_CURRENT_SOURCE_DIR}/../../../3rdparty/composable_kernel/include) - endif() - TE_GetHipifiedSources("${test_cuda_sources}" ${CMAKE_CURRENT_SOURCE_DIR} test_hip_sources) TE_AddHipifyDeps("${test_cuda_sources}" ${CMAKE_CURRENT_SOURCE_DIR}) message("${message_line}") message(STATUS "test_operator hipified sources: ${test_hip_sources}") set_target_properties(test_operator PROPERTIES SOURCES "${test_hip_sources}") + target_include_directories(test_operator BEFORE PRIVATE + ${CMAKE_CURRENT_SOURCE_DIR}/../../../3rdparty/composable_kernel/include) + endif() # Find required packages diff --git a/tests/cpp/operator/test_cast_mxfp4_transpose.cu b/tests/cpp/operator/test_cast_mxfp4_transpose.cu index a940c11cc1..27525fc184 100644 --- a/tests/cpp/operator/test_cast_mxfp4_transpose.cu +++ b/tests/cpp/operator/test_cast_mxfp4_transpose.cu @@ -463,11 +463,13 @@ TEST_P(FusedCastTransposeMXFP4TestSuite, TestFusedCastTransposeMXFP4) { // Forward activations auto OP = &identity; switch (Act_type) { + case ActivationType::Identity: OP = &identity; break; case ActivationType::GeLU: OP = &gelu; break; case ActivationType::SiLU: OP = &silu; break; case ActivationType::ReLU: OP = &relu; break; case ActivationType::QGeLU: OP = &qgelu; break; case ActivationType::SReLU: OP = &srelu; break; + default: GTEST_FAIL() << "Unsupported activation type"; break; } TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, diff --git a/tests/cpp/operator/test_cast_mxfp8.cu b/tests/cpp/operator/test_cast_mxfp8.cu index 738c25d7d4..1e72373e24 100644 --- a/tests/cpp/operator/test_cast_mxfp8.cu +++ b/tests/cpp/operator/test_cast_mxfp8.cu @@ -658,11 +658,13 @@ TEST_P(FusedCastMXFP8TestSuite, TestFusedCastMXFP8) { // Forward activations auto OP = &identity; switch (Act_type) { + case ActivationType::Identity: OP = &identity; break; case ActivationType::GeLU: OP = &gelu; break; case ActivationType::SiLU: OP = &silu; break; case ActivationType::ReLU: OP = &relu; break; case ActivationType::QGeLU: OP = &qgelu; break; case ActivationType::SReLU: OP = &srelu; break; + default: GTEST_FAIL() << "Unsupported activation type"; break; } TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, @@ -681,11 +683,13 @@ TEST_P(FusedCastMXFP8TestSuite, TestFusedCastMXFP8) { } else { auto OP = &identity; switch (Act_type) { + case ActivationType::Identity: OP = &identity; break; case ActivationType::GeLU: OP = &dgelu; break; case ActivationType::SiLU: OP = &dsilu; break; case ActivationType::ReLU: OP = &drelu; break; case ActivationType::QGeLU: OP = &dqgelu; break; case ActivationType::SReLU: OP = &dsrelu; break; + default: GTEST_FAIL() << "Unsupported activation type"; break; } TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, TRANSFORMER_ENGINE_TYPE_SWITCH_FP8_ONLY(output_type, OutputType, diff --git a/tests/cpp/operator/test_cast_nvfp4_transpose.cu b/tests/cpp/operator/test_cast_nvfp4_transpose.cu index 40c40d86f9..a92580de07 100644 --- a/tests/cpp/operator/test_cast_nvfp4_transpose.cu +++ b/tests/cpp/operator/test_cast_nvfp4_transpose.cu @@ -1453,11 +1453,13 @@ TEST_P(FusedCastTransposeNVFP4TestSuite, TestFusedCastTransposeNVFP4) { // Forward activations auto OP = &identity; switch (Act_type) { + case ActivationType::Identity: OP = &identity; break; case ActivationType::GeLU: OP = &gelu; break; case ActivationType::SiLU: OP = &silu; break; case ActivationType::ReLU: OP = &relu; break; case ActivationType::QGeLU: OP = &qgelu; break; case ActivationType::SReLU: OP = &srelu; break; + default: GTEST_FAIL() << "Unsupported activation type"; break; } TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, diff --git a/tests/cpp/operator/test_cublaslt_gemm.cu b/tests/cpp/operator/test_cublaslt_gemm.cu index aeb40d8967..85ab0247ca 100644 --- a/tests/cpp/operator/test_cublaslt_gemm.cu +++ b/tests/cpp/operator/test_cublaslt_gemm.cu @@ -515,6 +515,47 @@ std::pair getTestTolerances(const DType type, bool use_fp8, bool } +void checkMxFP8Support(const TestParams& params, const cudaDeviceProp& prop, bool &use_mxfp8, bool &use_hipkittens_mxfp8) { + use_mxfp8 = (params.scaling_mode == NVTEScalingMode::NVTE_MXFP8_1D_SCALING); + if (!use_mxfp8) { + use_hipkittens_mxfp8 = false; + return; + } +#ifdef __HIP_PLATFORM_AMD__ + if (!(prop.major == 9 && prop.minor >= 5) && !(prop.major >= 12)) { + GTEST_SKIP() << "MXFP8 requires gfx950 or newer"; + } +#endif + if (params.m % 16 || params.n % 16) { + GTEST_SKIP() << "MXFP8 requires M & N to be multiples of 16"; + } + +#ifdef __HIP_PLATFORM_AMD__ + const size_t required_k_multiple = (prop.major == 12 && prop.minor == 5) ? 32 : 128; +#else + const size_t required_k_multiple = 128; +#endif + if (params.k % required_k_multiple) { + GTEST_SKIP() << "MXFP8 requires K to be a multiple of " << required_k_multiple; + } + + use_hipkittens_mxfp8 = !params.force_hipblaslt; + if (!use_hipkittens_mxfp8) { + return; + } +#ifdef __HIP_PLATFORM_AMD__ + if (!(prop.major == 9 && (prop.minor == 4 || prop.minor == 5))) { + GTEST_SKIP() << "HipKittens requires gfx942 or gfx950"; + } + if (params.m % 256 || params.n % 256 || params.k < 256) { + GTEST_SKIP() << "HipKittens requires M and N 256-aligned, K >= 256"; + } +#else + GTEST_SKIP() << "HipKittens requires ROCm"; +#endif +} + + template void performTest(const TestParams& params) { DType atype = TypeInfo::dtype; @@ -524,30 +565,19 @@ void performTest(const TestParams& params) { DType dtype = TypeInfo::dtype; const bool has_fp8 = isFp8Type(atype) || isFp8Type(btype); - const bool use_mxfp8 = params.scaling_mode == NVTEScalingMode::NVTE_MXFP8_1D_SCALING; - const bool use_hipkittens_mxfp8 = use_mxfp8 && !params.force_hipblaslt; cudaDeviceProp prop; (void)cudaGetDeviceProperties(&prop, 0); + bool use_mxfp8 = false; + bool use_hipkittens_mxfp8 = false; + checkMxFP8Support(params, prop, use_mxfp8, use_hipkittens_mxfp8); + if (use_mxfp8) { if (!has_fp8) { GTEST_SKIP() << "MXFP8 scaling mode requires Float8 types"; } - if (params.m % 16 || params.n % 16) { - GTEST_SKIP() << "MXFP8 requires M & N to be multiples of 16"; - } - size_t required_k_multiple = 128; - #ifdef __HIP_PLATFORM_AMD__ - required_k_multiple = (prop.major == 12 && prop.minor == 5) ? 32 : 128; - #endif - if (params.k % required_k_multiple) { - GTEST_SKIP() << "MXFP8 requires K to be a multiple of " << required_k_multiple; - } - if (use_hipkittens_mxfp8 && (params.m % 256 || params.n % 256 || params.k < 256)) { - GTEST_SKIP() << "HipKittens requires M and N 256-aligned, K >= 256"; - } } #ifdef __HIP_PLATFORM_AMD__ @@ -580,10 +610,6 @@ void performTest(const TestParams& params) { if (!fp8_supported) { GTEST_SKIP() << "FP8 is not supported in current config"; } - const bool mxfp8_supported = (prop.major == 9 && prop.minor >= 5) || prop.major >= 12; - if (use_mxfp8 && !mxfp8_supported) { - GTEST_SKIP() << "MXFP8 is not supported in current config"; - } if (!use_hipkittens_mxfp8 && params.use_bias) { GTEST_SKIP() << "MXFP8 GEMM with bias is not supported by hipBLASLt"; } @@ -759,28 +785,11 @@ void performDqTest(const TestParams ¶ms) { cudaDeviceProp prop; (void)cudaGetDeviceProperties(&prop, 0); - if (params.m % 16 || params.n % 16) { - GTEST_SKIP() << "MXFP8 requires M & N to be multiples of 16"; - } - size_t required_k_multiple = 128; -#ifdef __HIP_PLATFORM_AMD__ - required_k_multiple = (prop.major == 12 && prop.minor == 5) ? 32 : 128; -#endif - if (params.k % required_k_multiple) { - GTEST_SKIP() << "MXFP8 requires K to be a multiple of " << required_k_multiple; - } - - bool mxfp8_supported = (prop.major == 9 && prop.minor >= 5) || prop.major >= 12; - const bool use_hipkittens_mxfp8 = !params.force_hipblaslt; - if (!mxfp8_supported) { - GTEST_SKIP() << "MXFP8 is not supported in current config"; - } + bool _unused = false; + checkMxFP8Support(params, prop, _unused, _unused); if (params.use_bias || params.use_gelu) { GTEST_SKIP() << "DqGEMMTestSuite does not yet have reference for bias/gelu epilogues"; } - if (use_hipkittens_mxfp8 && (params.m % 256 || params.n % 256 || params.k % 128 || params.k < 256)) { - GTEST_SKIP() << "HipKittens requires M and N 256-aligned, K >= 256"; - } // hipBLASLt on gfx950 produces incorrect results for certain MXFP8 // GEMMs with non-TN layouts. @@ -940,7 +949,7 @@ INSTANTIATE_TEST_SUITE_P(OperatorTest, GEMMTestSuite, ::testing::Values(false, true), //use_gelu ::testing::ValuesIn(kLayouts), //transa,transb ::testing::Values(false), //use mxfp8 - ::testing::Values(false)), //force hipblaslt + ::testing::Values(true)), //force hipblaslt GEMMTestName); INSTANTIATE_TEST_SUITE_P(OperatorTestFP8, FP8GEMMTestSuite, @@ -949,7 +958,7 @@ INSTANTIATE_TEST_SUITE_P(OperatorTestFP8, FP8GEMMTestSuite, ::testing::Values(false, true), //use_gelu ::testing::ValuesIn(kLayouts), //transa,transb ::testing::Values(false), //use mxfp8 - ::testing::Values(false)), //force hipblaslt + ::testing::Values(true)), //force hipblaslt GEMMTestName); INSTANTIATE_TEST_SUITE_P(OperatorTestMXFP8, FP8GEMMTestSuite, diff --git a/tests/pytorch/attention/test_attention.py b/tests/pytorch/attention/test_attention.py index 1bd7ea0272..bcd2c48cf9 100644 --- a/tests/pytorch/attention/test_attention.py +++ b/tests/pytorch/attention/test_attention.py @@ -310,9 +310,7 @@ def test_dot_product_attention( # Skip if only unfused backend is supported # Double-count the CK backend since we want to compare V2/V3 kernels - # Issue #16948 CK V2 is disabled for gfx1250 - has_ck_backend = ( IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends and - get_device_compute_capability() != (12, 5) ) + has_ck_backend = IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends if not has_ck_backend and ( len(fused_attn_backends) + flash_attn_supported + unfused_attn_supported ) < 2: @@ -1669,9 +1667,7 @@ def test_transformer_layer( # Skip if only unfused backend is supported # Double-count the CK backend since we want to compare V2/V3 kernels - # Issue #16948 CK V2 is disabled for gfx1250 - has_ck_backend = ( IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends and - get_device_compute_capability() != (12, 5) ) + has_ck_backend = IS_HIP_EXTENSION and FusedAttnBackend["CK"] in fused_attn_backends if not has_ck_backend and ( len(fused_attn_backends) + flash_attn_supported + unfused_attn_supported ) < 2: diff --git a/tests/pytorch/test_grouped_linear.py b/tests/pytorch/test_grouped_linear.py index 51c16769a3..2f76cf97e3 100644 --- a/tests/pytorch/test_grouped_linear.py +++ b/tests/pytorch/test_grouped_linear.py @@ -1241,6 +1241,7 @@ def _apply_grouped_bias_ref( return out +@pytest.mark.skipif(IS_HIP_EXTENSION, reason="GroupedTensor grouped GEMM is not supported on ROCm.") @pytest.mark.parametrize( "z, m, n, k", [ @@ -1255,8 +1256,6 @@ def _apply_grouped_bias_ref( @pytest.mark.parametrize("accumulate", [False, True]) @pytest.mark.parametrize("use_bias_scale", [False, True]) def test_grouped_gemm_grouped_tensor(z, m, n, k, case, layout, accumulate, use_bias_scale) -> None: - if IS_HIP_EXTENSION: - pytest.skip("GroupedTensor grouped GEMM needs nvte_grouped_gemm, unsupported on ROCm.") if torch.cuda.get_device_capability() < (9, 0): pytest.skip("Grouped GEMM requires Hopper (SM90) or newer.") if torch.cuda.get_device_capability() < (10, 0): @@ -1405,6 +1404,7 @@ def test_grouped_gemm_grouped_tensor(z, m, n, k, case, layout, accumulate, use_b torch.testing.assert_close(o, o_ref, **tols) +@pytest.mark.skipif(IS_HIP_EXTENSION, reason="GroupedTensor grouped GEMM is not supported on ROCm.") @pytest.mark.parametrize("layout", ["TN", "NN", "NT"]) @pytest.mark.parametrize("accumulate", [False, True]) @pytest.mark.parametrize("quant_type", ["bf16", "mxfp8"]) @@ -1566,6 +1566,7 @@ def _per_tensor_quantize_mxfp8( return [quantizer(t) for t in tensors] +@pytest.mark.skipif(IS_HIP_EXTENSION, reason="GroupedTensor grouped GEMM is not supported on ROCm.") @pytest.mark.parametrize( "shape", [ diff --git a/transformer_engine/common/cast/mxfp4/quantize_mxfp4.cuh b/transformer_engine/common/cast/mxfp4/quantize_mxfp4.cuh index 42e6276206..6d16473b69 100644 --- a/transformer_engine/common/cast/mxfp4/quantize_mxfp4.cuh +++ b/transformer_engine/common/cast/mxfp4/quantize_mxfp4.cuh @@ -34,9 +34,9 @@ void quantize(const Tensor &input, const Tensor *act_input, const Tensor *noop, int device; NVTE_CHECK_CUDA(hipGetDevice(&device)); NVTE_CHECK_CUDA(hipGetDeviceProperties(&prop, device)); - NVTE_CHECK(prop.major == 9 && prop.minor == 5, - "MXFP4 quantization requires gfx950 (detected gfx", - prop.major, prop.minor * 10, ")"); + NVTE_CHECK((prop.major == 9 && prop.minor == 5) || prop.major >= 12, + "MXFP4 quantization requires gfx950 and newer (detected gfx", + prop.major, prop.minor, "x)"); } int M = static_cast(input.flat_first_dim()); diff --git a/transformer_engine/pytorch/quantization.py b/transformer_engine/pytorch/quantization.py index 666b38f085..746cadf519 100644 --- a/transformer_engine/pytorch/quantization.py +++ b/transformer_engine/pytorch/quantization.py @@ -57,6 +57,7 @@ _MXFP8_SUPPORT: Optional[Tuple[bool, str]] = None _NVFP4_SUPPORT: Optional[Tuple[bool, str]] = None _FP8_BLOCK_SCALING_SUPPORT: Optional[Tuple[bool, str]] = None +_MXFP4_SUPPORT: Optional[Tuple[bool, str]] = None @dataclasses.dataclass(frozen=True) @@ -177,7 +178,7 @@ def _compute_mxfp8_support() -> Tuple[bool, str]: if os.getenv("NVTE_ROCM_ENABLE_MXFP8", "0") == "0": return False, "MXFP8 support is not enabled." gpu_arch = get_device_compute_capability() - if gpu_arch in ((9, 5), (12, 5)): + if gpu_arch in ((9, 5),): #TODO: enable for gfx1250 when GEMM is available return True, "" return False, "Device arch gfx95x or newer is required for MXFP8 execution." if get_device_compute_capability() >= (12, 0): @@ -203,7 +204,7 @@ def _compute_fp8_block_scaling_support() -> Tuple[bool, str]: """Return if fp8 block scaling support is available""" if IS_HIP_EXTENSION: gpu_arch = get_device_compute_capability() - if gpu_arch in ((9, 4), (9, 5)): # TODO: enabled for gfx1250 when ready + if gpu_arch in ((9, 4), (9, 5)): # TODO: enable for gfx1250 when GEMM is available return True, "" return False, "Device arch gfx94x or newer is required for FP8 block scaling execution." if get_device_compute_capability() >= (9, 0) and float(torch.version.cuda) >= 12.9: @@ -213,6 +214,15 @@ def _compute_fp8_block_scaling_support() -> Tuple[bool, str]: "FP8 block scaled GEMM requires compute capability 9.0 or higher and CUDA >= 12.9.", ) +def _compute_mxfp4_support() -> Tuple[bool, str]: + """Return if mxfp4 support is available""" + if IS_HIP_EXTENSION: + gpu_arch = get_device_compute_capability() + if gpu_arch in ((9, 5),): # TODO: enable for gfx1250 when GEMM is available + return True, "" + return False, "Device arch gfx95x or newer is required for MXFP4 execution." + return False, "Only ROCm supports MXFP4" + @torch.compiler.assume_constant_result def check_fp8_support() -> Tuple[bool, str]: """Return if fp8 support is available.""" @@ -249,15 +259,13 @@ def check_fp8_block_scaling_support() -> Tuple[bool, str]: return _FP8_BLOCK_SCALING_SUPPORT -@functools.lru_cache(maxsize=None) +@torch.compiler.assume_constant_result def check_mxfp4_support() -> Tuple[bool, str]: - """Return if mxfp4 support is available""" - if IS_HIP_EXTENSION: - gpu_arch = get_device_compute_capability() - if gpu_arch == (9, 5): # TODO: enable for gfx1250 when ready - return True, "" - return False, "Device arch gfx95x or newer is required for MXFP4 execution." - return False, "Only ROCm gfx950 supports MXFP4" + """Return if MXFP4 support is available.""" + global _MXFP4_SUPPORT + if _MXFP4_SUPPORT is None: + _MXFP4_SUPPORT = _compute_mxfp4_support() + return _MXFP4_SUPPORT def check_recipe_support(recipe: Recipe) -> None: