Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
56 commits
Select commit Hold shift + click to select a range
d89f90f
Updated QoLA (to port CK receipt patch) and TE manifest
Micky774 Apr 28, 2026
312212f
Updated manifest
Micky774 Apr 28, 2026
89f6983
Corrected AITER mha args validation against pinned commit
Micky774 Apr 28, 2026
c417e40
Updated cmake w/ dubious ownership protection
Micky774 Apr 28, 2026
90e3a5a
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Apr 29, 2026
c7ecaf7
Corrected logging
Micky774 Apr 29, 2026
da6e9a6
Updated qola to build aiter w/ new third_party spec
Micky774 Apr 29, 2026
e7ed124
Added guards against AITER known buggy implementations
Micky774 Apr 29, 2026
6241f99
Updated build
Micky774 May 12, 2026
82ebdb6
Drop guard for corrected bug
Micky774 May 19, 2026
68a2eaf
Merge branch 'dev' into zain/qola/aiter-update
Micky774 May 28, 2026
f9ab59c
Update AITER commit, adopt new API
Micky774 May 28, 2026
fec80b8
Added guard against graph-unsafe CK V2 kernels
Micky774 Jun 5, 2026
6eb3b23
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jun 5, 2026
8fd79f9
Updated with AOT memory handling (dynamic on host)
Micky774 Jun 16, 2026
ab68ca9
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jun 16, 2026
f521ef0
Added entry-point for workspace calc in QoLA
Micky774 Jun 18, 2026
e5ebe9c
Updated AITER commit to one w/ CK graph fix cherry-pick
Micky774 Jun 19, 2026
c08548e
Updated determnistic arg plumbing
Micky774 Jun 19, 2026
292870e
Updated CK JIT
Micky774 Jun 22, 2026
3476f53
Updated to gfx-1250 qola branch
Micky774 Jun 22, 2026
79c400d
Added pinned host memory pool allocation scheme
Micky774 Jun 23, 2026
3a98eda
Update CK JIT and prebuiuld cache list
ipanfilo Jun 24, 2026
bb0797a
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jun 24, 2026
bf2ae3c
Made error more descriptive
Micky774 Jun 26, 2026
8a91a58
Removed outdated graph safety guards
Micky774 Jun 26, 2026
679283a
Marked unused param in get attn backend
Micky774 Jun 26, 2026
948c077
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jun 26, 2026
eaa860d
Updated qola commit
Micky774 Jun 26, 2026
48e73b2
Merge branch 'dev' into zain/qola/aiter-update
ipanfilo Jun 29, 2026
9172b4b
PR feedback
Micky774 Jun 29, 2026
485ba06
Fix bwd workspace underallocation for d_v \!= d_qk configs
Micky774 Jun 29, 2026
fe94903
Fixed flaky determinism failures for padded configs
Micky774 Jul 1, 2026
eaf3cce
Corrected comment
Micky774 Jul 1, 2026
2411723
Moved workspace zeroing from TE to QoLA CK patch
Micky774 Jul 1, 2026
9860a92
Host memory zero on TE side
Micky774 Jul 6, 2026
6a23ce3
Reduce bound split size
Micky774 Jul 6, 2026
52b39d5
Corrected memory allocation upper bound
Micky774 Jul 7, 2026
9c8fb58
Added ignored comment to ignored cuda_graph arg
Micky774 Jul 8, 2026
729dfe2
Updated to leverage CK patch for device-side workspace prep
Micky774 Jul 9, 2026
270c358
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jul 9, 2026
49e65f5
Zeroed out dk_expanded buffer, and enforced stricter aiter zeroing
Micky774 Jul 9, 2026
db13f73
Added back dropped seqlen correction
Micky774 Jul 10, 2026
0f7f531
Remove unnecessary zeroing
Micky774 Jul 10, 2026
4fe3a7b
Update qola patch to CK dq_acc zeroing
Micky774 Jul 10, 2026
cf46085
Lower bound sizing fix
Micky774 Jul 13, 2026
9c2cccb
Corrected JAX group mode sizing issue
Micky774 Jul 17, 2026
e3df798
Corrected zeroing via AITER patch
Micky774 Jul 17, 2026
314777b
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jul 17, 2026
3346d3c
Corrected patch
Micky774 Jul 17, 2026
c7b8e1d
Kept default nullptr for nvidia builds
Micky774 Jul 22, 2026
e2f8832
Updated prebuilt list for gfx9 kernels
Micky774 Jul 22, 2026
0e40940
Updated ck-jit to main
Micky774 Jul 22, 2026
d65de00
Merge branch 'dev' into zain/qola/aiter-update
Micky774 Jul 22, 2026
efc7fbd
Added comment
Micky774 Jul 23, 2026
9b5b06b
Moved to qola main after gfx1250 and determinism merges
Micky774 Jul 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion 3rdparty/ck_jit
Comment thread
ipanfilo marked this conversation as resolved.
1,142 changes: 601 additions & 541 deletions ci/ck_jit_prebuild.txt
Comment thread
ipanfilo marked this conversation as resolved.

