Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
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
2 changes: 1 addition & 1 deletion 3rdparty/ck_jit
2 changes: 1 addition & 1 deletion transformer_engine/common/ck_fused_attn/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 3 additions & 2 deletions transformer_engine/common/ck_fused_attn/qola_manifest.toml
Original file line number Diff line number Diff line change
@@ -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"]
Expand All @@ -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"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
Loading