Skip to content

Commit 34cb16f

Browse files
authored
feat(ck-tile): multi-ABD GEMM TE to dispatcher bridge (#9305)
ISSUE ID: #8997 ## Motivation The CK Tile dispatcher could already generate and launch regular GEMM through the TileEngine → Dispatcher bridge, but it had no path for the multi-tensor **gemm_multi_abd** op. Multi-ABD is used when a GEMM needs to combine several A and B operands and fuse several D operands in the epilogue (`E = cde_op(a_op(As) @ b_op(Bs), {Ds})`), which is a real Old-TE capability with no dispatcher equivalent. This PR closes that gap so Python callers can drive multi_abd through the dispatcher at parity with the legacy Tile Engine version, without touching C++. It follows the divergent-ABI pattern established by the grouped bridge (#9000) because multi_abd needs **arrays** of A/B/D device pointers, not the single-pointer regular GEMM ABI. The capability set matches the Old-TE `gemm_multi_abd_instance_builder.py` exactly: `fp16`, `rcrr` layout, configurable A/B/D tensor counts, and the element-wise op set `{PassThrough, AddScale, MultiDMultiply, MultiDAdd}`. ## Test Plan - Run the CPU-only unit tests (no GPU required): `python3 -m pytest dispatcher/tests/test_multi_abd_bridge.py -v` - On-GPU numeric verification through the bridge launch path (gfx942 / MI300X), 512x512x512 fp16 rcrr, across the default 2/2/2 all-PassThrough config and non-PassThrough element-wise ops. - Confirm the CI and default config expansions yield the expected kernel counts. ## Test Result - CPU-only unit tests pass (10 passed). - Numeric verification (bridge launch path), 512x512x512 fp16 rcrr: - default 2/2/2 all-PassThrough: `max_rel = 2.9e-4` - CDE = MultiDAdd: `max_rel = 5.7e-4` - A-op = MultiDAdd: `max_rel = 4.1e-4` - all far below the fp16 tolerance (2e-2); 0 failed measurements. - CI config expands to 16 arch-valid kernels; `default_config.json` → 8896. - `standard` variant `expand_sweep` regression clean. - clang-format-18 (18.1.8) clean on `gemm_multi_abd_ctypes_lib.cpp`. - Serialized A/B perf-parity vs Old-TE (MI300X / gfx942, fp16 rcrr, 16 stems × 5 shapes = 80 rows, interleaved, fair 50/100/flush/rotating both sides): **at parity** — median gap -0.24%, mean -0.66%, 100% within ±15%, 87.5% within ±5% (range [-9.57%, +5.35%]). See the parity comment for details. --- **Related PRs / references (TileEngine → Dispatcher GEMM bridge series):** #8997 (regular GEMM fp16/bf16 all-layout), #9000 (grouped GEMM), #9028 (stream-K), #8887 (fp8/bf8/int8). This PR is a sibling in the same bridge effort tracked across those PRs. --------- Co-authored-by: Muhammed Emin Ozturk <3836908+ozturkosu@users.noreply.github.com>
1 parent af3ef9a commit 34cb16f

9 files changed

Lines changed: 2072 additions & 116 deletions

File tree

projects/composablekernel/dispatcher/bindings/ctypes/CMakeLists.txt

Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -276,6 +276,91 @@ else()
276276
)
277277
endif()
278278

