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