From 185540077f375a5d14c617aab4c3613b8b59eea6 Mon Sep 17 00:00:00 2001 From: Ilya Panfilov Date: Tue, 21 Jul 2026 23:22:54 -0400 Subject: [PATCH] Add QoLA generated CK WS prepare device func support --- ck_api_rewrite.py | 7 +++++++ ck_jit_runtime.cpp | 29 ++++++++++++++++++++++++++--- ck_post_build.py | 3 ++- 3 files changed, 35 insertions(+), 4 deletions(-) diff --git a/ck_api_rewrite.py b/ck_api_rewrite.py index 57da3a9..23d4fc6 100644 --- a/ck_api_rewrite.py +++ b/ck_api_rewrite.py @@ -110,6 +110,11 @@ def _detect_bwd_scheme_and_lambda(): r'&fmha_bwd_dq_dk_dv_dq_prepare_ws_host_\s*<[^>]+>', 'reinterpret_cast(ck_jit_bwd_get_prepare_ws_func(dq_))', body) + # QoLA patched CK uses device workspace prepare function. + body = _re.sub( + r'&fmha_bwd_dq_dk_dv_dq_prepare_ws_device_\s*<[^>]+>', + 'reinterpret_cast(ck_jit_bwd_get_prepare_ws_device_func(dq_))', + body) body = _re.sub( r'fmha_bwd_dq_dk_dv_needs_zero_dq_acc_\s*<[^>]+>\s*\(\s*\)', 'ck_jit_bwd_needs_zero_dq_acc(dq_)', @@ -467,6 +472,8 @@ def rewrite_api_file(src_path, dst_path, api_kind): "size_t ck_jit_bwd_dq_ws_device_upper_bound(const char*, ck_tile::index_t,\n" " ck_tile::index_t, ck_tile::index_t, ck_tile::index_t, ck_tile::index_t);\n" "void* ck_jit_bwd_get_prepare_ws_func(const char*);\n" + "// QoLA patched CK uses device workspace prepare function:\n" + "void* ck_jit_bwd_get_prepare_ws_device_func(const char*);\n" ), "fwd_splitkv": f"float ck_jit_fwd_splitkv_call(const char*, const char*, {_sc}, fmha_fwd_splitkv_args);\n", "batch_prefill":f"float ck_jit_batch_prefill_call(const char*, {_sc}, fmha_batch_prefill_args);\n", diff --git a/ck_jit_runtime.cpp b/ck_jit_runtime.cpp index 15ac943..9f26500 100644 --- a/ck_jit_runtime.cpp +++ b/ck_jit_runtime.cpp @@ -128,6 +128,7 @@ struct BwdDqBlobState : BlobState { void* fn_ws_host_size = nullptr; // dq_ws_host_size_<> void* fn_ws_device_upper_bound = nullptr; // dq_ws_device_upper_bound_<> void* fn_prepare_ws_host = nullptr; // dq_prepare_ws_host_<> + void* fn_prepare_ws_device = nullptr; // dq_prepare_ws_device_<>, added by QoLA #endif }; @@ -648,15 +649,18 @@ static void resolve_bwd_dq_meta(const char* dq_dk_dv_blob, BwdDqBlobState& state // Workspace-based API helpers — CK commit 2c677e84 // "[CK_TILE] Use Unified Workspace for FMHA BWD" // -// Resolves three new symbols from the dq_dk_dv blob: +// Resolves the following new symbols from the dq_dk_dv blob: // fmha_bwd_dq_dk_dv_dq_ws_host_size_(int batch) → size_t // fmha_bwd_dq_dk_dv_dq_ws_device_upper_bound_(...) → size_t // fmha_bwd_dq_dk_dv_dq_prepare_ws_host_(void*,...) → size_t +// fmha_bwd_dq_dk_dv_dq_prepare_ws_device_(void*,...) → void (launches kernel) +// The last one is added to CK by QoLA patch // // ELF mangled-name prefixes (Itanium ABI, template function length prefix): // _Z34fmha_bwd_dq_dk_dv_dq_ws_host_size_I (34 chars) // _Z43fmha_bwd_dq_dk_dv_dq_ws_device_upper_bound_I (43 chars) // _Z37fmha_bwd_dq_dk_dv_dq_prepare_ws_host_I (37 chars) +// _Z37fmha_bwd_dq_dk_dv_dq_prepare_ws_device_I (39 chars) // --------------------------------------------------------------------------- static void resolve_bwd_dq_ws_meta(const char* dq_dk_dv_blob, BwdDqBlobState& state) { @@ -675,9 +679,15 @@ static void resolve_bwd_dq_ws_meta(const char* dq_dk_dv_blob, BwdDqBlobState& st state.fn_ws_device_upper_bound = find_sym_by_prefix( state.handle, state.so_path.c_str(), "_Z43fmha_bwd_dq_dk_dv_dq_ws_device_upper_bound_I"); - state.fn_prepare_ws_host = find_sym_by_prefix( + state.fn_prepare_ws_device = find_sym_by_prefix( state.handle, state.so_path.c_str(), - "_Z37fmha_bwd_dq_dk_dv_dq_prepare_ws_host_I"); + "_Z39fmha_bwd_dq_dk_dv_dq_prepare_ws_device_I"); + // QoLA patched CK does not use WS prepare host function, so resolve it conditonally + if (state.fn_prepare_ws_device == nullptr) { + state.fn_prepare_ws_host = find_sym_by_prefix( + state.handle, state.so_path.c_str(), + "_Z37fmha_bwd_dq_dk_dv_dq_prepare_ws_host_I"); + } }); } #endif // CK_JIT_BWD_WORKSPACE_V2 @@ -731,6 +741,19 @@ void* ck_jit_bwd_get_prepare_ws_func(const char* dq_dk_dv_blob) } return state->fn_prepare_ws_host; } + +__attribute__((visibility("hidden"))) +void* ck_jit_bwd_get_prepare_ws_device_func(const char* dq_dk_dv_blob) +{ + BwdDqBlobState* state = get_bwd_dq_dk_dv_state(dq_dk_dv_blob); + resolve_bwd_dq_ws_meta(dq_dk_dv_blob, *state); + if (!state->fn_prepare_ws_device) { + ::fprintf(stderr, "[CK-JIT] ERROR: dq_prepare_ws_device symbol not found in %s\n", + dq_dk_dv_blob); + return nullptr; + } + return state->fn_prepare_ws_device; +} #endif // CK_JIT_BWD_WORKSPACE_V2 __attribute__((visibility("hidden"))) diff --git a/ck_post_build.py b/ck_post_build.py index dc225b7..8a53568 100644 --- a/ck_post_build.py +++ b/ck_post_build.py @@ -193,7 +193,8 @@ def _build_runtime_cmd(hipcc, is_fwd, ck_include, ck_fmha_include, # "[CK_TILE] Use Unified Workspace for FMHA BWD". # When present, -DCK_JIT_BWD_WORKSPACE_V2=1 enables the new runtime helpers # (ck_jit_bwd_dq_ws_host_size, ck_jit_bwd_dq_ws_device_upper_bound, - # ck_jit_bwd_get_prepare_ws_func) in ck_jit_runtime.cpp. + # ck_jit_bwd_get_prepare_ws_func, ck_jit_bwd_get_prepare_ws_device_func) in ck_jit_runtime.cpp. + # ck_jit_bwd_get_prepare_ws_device_func - is not a part of CK but added by QoLA patches _fmha_bwd_hpp = os.path.join(ck_fmha_include, "fmha_bwd.hpp") if ck_fmha_include else "" if _fmha_bwd_hpp and os.path.isfile(_fmha_bwd_hpp): try: