Skip to content
Merged
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
7 changes: 7 additions & 0 deletions ck_api_rewrite.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,6 +110,11 @@ def _detect_bwd_scheme_and_lambda():
r'&fmha_bwd_dq_dk_dv_dq_prepare_ws_host_\s*<[^>]+>',
'reinterpret_cast<PrepareWorkspaceHostFunc>(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<PrepareWorkspaceDeviceFunc>(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_)',
Expand Down Expand Up @@ -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",
Expand Down
29 changes: 26 additions & 3 deletions ck_jit_runtime.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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
};

Expand Down Expand Up @@ -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_<T,Arch>(int batch) → size_t
// fmha_bwd_dq_dk_dv_dq_ws_device_upper_bound_<T,Arch>(...) → size_t
// fmha_bwd_dq_dk_dv_dq_prepare_ws_host_<T,Arch>(void*,...) → size_t
// fmha_bwd_dq_dk_dv_dq_prepare_ws_device_<T,Arch>(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)
{
Expand All @@ -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
Expand Down Expand Up @@ -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")))
Expand Down
3 changes: 2 additions & 1 deletion ck_post_build.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down