Large diffs are not rendered by default.

70 changes: 34 additions & 36 deletions transformer_engine/common/ck_fused_attn/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -65,48 +65,45 @@ if("${AITER_SHA}" STREQUAL "")
message(FATAL_ERROR
"Failed to parse 'aiter_commit = \"...\"' line in ${__QOLA_MANIFEST}.")
endif()

if(Python_EXECUTABLE)
if (NOT __SKIP_AITER_CHECKOUT)
execute_process(
COMMAND sh -c
"PYTHONPATH=\"${__QOLA_DIR}:$PYTHONPATH\" '${Python_EXECUTABLE}' -m qola.cli checkout \
--manifest '${__QOLA_MANIFEST}' \
--aiter-root '${__AITER_SOURCE_DIR}'"
RESULT_VARIABLE AITER_CHECKOUT_RESULT
OUTPUT_VARIABLE AITER_CHECKOUT_OUTPUT
ERROR_VARIABLE AITER_CHECKOUT_ERROR
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_STRIP_TRAILING_WHITESPACE
)
if(NOT AITER_CHECKOUT_RESULT EQUAL 0)
message(FATAL_ERROR
"Failed to sync AITER source tree at ${__AITER_SOURCE_DIR} to "
"manifest-pinned commit ${AITER_SHA}.\n"
"${AITER_CHECKOUT_OUTPUT}\n${AITER_CHECKOUT_ERROR}")
endif()
message(STATUS "[AITER] Synced ${__AITER_SOURCE_DIR} to ${AITER_SHA}")
endif()

if (NOT __SKIP_AITER_CHECKOUT)
execute_process(
COMMAND ${Python_EXECUTABLE} ${CMAKE_CURRENT_LIST_DIR}/check_aiter_mha_args.py
--mode both
--te-dir "${CMAKE_CURRENT_LIST_DIR}/../../.."
--aiter-root "${__AITER_SOURCE_DIR}"
RESULT_VARIABLE AITER_ARG_CHECK_RESULT
OUTPUT_VARIABLE AITER_ARG_CHECK_OUTPUT
ERROR_VARIABLE AITER_ARG_CHECK_ERROR
COMMAND sh -c
"PYTHONPATH=\"${__QOLA_DIR}:$PYTHONPATH\" '${Python_EXECUTABLE}' -m qola.cli checkout \
--manifest '${__QOLA_MANIFEST}' \
--aiter-root '${__AITER_SOURCE_DIR}'"
RESULT_VARIABLE AITER_CHECKOUT_RESULT
OUTPUT_VARIABLE AITER_CHECKOUT_OUTPUT
ERROR_VARIABLE AITER_CHECKOUT_ERROR
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_STRIP_TRAILING_WHITESPACE
)

if(NOT AITER_ARG_CHECK_RESULT EQUAL 0)
if(NOT AITER_CHECKOUT_RESULT EQUAL 0)
message(FATAL_ERROR
"AITER API validation failed in check_aiter_mha_args.py.\n"
"${AITER_ARG_CHECK_OUTPUT}\n${AITER_ARG_CHECK_ERROR}")
"Failed to sync AITER source tree at ${__AITER_SOURCE_DIR} to "
"manifest-pinned commit ${AITER_SHA}.\n"
"${AITER_CHECKOUT_OUTPUT}\n${AITER_CHECKOUT_ERROR}")
endif()
message(STATUS "AITER API validation passed via check_aiter_mha_args.py")
message(STATUS "[AITER] Synced ${__AITER_SOURCE_DIR} to ${AITER_SHA}")
endif()

