Skip to content
Closed
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
6 changes: 3 additions & 3 deletions ck_api_rewrite.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,8 +107,8 @@ def _detect_bwd_scheme_and_lambda():
'ck_jit_bwd_dq_ws_device_upper_bound(dq_, ',
body)
body = _re.sub(
r'&fmha_bwd_dq_dk_dv_dq_prepare_ws_host_\s*<[^>]+>',
'reinterpret_cast<PrepareWorkspaceHostFunc>(ck_jit_bwd_get_prepare_ws_func(dq_))',
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*\)',
Expand Down Expand Up @@ -466,7 +466,7 @@ def rewrite_api_file(src_path, dst_path, api_kind):
"size_t ck_jit_bwd_dq_ws_host_size(const char*, ck_tile::index_t);\n"
"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"
"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
22 changes: 11 additions & 11 deletions ck_jit_runtime.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -127,7 +127,7 @@ struct BwdDqBlobState : BlobState {
std::once_flag ws_meta_init_flag;
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_<>
#endif
};

Expand Down Expand Up @@ -649,14 +649,14 @@ static void resolve_bwd_dq_meta(const char* dq_dk_dv_blob, BwdDqBlobState& state
// "[CK_TILE] Use Unified Workspace for FMHA BWD"
//
// Resolves three 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_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_device_<T,Arch>(void*,...) → void (launches kernel)
//
// 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)
// _Z39fmha_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 +675,9 @@ 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");
});
}
#endif // CK_JIT_BWD_WORKSPACE_V2
Expand Down Expand Up @@ -720,16 +720,16 @@ size_t ck_jit_bwd_dq_ws_device_upper_bound(const char* dq_dk_dv_blob,
}

__attribute__((visibility("hidden")))
void* ck_jit_bwd_get_prepare_ws_func(const char* dq_dk_dv_blob)
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_host) {
::fprintf(stderr, "[CK-JIT] ERROR: dq_prepare_ws_host symbol not found in %s\n",
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_host;
return state->fn_prepare_ws_device;
}
#endif // CK_JIT_BWD_WORKSPACE_V2

Expand Down
2 changes: 1 addition & 1 deletion ck_post_build.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,7 +193,7 @@ 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_device_func) in ck_jit_runtime.cpp.
_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