279+
# =============================================================================
280+
# TileEngine -> Dispatcher Bridge ctypes Libraries (append-only shared block)
281+
# =============================================================================
282+
#
283+
# The bridge PRs (#9305 gemm_multi_abd, #9306 batched_gemm, #9328
284+
# batched_contraction) each add a registry-bypass ctypes .so target. To keep the
285+
# branches mutually conflict-proof, this block is BYTE-IDENTICAL on every branch
286+
# that touches this file. Each target is doubly guarded:
287+
# (1) on its generated kernel-header glob, AND
288+
# (2) on if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/<op>_ctypes_lib.cpp)
289+
# so a branch that does not ship a given <op>_ctypes_lib.cpp simply skips that
290+
# target at configure time (no error). The Python bridges build these .so files
291+
# at runtime via hipcc; these CMake targets exist for CMake-driven / CI builds.
292+
293+
# --- GEMM Multi-ABD (registry-bypass, array-pointer ABI) ---------------------
294+
if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/gemm_multi_abd_ctypes_lib.cpp)
295+
file(GLOB GEMM_MULTI_ABD_KERNEL_HEADERS
296+
"${CMAKE_BINARY_DIR}/generated_kernels/gemm_*_multiabd_*.hpp")
297+
if(GEMM_MULTI_ABD_KERNEL_HEADERS)
298+
list(SORT GEMM_MULTI_ABD_KERNEL_HEADERS)
299+
list(GET GEMM_MULTI_ABD_KERNEL_HEADERS 0 GEMM_MULTI_ABD_KERNEL_HEADER)
300+
message(STATUS "Found GEMM Multi-ABD kernel for ctypes lib: ${GEMM_MULTI_ABD_KERNEL_HEADER}")
301+
add_ctypes_library(dispatcher_gemm_multi_abd_lib
302+
gemm_multi_abd_ctypes_lib.cpp
303+
KERNEL_HEADER ${GEMM_MULTI_ABD_KERNEL_HEADER}
304+
)
305+
target_compile_definitions(dispatcher_gemm_multi_abd_lib PRIVATE
306+
CK_TILE_SINGLE_KERNEL_INCLUDE)
307+
else()
308+
message(STATUS
309+
"No GEMM Multi-ABD kernel found for ctypes lib - skipping dispatcher_gemm_multi_abd_lib "
310+
"(built at runtime by the Python bridge once a kernel header exists)")
311+
endif()
312+
endif()
313+
314+
# --- Batched GEMM (registry-bypass, batch_count + per-batch strides ABI) ------
315+
if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/batched_gemm_ctypes_lib.cpp)
316+
file(GLOB BATCHED_GEMM_KERNEL_HEADERS "${CMAKE_BINARY_DIR}/generated_kernels/gemm_*_batched.hpp")
317+
if(BATCHED_GEMM_KERNEL_HEADERS)
318+
list(SORT BATCHED_GEMM_KERNEL_HEADERS)
319+
list(GET BATCHED_GEMM_KERNEL_HEADERS 0 BATCHED_GEMM_KERNEL_HEADER)
320+
message(STATUS "Found Batched GEMM kernel for ctypes lib: ${BATCHED_GEMM_KERNEL_HEADER}")
321+
add_library(dispatcher_batched_gemm_lib SHARED batched_gemm_ctypes_lib.cpp)
322+
target_include_directories(dispatcher_batched_gemm_lib PRIVATE
323+
${PROJECT_SOURCE_DIR}/include
324+
${PROJECT_SOURCE_DIR}/dispatcher/include
325+
)
326+
target_link_libraries(dispatcher_batched_gemm_lib PRIVATE hip::device)
327+
target_compile_options(dispatcher_batched_gemm_lib PRIVATE
328+
-include ${BATCHED_GEMM_KERNEL_HEADER}
329+
)
330+
target_compile_definitions(dispatcher_batched_gemm_lib PRIVATE CK_TILE_SINGLE_KERNEL_INCLUDE)
331+
set_target_properties(dispatcher_batched_gemm_lib PROPERTIES
332+
POSITION_INDEPENDENT_CODE ON
333+
CXX_STANDARD 17
334+
)
335+
else()
336+
message(STATUS "No Batched GEMM kernel found for ctypes lib - skipping dispatcher_batched_gemm_lib")
337+
endif()
338+
endif()
339+
340+
# --- Batched-Contraction (registry-bypass, force-included kernel) ------------
341+
if(EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/batched_contraction_ctypes_lib.cpp)
342+
file(GLOB BC_KERNEL_HEADERS "${CMAKE_BINARY_DIR}/generated_kernels/batched_contraction_*.hpp")
343+
if(BC_KERNEL_HEADERS)
344+
list(SORT BC_KERNEL_HEADERS)
345+
list(GET BC_KERNEL_HEADERS 0 BC_KERNEL_HEADER)
346+
message(STATUS "Found batched-contraction kernel for ctypes lib: ${BC_KERNEL_HEADER}")
347+
add_ctypes_library(dispatcher_batched_contraction_lib
348+
batched_contraction_ctypes_lib.cpp
349+
KERNEL_HEADER ${BC_KERNEL_HEADER}
350+
)
351+
# The generated header only exports SelectedKernel/KERNEL_NAME/CONTRACTION_KEY_*
352+
# under this define, so the force-include path requires it.
353+
target_compile_definitions(dispatcher_batched_contraction_lib PRIVATE
354+
CK_TILE_SINGLE_KERNEL_INCLUDE
355+
)
356+
else()
357+
message(STATUS
358+
"No batched-contraction kernel found for ctypes lib - skipping "
359+
"dispatcher_batched_contraction_lib (built at runtime by the Python "
360+
"bridge once a kernel header exists)")
361+
endif()
362+
endif()
363+
279364
# =============================================================================
280365
# GPU Helper Executable
281366
# =============================================================================

0 commit comments

Comments
 (0)