execute_process(
COMMAND ${Python_EXECUTABLE} ${CMAKE_CURRENT_LIST_DIR}/check_aiter_mha_args.py
--mode both
--te-dir "${CMAKE_CURRENT_LIST_DIR}/../../.."
--aiter-root "${__AITER_SOURCE_DIR}"
RESULT_VARIABLE AITER_ARG_CHECK_RESULT
OUTPUT_VARIABLE AITER_ARG_CHECK_OUTPUT
ERROR_VARIABLE AITER_ARG_CHECK_ERROR
OUTPUT_STRIP_TRAILING_WHITESPACE
ERROR_STRIP_TRAILING_WHITESPACE
)

if(NOT AITER_ARG_CHECK_RESULT EQUAL 0)
message(FATAL_ERROR
"AITER API validation failed in check_aiter_mha_args.py.\n"
"${AITER_ARG_CHECK_OUTPUT}\n${AITER_ARG_CHECK_ERROR}")
endif()
message(STATUS "AITER API validation passed via check_aiter_mha_args.py")

# Sanity-check the resolved include directories now that `qola checkout` has
# materialized the AITER tree.
Expand Down Expand Up @@ -224,7 +221,8 @@ endforeach()
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})
-DCK_TILE_FLOAT_TO_BFLOAT16_DEFAULT=${CK_FUSED_ATTN_FLOAT_TO_BFLOAT16_DEFAULT}
-DENABLE_CK=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
Original file line number Diff line number Diff line change
Expand Up @@ -113,7 +113,6 @@ struct CkAttnBwdArgs : CKAttnCommonArgs {
// dQ
void* dq_ptr = nullptr;
uint64_t stride_b_dq = 0, stride_h_dq = 0, stride_s_dq = 0;
void* dq_acc_ptr = nullptr;

// dK / dV expanded (MQA/GQA reduction inputs; null when h==hg)
void* dk_expanded_ptr = nullptr;
Expand All @@ -134,6 +133,13 @@ struct CkAttnBwdArgs : CKAttnCommonArgs {
// Workspace shared with forward LSE
void* lse_workspace_ptr = nullptr;

// AOT scratch for AITER's internal bwd allocations (launcher metadata + dq_acc
// accumulator). Carved from the caller's workspace and handed to aiter through
// the workspace_alloc callback; ck_attn_bwd_workspace_size() reports the bytes
// to reserve. aiter_workspace_bytes bounds the bump allocator.
void* aiter_workspace_ptr = nullptr;
size_t aiter_workspace_bytes = 0;

// V3 ASM kernel selection
bool deterministic = false;
bool uses_bwd_v3 = false;
Expand All @@ -143,6 +149,10 @@ struct CkAttnBwdArgs : CKAttnCommonArgs {
hipError_t ck_attn_fwd(const CKAttnFwdArgs& args, hipStream_t stream);
hipError_t ck_attn_bwd(const CkAttnBwdArgs& args, hipStream_t stream);

// Bytes of AOT device scratch ck_attn_bwd needs for AITER's internal bwd
// workspace (launcher metadata + dq_acc), covering both the v2 (CK launcher) and
// v3 (asm) dispatch paths. Pure host-side computation; no kernel launch.
size_t ck_attn_bwd_workspace_size(const CkAttnBwdArgs& args);
// Probe whether AITER's v3 (asm) path will run for the given config, without
// launching a kernel (backed by AITER's v3_api_check dry-run). Returns true iff
// the v3 path is selected; false means the CK v2 path (or no support) would run.
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
[qola]
aiter_commit = "d32b0cb62ecc32bb1a858e8437d58eb9b3856af6" # pinned AITER submodule commit
aiter_commit = "e95f7d4e3924e848c1daf4df449e27822b332a09" # pinned AITER submodule commit
namespace = "te"
rocm_versions = ["7.2"]

Expand All @@ -9,9 +9,11 @@ architectures = ["gfx950", "gfx942"]
[[modules]]
name = "libmha_fwd"
mode = "cpp_itfs"
receipt = 700
drop_srcs = ["mha_fwd_split.cu", "mha_fwd_batch_prefill.cu"]
drop_directions = ["fwd_splitkv", "batch_prefill"]

[[modules]]
name = "libmha_bwd"
mode = "cpp_itfs"
receipt = 700
Loading
Loading