From 4000fc8fd8f5725dcc596753222c9040a64bb024 Mon Sep 17 00:00:00 2001 From: Meekail Zain Date: Thu, 9 Jul 2026 18:33:26 +0000 Subject: [PATCH] Updated to use new device workspace prep function --- ck_api_rewrite.py | 6 +++--- ck_jit_runtime.cpp | 22 +++++++++++----------- ck_post_build.py | 2 +- 3 files changed, 15 insertions(+), 15 deletions(-) diff --git a/ck_api_rewrite.py b/ck_api_rewrite.py index 57da3a9..2a883d4 100644 --- a/ck_api_rewrite.py +++ b/ck_api_rewrite.py @@ -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(ck_jit_bwd_get_prepare_ws_func(dq_))', + 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*\)', @@ -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", diff --git a/ck_jit_runtime.cpp b/ck_jit_runtime.cpp index 15ac943..1983c55 100644 --- a/ck_jit_runtime.cpp +++ b/ck_jit_runtime.cpp @@ -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 }; @@ -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_(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_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_device_(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) { @@ -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 @@ -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 diff --git a/ck_post_build.py b/ck_post_build.py index dc225b7..5af656e 100644 --- a/ck_post_build.py +++ b/ck_post_build.py @@ -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: