From 18f445122af4f0ff09c7144c10c2d1a245a11388 Mon Sep 17 00:00:00 2001 From: Alex Magro Date: Sat, 8 Aug 2026 02:12:33 +0000 Subject: [PATCH] gfx942-only build bugfix and kittens refactor --- transformer_engine/common/CMakeLists.txt | 7 +- .../common/gemm/kittens/CMakeLists.txt | 90 ++++----- .../gemm/kittens/cdna3/blockwise_fp8_gemm.cpp | 25 +-- .../cdna3/blockwise_fp8_gemm_helper.cuh | 27 +-- .../gemm/kittens/cdna4/blockwise_fp8_gemm.cpp | 26 +-- .../cdna4/blockwise_fp8_gemm_helper.cuh | 28 +-- .../common/gemm/kittens/cdna4/mxfp8_gemm.cpp | 158 +++++++++------- .../common/gemm/kittens/cdna4/mxfp8_gemm.h | 37 ---- .../common/gemm/kittens/kittens_common.h | 178 +++++++++++++++--- .../gemm/kittens/kittens_kernel_common.cuh | 53 ++++++ transformer_engine/common/gemm/rocm_gemm.cu | 76 ++++---- 11 files changed, 392 insertions(+), 313 deletions(-) delete mode 100644 transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.h create mode 100644 transformer_engine/common/gemm/kittens/kittens_kernel_common.cuh diff --git a/transformer_engine/common/CMakeLists.txt b/transformer_engine/common/CMakeLists.txt index e646b72c3..c56175fcc 100644 --- a/transformer_engine/common/CMakeLists.txt +++ b/transformer_engine/common/CMakeLists.txt @@ -717,12 +717,7 @@ else() # USE_ROCM endif() if(USE_HIPKITTENS_GEMM) target_compile_definitions(transformer_engine PUBLIC USE_HIPKITTENS_GEMM) - if(KITTENS_HAVE_CDNA3) - target_compile_definitions(transformer_engine PUBLIC KITTENS_HAVE_CDNA3) - endif() - if(KITTENS_HAVE_CDNA4) - target_compile_definitions(transformer_engine PUBLIC KITTENS_HAVE_CDNA4) - endif() + target_compile_definitions(transformer_engine PUBLIC ${KITTENS_HAVE_DEFS}) list(APPEND transformer_engine_LINKER_LIBS kittens_gemm) endif() target_link_libraries(transformer_engine PUBLIC ${transformer_engine_LINKER_LIBS}) diff --git a/transformer_engine/common/gemm/kittens/CMakeLists.txt b/transformer_engine/common/gemm/kittens/CMakeLists.txt index f09a7c7ca..39c42b3a6 100644 --- a/transformer_engine/common/gemm/kittens/CMakeLists.txt +++ b/transformer_engine/common/gemm/kittens/CMakeLists.txt @@ -3,23 +3,26 @@ cmake_minimum_required(VERSION 3.21) -list(FIND CMAKE_HIP_ARCHITECTURES "gfx942" _gfx942_index) -list(FIND CMAKE_HIP_ARCHITECTURES "gfx950" _gfx950_index) +set(KITTENS_SUPPORTED_ARCHS gfx942 gfx950) + +set(_kittens_enabled_archs "") +foreach(_arch IN LISTS KITTENS_SUPPORTED_ARCHS) + if(_arch IN_LIST CMAKE_HIP_ARCHITECTURES) + list(APPEND _kittens_enabled_archs ${_arch}) + endif() +endforeach() include(CheckCXXCompilerFlag) check_cxx_compiler_flag("-std=c++20" HAS_CXX20) -if(_gfx942_index EQUAL -1 AND _gfx950_index EQUAL -1) - message(STATUS "HipKittens GEMM disabled (neither gfx942 nor gfx950 in CMAKE_HIP_ARCHITECTURES)") +if(NOT _kittens_enabled_archs) + message(STATUS "HipKittens GEMM disabled (none of ${KITTENS_SUPPORTED_ARCHS} in CMAKE_HIP_ARCHITECTURES)") set(USE_HIPKITTENS_GEMM OFF PARENT_SCOPE) elseif(NOT HAS_CXX20) message(WARNING "HipKittens GEMMs require C++20") set(USE_HIPKITTENS_GEMM OFF PARENT_SCOPE) else() - set(HIPKITTENS_CDNA3_INCLUDE_DIR - "${CMAKE_CURRENT_SOURCE_DIR}/../../../../3rdparty/hipkittens/include") - set(HIPKITTENS_CDNA4_INCLUDE_DIR - "${CMAKE_CURRENT_SOURCE_DIR}/../../../../3rdparty/hipkittens/include") + set(HIPKITTENS_INCLUDE_DIR "${CMAKE_CURRENT_SOURCE_DIR}/../../../../3rdparty/hipkittens/include") set(CMAKE_CXX_STANDARD 20) project(kittens_gemm LANGUAGES HIP CXX) @@ -28,53 +31,52 @@ else() include_directories("${ROCM_PATH}/include/hip") set(_kittens_arch_objs "") + set(_kittens_have_defs "") - if(NOT _gfx942_index EQUAL -1) - if(NOT EXISTS "${HIPKITTENS_CDNA3_INCLUDE_DIR}/kittens.cuh") - message(FATAL_ERROR - "Could not find HipKittens (CDNA3) headers at ${HIPKITTENS_CDNA3_INCLUDE_DIR}. " - "Try running 'git submodule update --init --recursive'.") + function(kittens_add_arch) + cmake_parse_arguments(A "" "TAG;GFX" "SOURCES;FLAGS" ${ARGN}) + if(NOT "${A_GFX}" IN_LIST _kittens_enabled_archs) + return() endif() - add_library(kittens_gemm_cdna3 OBJECT cdna3/blockwise_fp8_gemm.cpp) - set_source_files_properties(cdna3/blockwise_fp8_gemm.cpp PROPERTIES LANGUAGE HIP) - set_target_properties(kittens_gemm_cdna3 PROPERTIES - HIP_ARCHITECTURES "gfx942" POSITION_INDEPENDENT_CODE ON) - target_include_directories(kittens_gemm_cdna3 PRIVATE - ${HIP_INCLUDE_DIRS} "${HIPKITTENS_CDNA3_INCLUDE_DIR}") - target_compile_options(kittens_gemm_cdna3 PRIVATE - -DKITTENS_CDNA3 -fno-gpu-rdc -O3) - target_link_libraries(kittens_gemm_cdna3 PRIVATE hip::host hip::device) - list(APPEND _kittens_arch_objs $) - set(KITTENS_HAVE_CDNA3 ON PARENT_SCOPE) - endif() - - if(NOT _gfx950_index EQUAL -1) - if(NOT EXISTS "${HIPKITTENS_CDNA4_INCLUDE_DIR}/kittens.cuh") + if(NOT EXISTS "${HIPKITTENS_INCLUDE_DIR}/kittens.cuh") message(FATAL_ERROR - "Could not find HipKittens (CDNA4) headers at ${HIPKITTENS_CDNA4_INCLUDE_DIR}. " + "Could not find HipKittens headers at ${HIPKITTENS_INCLUDE_DIR}. " "Try running 'git submodule update --init --recursive'.") endif() - add_library(kittens_gemm_cdna4 OBJECT cdna4/blockwise_fp8_gemm.cpp cdna4/mxfp8_gemm.cpp) - set_source_files_properties(cdna4/mxfp8_gemm.cpp - PROPERTIES LANGUAGE HIP COMPILE_FLAGS "-gline-tables-only") - set_source_files_properties(cdna4/blockwise_fp8_gemm.cpp - PROPERTIES LANGUAGE HIP COMPILE_FLAGS "-gline-tables-only -ffast-math") - set_target_properties(kittens_gemm_cdna4 PROPERTIES - HIP_ARCHITECTURES "gfx950" POSITION_INDEPENDENT_CODE ON) - target_include_directories(kittens_gemm_cdna4 PRIVATE - ${HIP_INCLUDE_DIRS} "${HIPKITTENS_CDNA4_INCLUDE_DIR}") - target_compile_options(kittens_gemm_cdna4 PRIVATE - -DKITTENS_CDNA4 -fno-gpu-rdc -O3) - target_link_libraries(kittens_gemm_cdna4 PRIVATE hip::host hip::device) - list(APPEND _kittens_arch_objs $) - set(KITTENS_HAVE_CDNA4 ON PARENT_SCOPE) - endif() + + string(TOUPPER "${A_TAG}" _tag_upper) + set(_target "kittens_gemm_${A_TAG}") + + add_library(${_target} OBJECT ${A_SOURCES}) + set_source_files_properties(${A_SOURCES} PROPERTIES LANGUAGE HIP) + set_target_properties(${_target} PROPERTIES + HIP_ARCHITECTURES "${A_GFX}" POSITION_INDEPENDENT_CODE ON) + target_include_directories(${_target} PRIVATE + ${HIP_INCLUDE_DIRS} "${HIPKITTENS_INCLUDE_DIR}") + target_compile_options(${_target} PRIVATE + -DKITTENS_${_tag_upper} -fno-gpu-rdc -O3 ${A_FLAGS}) + target_link_libraries(${_target} PRIVATE hip::host hip::device) + + set(_kittens_arch_objs ${_kittens_arch_objs} $ PARENT_SCOPE) + set(_kittens_have_defs ${_kittens_have_defs} KITTENS_HAVE_${_tag_upper} PARENT_SCOPE) + endfunction() + + kittens_add_arch(TAG cdna3 GFX gfx942 + SOURCES cdna3/blockwise_fp8_gemm.cpp) + + kittens_add_arch(TAG cdna4 GFX gfx950 + SOURCES cdna4/blockwise_fp8_gemm.cpp cdna4/mxfp8_gemm.cpp + FLAGS -gline-tables-only) + set_source_files_properties(cdna4/blockwise_fp8_gemm.cpp + PROPERTIES COMPILE_FLAGS "-ffast-math") add_library(kittens_gemm SHARED ${_kittens_arch_objs}) set_target_properties(kittens_gemm PROPERTIES LINKER_LANGUAGE HIP) target_include_directories(kittens_gemm PRIVATE ${HIP_INCLUDE_DIRS}) target_link_libraries(kittens_gemm PUBLIC hip::host hip::device) + set(KITTENS_HAVE_DEFS "${_kittens_have_defs}" PARENT_SCOPE) + install(TARGETS kittens_gemm DESTINATION ${CMAKE_INSTALL_PREFIX}/transformer_engine/lib) endif() diff --git a/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm.cpp b/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm.cpp index 90b5fb286..86e5517ad 100644 --- a/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm.cpp +++ b/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm.cpp @@ -6,8 +6,9 @@ #include #include "kittens.cuh" #include "../kittens_common.h" +#include "../kittens_kernel_common.cuh" -namespace { +namespace te_kittens::cdna3 { #include "blockwise_fp8_gemm_helper.cuh" @@ -354,16 +355,6 @@ void micro_tk(const micro_globals g) { store_output(g.c.raw_ptr, C_accum[1], row * 4 + warp_row + WARPS_ROW, col * 4 + warp_col, M, N); } -#define BOOL_SWITCH(val, NAME, ...) \ - if (val) { constexpr bool NAME = true; __VA_ARGS__ } \ - else { constexpr bool NAME = false; __VA_ARGS__ } - -static GemmEpilogue select_epilogue(bool has_bias, bool has_gelu, bool has_beta) { - if (has_gelu) return has_beta ? GemmEpilogue::GELU_AUX_BETA : GemmEpilogue::GELU_AUX; - if (has_bias) return has_beta ? GemmEpilogue::BIAS_BETA : GemmEpilogue::BIAS; - return has_beta ? GemmEpilogue::BETA : GemmEpilogue::DEFAULT; -} - template static void dispatch_micro_epilogue(micro_globals g) { @@ -374,8 +365,8 @@ static void dispatch_micro_epilogue(micro_globals g) { hipFuncSetAttribute((void*)kern, hipFuncAttributeMaxDynamicSharedMemorySize, mem_size); kern<<>>(g); }; - BOOL_SWITCH(is_partial_m, IS_PARTIAL_M, - BOOL_SWITCH(is_partial_n, IS_PARTIAL_N, + KITTENS_BOOL_SWITCH(is_partial_m, IS_PARTIAL_M, + KITTENS_BOOL_SWITCH(is_partial_n, IS_PARTIAL_N, launch(micro_tk); ) ) @@ -405,7 +396,7 @@ template static void dispatch_micro(micro_globals g, bool has_bias, bool has_gelu, bool has_beta, bool has_partial_k) { - BOOL_SWITCH(has_partial_k, IS_PARTIAL_K, + KITTENS_BOOL_SWITCH(has_partial_k, IS_PARTIAL_K, dispatch_micro_k(g, has_bias, has_gelu, has_beta); ) } @@ -465,11 +456,9 @@ class BlockwiseGemmCdna3 final : public BlockwiseGemmBackend { } }; -#undef BOOL_SWITCH - -} +} // namespace te_kittens::cdna3 BlockwiseGemmBackend *BlockwiseGemmBackend::get_cdna3() { - static BlockwiseGemmCdna3 impl; + static te_kittens::cdna3::BlockwiseGemmCdna3 impl; return &impl; } diff --git a/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm_helper.cuh b/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm_helper.cuh index 413b126c8..cabed0574 100644 --- a/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm_helper.cuh +++ b/transformer_engine/common/gemm/kittens/cdna3/blockwise_fp8_gemm_helper.cuh @@ -8,6 +8,7 @@ #include #include "kittens.cuh" #include "../../../util/math.h" +using namespace te_kittens::blockwise; // NOLINT(build/namespaces) typedef int int32x4_lds_t __attribute__((ext_vector_type(4))); struct __attribute__((packed)) buf_res { const void *ptr; uint32_t range; uint32_t config; }; @@ -87,12 +88,6 @@ __device__ inline float rtne_bias(float v) { return __builtin_bit_cast(float, bits); } -__device__ inline float read_elem(const void *p, int dtype, int idx) { - if (dtype == 6) return __bfloat162float(reinterpret_cast(p)[idx]); - if (dtype == 5) return __half2float(reinterpret_cast(p)[idx]); - return reinterpret_cast(p)[idx]; -} - template __device__ inline float rtne_cast_roundtrip(float v) { if constexpr (std::is_same_v) { @@ -132,26 +127,6 @@ __device__ inline void store_output(OType *c_ptr, const AccType &Cacc, } } -enum struct GemmEpilogue { - DEFAULT, - BIAS, - GELU_AUX, - BETA, - BIAS_BETA, - GELU_AUX_BETA, -}; - -__host__ __device__ inline constexpr bool epilogue_has_bias(GemmEpilogue e) { - return e == GemmEpilogue::BIAS || e == GemmEpilogue::BIAS_BETA; -} -__host__ __device__ inline constexpr bool epilogue_has_gelu(GemmEpilogue e) { - return e == GemmEpilogue::GELU_AUX || e == GemmEpilogue::GELU_AUX_BETA; -} -__host__ __device__ inline constexpr bool epilogue_has_beta(GemmEpilogue e) { - return e == GemmEpilogue::BETA || e == GemmEpilogue::BIAS_BETA - || e == GemmEpilogue::GELU_AUX_BETA; -} - template __device__ inline void apply_epilogue( AccType &Cacc, int Rtile, int Ctile, int M, int N, diff --git a/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm.cpp b/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm.cpp index 0c217cecc..1cadaf47d 100644 --- a/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm.cpp +++ b/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm.cpp @@ -9,9 +9,10 @@ #include #include "kittens.cuh" #include "../kittens_common.h" +#include "../kittens_kernel_common.cuh" -namespace { +namespace te_kittens::cdna4 { #include "blockwise_fp8_gemm_helper.cuh" @@ -941,16 +942,6 @@ void micro_tk_partial_k(micro_globals } -#define BOOL_SWITCH(val, NAME, ...) \ - if (val) { constexpr bool NAME = true; __VA_ARGS__ } \ - else { constexpr bool NAME = false; __VA_ARGS__ } - -static GemmEpilogue select_epilogue(bool has_bias, bool has_gelu, bool has_beta) { - if (has_gelu) return has_beta ? GemmEpilogue::GELU_AUX_BETA : GemmEpilogue::GELU_AUX; - if (has_bias) return has_beta ? GemmEpilogue::BIAS_BETA : GemmEpilogue::BIAS; - return has_beta ? GemmEpilogue::BETA : GemmEpilogue::DEFAULT; -} - template static void dispatch_micro_kernel(micro_globals_fp8 g) { @@ -991,13 +982,12 @@ static void dispatch_micro_epilogue(int cbsz, int blgp, bool has_bias, bool has_ template static void dispatch_micro(bool is_1d2d, int cbsz, int blgp, bool has_bias, bool has_gelu, bool has_beta, bool has_partial_k, micro_globals_fp8 g) { - BOOL_SWITCH(is_1d2d, IS_1D2D, - BOOL_SWITCH(has_partial_k, IS_PARTIAL_K, + KITTENS_BOOL_SWITCH(is_1d2d, IS_1D2D, + KITTENS_BOOL_SWITCH(has_partial_k, IS_PARTIAL_K, dispatch_micro_epilogue(cbsz, blgp, has_bias, has_gelu, has_beta, g); ) ) } -#undef BOOL_SWITCH template static void launch_pow2_kernel(const pow2_kernel_args &a) { @@ -1049,7 +1039,7 @@ static void launch_pow2(int cbsz, int blgp, bool has_bias, bool has_gelu, bool h const int padM = tiles_M * BLOCK_M; const int padN = tiles_N * BLOCK_N; - const size_t sa_bytes = align_up_pow2ws((size_t)k_iters * padM * sizeof(uint32_t)); + const size_t sa_bytes = kittens_align_up((size_t)k_iters * padM * sizeof(uint32_t), 256); uint32_t *packed_sa = reinterpret_cast(workspace); uint32_t *packed_sb = reinterpret_cast((uint8_t *)workspace + sa_bytes); @@ -1117,7 +1107,7 @@ class BlockwiseGemmCdna4 final : public BlockwiseGemmBackend { const int k_iters = K / BLOCK_K; const int padM = ((kM + BLOCK_M - 1) / BLOCK_M) * BLOCK_M; const int padN = ((kN + BLOCK_N - 1) / BLOCK_N) * BLOCK_N; - const size_t pow2_ws_bytes = align_up_pow2ws((size_t)k_iters * padM * sizeof(uint32_t)) + + const size_t pow2_ws_bytes = kittens_align_up((size_t)k_iters * padM * sizeof(uint32_t), 256) + (size_t)k_iters * padN * sizeof(uint32_t); void *owned_ws = nullptr; if (use_pow2 && !has_partial_k && @@ -1163,9 +1153,9 @@ class BlockwiseGemmCdna4 final : public BlockwiseGemmBackend { } }; -} +} // namespace te_kittens::cdna4 BlockwiseGemmBackend *BlockwiseGemmBackend::get_cdna4() { - static BlockwiseGemmCdna4 impl; + static te_kittens::cdna4::BlockwiseGemmCdna4 impl; return &impl; } diff --git a/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm_helper.cuh b/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm_helper.cuh index 7207a97fd..b100d12a5 100644 --- a/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm_helper.cuh +++ b/transformer_engine/common/gemm/kittens/cdna4/blockwise_fp8_gemm_helper.cuh @@ -8,6 +8,7 @@ #include #include "kittens.cuh" #include "../../../util/math.h" +using namespace te_kittens::blockwise; // NOLINT(build/namespaces) template struct RowScale { float2 v[HEIGHT][2]; }; @@ -77,12 +78,6 @@ __device__ inline void store_output(OType *c_ptr, const AccType &acc, } } -__device__ inline float read_elem(const void *p, int dtype, int idx) { - if (dtype == 6) return __bfloat162float(reinterpret_cast(p)[idx]); - if (dtype == 5) return __half2float(reinterpret_cast(p)[idx]); - return reinterpret_cast(p)[idx]; -} - template __device__ inline float round_to_out_dtype(float v) { if constexpr (std::is_same_v) { @@ -94,26 +89,6 @@ __device__ inline float round_to_out_dtype(float v) { } } -enum struct GemmEpilogue { - DEFAULT, - BIAS, - GELU_AUX, - BETA, - BIAS_BETA, - GELU_AUX_BETA, -}; - -__host__ __device__ inline constexpr bool epilogue_has_bias(GemmEpilogue e) { - return e == GemmEpilogue::BIAS || e == GemmEpilogue::BIAS_BETA; -} -__host__ __device__ inline constexpr bool epilogue_has_gelu(GemmEpilogue e) { - return e == GemmEpilogue::GELU_AUX || e == GemmEpilogue::GELU_AUX_BETA; -} -__host__ __device__ inline constexpr bool epilogue_has_beta(GemmEpilogue e) { - return e == GemmEpilogue::BETA || e == GemmEpilogue::BIAS_BETA - || e == GemmEpilogue::GELU_AUX_BETA; -} - template __device__ inline void apply_epilogue( AccType &acc, int m_off, int n_off, int M, int N, @@ -455,4 +430,3 @@ static void launch_pack_scales_pow2(const float *scales, uint32_t *packed, int p pack_scales_pow2_kernel<<>>(scales, packed, padded_dim, real_dim, scale_K, k_iters, scale_block); } -static inline size_t align_up_pow2ws(size_t x) { return (x + 255) & ~size_t(255); } diff --git a/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.cpp b/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.cpp index 795d6dfab..df343efca 100644 --- a/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.cpp +++ b/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.cpp @@ -4,10 +4,12 @@ *************************************************************************/ #include "kittens.cuh" -#include "mxfp8_gemm.h" +#include "../kittens_common.h" +#include "../kittens_kernel_common.cuh" #include #include +namespace te_kittens::cdna4 { constexpr int NUM_WARPS = 8; constexpr int NUM_THREADS = NUM_WARPS * kittens::WARP_THREADS; @@ -1108,10 +1110,6 @@ __global__ __launch_bounds__(NUM_THREADS, 2) void mxfp8_gemm_nt_kernel(const gl_ block_m, block_row, block_col, warp_m, warp_n); } -#define HK_BOOL_SWITCH(val, NAME, ...) \ - if (val) { constexpr bool NAME = true; __VA_ARGS__ } \ - else { constexpr bool NAME = false; __VA_ARGS__ } - template static void launch_gemm_typed( const void *A, const void *B, void *C, @@ -1251,10 +1249,6 @@ static void launch_pack_scales_fused_varying(const uint8_t *const *d_scale_ptrs, nullptr, ln, dim, 0, max_k_iters, tiles_per_col, d_scale_ptrs, 0, d_k_iters_arr, d_output_offsets); } -static size_t hk_align_up(size_t x, size_t a) { - return (x + a - 1) & ~(a - 1); -} - static bool check_tn_constraints(int M, int N, int K) { return M % BLOCK_ROW == 0 && N % BLOCK_COL == 0 && K % BLOCK_K == 0 && K >= 256; } @@ -1286,8 +1280,8 @@ static bool mxfp8_gemm_impl( // Lane-native scale buffers: A = 256 words/tile, B = 512 (hi/lo pair). If they overflow the // caller's budget we return false and fall back to hipBLASLt. - size_t sa_bytes = hk_align_up((size_t)k_iters * tiles_M * 256 * sizeof(uint32_t), 256); - size_t sb_bytes = hk_align_up((size_t)k_iters * tiles_N * 512 * sizeof(uint32_t), 256); + size_t sa_bytes = kittens_align_up((size_t)k_iters * tiles_M * 256 * sizeof(uint32_t), 256); + size_t sb_bytes = kittens_align_up((size_t)k_iters * tiles_N * 512 * sizeof(uint32_t), 256); if (workspace_size < sa_bytes + sb_bytes) return false; auto *packed_sa = (uint32_t *)workspace; @@ -1324,7 +1318,7 @@ static int out_code(int dt) { } } -bool kittens_mxfp8_gemm( +static bool mxfp8_gemm( const void *A, const void *B, void *C, const void *scale_A, const void *scale_B, int M, int N, int K, @@ -1344,46 +1338,43 @@ bool kittens_mxfp8_gemm( bool accumulate = beta != 0.0f; bool result = false; - HK_BOOL_SWITCH(transa, TRANSA, - HK_BOOL_SWITCH(transb, TRANSB, - HK_BOOL_SWITCH(accumulate, ACCUMULATE, + KITTENS_BOOL_SWITCH(transa, TRANSA, + KITTENS_BOOL_SWITCH(transb, TRANSB, + KITTENS_BOOL_SWITCH(accumulate, ACCUMULATE, if constexpr (!(TRANSA && TRANSB)) { result = mxfp8_gemm_impl(A, B, C, scale_A, scale_B, M, N, K, a_fp8, b_fp8, bias, bias_dc, aux_gelu, out_dc, aux_dc, workspace, workspace_size, stream); } else { - assert(0 && "kittens_mxfp8_gemm: TT layout is not supported"); + assert(0 && "mxfp8_gemm: TT layout is not supported"); } ))) // NOLINT(*) return result; } -static size_t hk_align256(size_t x) { - return (x + 255) & ~(size_t)255; +// Traces why a grouped launch was declined +static void warn_fallback(const char *tag, const char *reason) { + static bool enabled = [] { + const char *v = std::getenv("NVTE_CUTLASS_GROUPED_GEMM_WARN_FALLBACK"); + return v && v[0] == '1'; + }(); + if (enabled) { + fprintf(stderr, "[%s] falling back: %s\n", tag, reason); + } } -bool kittens_grouped_mxfp8_gemm( +static bool grouped_mxfp8_gemm( const void *const *A_array, const void *const *B_array, void *const *C_array, const void *const *scale_A_array, const void *const *scale_B_array, int M, const int *N_array, int K, int num_experts, bool transa, bool transb, int a_dtype, int b_dtype, int out_dtype, void *workspace, size_t workspace_size, hipStream_t stream) { - auto warn = [](const char *reason) { - static bool enabled = [] { - const char *v = std::getenv("NVTE_CUTLASS_GROUPED_GEMM_WARN_FALLBACK"); - return v && v[0] == '1'; - }(); - if (enabled) { - fprintf(stderr, "[HK-grouped] falling back: %s\n", reason); - } - }; - - if (transa && transb) { warn("TT layout not supported"); return false; } - if (!transa && transb) { warn("NT layout: use kittens_grouped_mxfp8_wgrad"); return false; } - if (M % BLOCK_ROW != 0) { warn("M not 256-aligned"); return false; } - if (K % BLOCK_K != 0 || K < 256) { warn("K not 128-aligned or < 256"); return false; } - if (num_experts <= 0) { warn("num_experts <= 0"); return false; } + if (transa && transb) { warn_fallback("HK-grouped", "TT layout not supported"); return false; } + if (!transa && transb) { warn_fallback("HK-grouped", "NT layout: use grouped_mxfp8_wgrad"); return false; } + if (M % BLOCK_ROW != 0) { warn_fallback("HK-grouped", "M not 256-aligned"); return false; } + if (K % BLOCK_K != 0 || K < 256) { warn_fallback("HK-grouped", "K not 128-aligned or < 256"); return false; } + if (num_experts <= 0) { warn_fallback("HK-grouped", "num_experts <= 0"); return false; } int tiles_M = M / BLOCK_COL; int k_iters = K / BLOCK_K; @@ -1393,7 +1384,7 @@ bool kittens_grouped_mxfp8_gemm( int total_N = 0; int total_n_tiles = 0; for (int g = 0; g < num_experts; g++) { - if (N_array[g] % BLOCK_COL != 0) { warn("N_array not 256-aligned"); return false; } + if (N_array[g] % BLOCK_COL != 0) { warn_fallback("HK-grouped", "N_array not 256-aligned"); return false; } h_tile_offsets[g] = total_n_tiles; total_N += N_array[g]; total_n_tiles += N_array[g] / BLOCK_COL; @@ -1403,18 +1394,18 @@ bool kittens_grouped_mxfp8_gemm( int grid = tiles_M * total_n_tiles; if (grid == 0) return true; - size_t sa_pk_bytes = hk_align256((size_t)k_iters * num_experts * M * sizeof(uint32_t)); - size_t sb_pk_bytes = hk_align256((size_t)2 * k_iters * total_N * sizeof(uint32_t)); - size_t a_ptrs_bytes = hk_align256((size_t)num_experts * sizeof(void *)); - size_t b_ptrs_bytes = hk_align256((size_t)num_experts * sizeof(void *)); - size_t c_ptrs_bytes = hk_align256((size_t)num_experts * sizeof(void *)); - size_t sa_ptrs_bytes = hk_align256((size_t)num_experts * sizeof(void *)); - size_t offsets_bytes = hk_align256((size_t)(num_experts + 1) * sizeof(int)); - size_t sb_off_bytes = hk_align256((size_t)(num_experts + 1) * sizeof(int)); + size_t sa_pk_bytes = kittens_align_up((size_t)k_iters * num_experts * M * sizeof(uint32_t), 256); + size_t sb_pk_bytes = kittens_align_up((size_t)2 * k_iters * total_N * sizeof(uint32_t), 256); + size_t a_ptrs_bytes = kittens_align_up((size_t)num_experts * sizeof(void *), 256); + size_t b_ptrs_bytes = kittens_align_up((size_t)num_experts * sizeof(void *), 256); + size_t c_ptrs_bytes = kittens_align_up((size_t)num_experts * sizeof(void *), 256); + size_t sa_ptrs_bytes = kittens_align_up((size_t)num_experts * sizeof(void *), 256); + size_t offsets_bytes = kittens_align_up((size_t)(num_experts + 1) * sizeof(int), 256); + size_t sb_off_bytes = kittens_align_up((size_t)(num_experts + 1) * sizeof(int), 256); size_t total_ws = sa_pk_bytes + sb_pk_bytes + a_ptrs_bytes + b_ptrs_bytes + c_ptrs_bytes + sa_ptrs_bytes + offsets_bytes + sb_off_bytes; if (workspace_size < total_ws) { - warn("workspace too small"); return false; + warn_fallback("HK-grouped", "workspace too small"); return false; } uint8_t *ws = (uint8_t *)workspace; @@ -1438,8 +1429,8 @@ bool kittens_grouped_mxfp8_gemm( std::vector h_sb_tile_offsets(num_experts + 1); int sb_tile_cursor = 0; uint32_t *sb_cursor = sb_pk; - HK_BOOL_SWITCH(!transa, COLWISE_A, - HK_BOOL_SWITCH(transb, COLWISE_B, + KITTENS_BOOL_SWITCH(!transa, COLWISE_A, + KITTENS_BOOL_SWITCH(transb, COLWISE_B, // Pack weight scales: single fused launch for all experts launch_pack_scales_fused( (const uint8_t *const *)d_sa_ptrs, sa_pk, @@ -1804,25 +1795,15 @@ void mxfp8_wgrad_nt_kernel( kittens::store(C_local, oC, out_coord_C); kittens::store(C_local, oD, out_coord_D); } -bool kittens_grouped_mxfp8_wgrad(const void *const *A_array, const void *const *B_array, void *const *D_array, +static bool grouped_mxfp8_wgrad(const void *const *A_array, const void *const *B_array, void *const *D_array, const void *const *scale_A_array, const void *const *scale_B_array, int N, int K, const int *M_array, int num_experts, int a_dtype, int b_dtype, int out_dtype, bool accumulate, void *workspace, size_t workspace_size, hipStream_t stream) { - auto warn = [](const char *reason) { - static bool enabled = [] { - const char *v = std::getenv("NVTE_CUTLASS_GROUPED_GEMM_WARN_FALLBACK"); - return v && v[0] == '1'; - }(); - if (enabled) { - fprintf(stderr, "[HK-wgrad] falling back: %s\n", reason); - } - }; - - if (N % BLOCK_ROW != 0) { warn("N not 256-aligned"); return false; } - if (K % BLOCK_COL != 0) { warn("K not 256-aligned"); return false; } - if (num_experts <= 0) { warn("num_experts <= 0"); return false; } + if (N % BLOCK_ROW != 0) { warn_fallback("HK-wgrad", "N not 256-aligned"); return false; } + if (K % BLOCK_COL != 0) { warn_fallback("HK-wgrad", "K not 256-aligned"); return false; } + if (num_experts <= 0) { warn_fallback("HK-wgrad", "num_experts <= 0"); return false; } int tiles_M = N / BLOCK_ROW; int tiles_N = K / BLOCK_COL; @@ -1842,7 +1823,7 @@ bool kittens_grouped_mxfp8_wgrad(const void *const *A_array, const void *const * int M_g = M_array[g]; if (M_g == 0) continue; if (M_g % BLOCK_K != 0 || M_g < 256) { - warn("M_i not 128-aligned or < 256"); + warn_fallback("HK-wgrad", "M_i not 128-aligned or < 256"); return false; } int k_iters_g = M_g / BLOCK_K; @@ -1888,18 +1869,18 @@ bool kittens_grouped_mxfp8_wgrad(const void *const *A_array, const void *const * idx++; } - size_t sa_pk_bytes = hk_align256(total_sa_entries * sizeof(uint32_t)); - size_t sb_pk_bytes = hk_align256(total_sb_entries * sizeof(uint32_t)); - size_t info_bytes = hk_align256((size_t)num_active * sizeof(WgradExpertInfo)); - size_t sa_ptrs_bytes = hk_align256((size_t)num_active * sizeof(void *)); - size_t sb_ptrs_bytes = hk_align256((size_t)num_active * sizeof(void *)); - size_t ki_arr_bytes = hk_align256((size_t)num_active * sizeof(int)); - size_t sa_off_bytes = hk_align256((size_t)num_active * sizeof(int)); - size_t sb_off_bytes = hk_align256((size_t)num_active * sizeof(int)); + size_t sa_pk_bytes = kittens_align_up(total_sa_entries * sizeof(uint32_t), 256); + size_t sb_pk_bytes = kittens_align_up(total_sb_entries * sizeof(uint32_t), 256); + size_t info_bytes = kittens_align_up((size_t)num_active * sizeof(WgradExpertInfo), 256); + size_t sa_ptrs_bytes = kittens_align_up((size_t)num_active * sizeof(void *), 256); + size_t sb_ptrs_bytes = kittens_align_up((size_t)num_active * sizeof(void *), 256); + size_t ki_arr_bytes = kittens_align_up((size_t)num_active * sizeof(int), 256); + size_t sa_off_bytes = kittens_align_up((size_t)num_active * sizeof(int), 256); + size_t sb_off_bytes = kittens_align_up((size_t)num_active * sizeof(int), 256); size_t total_ws = sa_pk_bytes + sb_pk_bytes + info_bytes + sa_ptrs_bytes + sb_ptrs_bytes + ki_arr_bytes + sa_off_bytes + sb_off_bytes; if (workspace_size < total_ws) { - warn("workspace too small"); + warn_fallback("HK-wgrad", "workspace too small"); return false; } @@ -1941,7 +1922,7 @@ bool kittens_grouped_mxfp8_wgrad(const void *const *A_array, const void *const * gl_fp8_rt gl_B((kittens::fp8e4m3 *)B_array[0], nullptr, nullptr, (size_t)max_M, (size_t)K); auto launch_wgrad = [&](auto gl_D) { - HK_BOOL_SWITCH(accumulate, ACCUMULATE, + KITTENS_BOOL_SWITCH(accumulate, ACCUMULATE, mxfp8_wgrad_nt_kernel<<>>( gl_A, gl_B, gl_D, gl_SA, gl_SB, d_info, tiles_M, tiles_N, tiles_per_expert); @@ -1960,4 +1941,37 @@ bool kittens_grouped_mxfp8_wgrad(const void *const *A_array, const void *const * return true; } -#undef HK_BOOL_SWITCH +class MXFP8GemmCdna4 final : public MXFP8GemmBackend { + public: + bool gemm(const MXFP8GemmArgs &args) override { + return mxfp8_gemm(args.A, args.B, args.C, args.scale_A, args.scale_B, + args.M, args.N, args.K, args.transa, args.transb, + args.a_dtype, args.b_dtype, args.bias, args.bias_dtype, + args.aux_gelu, args.out_dtype, args.aux_dtype, args.beta, + args.workspace, args.workspace_size, args.stream); + } + + bool grouped_gemm(const MXFP8GroupedGemmArgs &args) override { + return grouped_mxfp8_gemm(args.A_array, args.B_array, args.C_array, + args.scale_A_array, args.scale_B_array, + args.M, args.N_array, args.K, args.num_experts, + args.transa, args.transb, + args.a_dtype, args.b_dtype, args.out_dtype, + args.workspace, args.workspace_size, args.stream); + } + + bool grouped_wgrad(const MXFP8WgradArgs &args) override { + return grouped_mxfp8_wgrad(args.A_array, args.B_array, args.D_array, + args.scale_A_array, args.scale_B_array, + args.N, args.K, args.M_array, args.num_experts, + args.a_dtype, args.b_dtype, args.out_dtype, args.accumulate, + args.workspace, args.workspace_size, args.stream); + } +}; + +} // namespace te_kittens::cdna4 + +MXFP8GemmBackend *MXFP8GemmBackend::get_cdna4() { + static te_kittens::cdna4::MXFP8GemmCdna4 impl; + return &impl; +} diff --git a/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.h b/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.h deleted file mode 100644 index c32269c4f..000000000 --- a/transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.h +++ /dev/null @@ -1,37 +0,0 @@ -/************************************************************************* - * Copyright (c) 2026, Advanced Micro Devices, Inc. All rights reserved. - * License for AMD contributions = MIT. See LICENSE for more information -*************************************************************************/ - -#pragma once - -#include -#include - -#include "../kittens_common.h" - -bool kittens_mxfp8_gemm( - const void *A, const void *B, void *C, - const void *scale_A, const void *scale_B, - int M, int N, int K, - bool transa, bool transb, - int a_dtype, int b_dtype, - const void *bias, int bias_dtype, - void *aux_gelu, int out_dtype, int aux_dtype, - float beta, void *workspace, size_t workspace_size, - hipStream_t stream); - -bool kittens_grouped_mxfp8_gemm( - const void *const *A_array, const void *const *B_array, void *const *C_array, - const void *const *scale_A_array, const void *const *scale_B_array, - int M, const int *N_array, int K, int num_experts, - bool transa, bool transb, int a_dtype, int b_dtype, int out_dtype, - void *workspace, size_t workspace_size, hipStream_t stream); - -bool kittens_grouped_mxfp8_wgrad( - const void *const *A_array, const void *const *B_array, void *const *D_array, - const void *const *scale_A_array, const void *const *scale_B_array, - int N, int K, const int *M_array, int num_experts, - int a_dtype, int b_dtype, int out_dtype, - bool accumulate, - void *workspace, size_t workspace_size, hipStream_t stream); diff --git a/transformer_engine/common/gemm/kittens/kittens_common.h b/transformer_engine/common/gemm/kittens/kittens_common.h index 5c129bcc7..4a08958aa 100644 --- a/transformer_engine/common/gemm/kittens/kittens_common.h +++ b/transformer_engine/common/gemm/kittens/kittens_common.h @@ -53,22 +53,25 @@ class BlockwiseGemmBackend { virtual ~BlockwiseGemmBackend() = default; virtual void run(const BlockwiseGemmArgs &args) = 0; + static BlockwiseGemmBackend *get() { + const int arch = transformer_engine::cuda::sm_arch(); +#ifdef KITTENS_HAVE_CDNA4 + if (arch == 95) { + return get_cdna4(); + } +#endif +#ifdef KITTENS_HAVE_CDNA3 + if (arch == 94) { + return get_cdna3(); + } +#endif + static_cast(arch); + return nullptr; + } + private: static BlockwiseGemmBackend *get_cdna3(); static BlockwiseGemmBackend *get_cdna4(); - - friend inline void kittens_blockwise_fp8_gemm( - const void *A, const void *B, void *C, - const void *scale_A, const void *scale_B, - int M, int N, int K, - int a_dtype, int b_dtype, - int a_scaling_mode, int b_scaling_mode, - int out_dtype, - const void *bias, int bias_dtype, - const void *gelu_aux, int gelu_aux_dtype, - const void *c_in, float beta, - void *workspace, size_t workspace_size, - hipStream_t stream); }; inline void kittens_blockwise_fp8_gemm( @@ -83,23 +86,11 @@ inline void kittens_blockwise_fp8_gemm( const void *c_in, float beta, void *workspace, size_t workspace_size, hipStream_t stream) { - const int arch = transformer_engine::cuda::sm_arch(); - - BlockwiseGemmBackend *backend = nullptr; -#ifdef KITTENS_HAVE_CDNA4 - if (arch == 95) { - backend = BlockwiseGemmBackend::get_cdna4(); - } -#endif -#ifdef KITTENS_HAVE_CDNA3 - if (arch == 94) { - backend = BlockwiseGemmBackend::get_cdna3(); - } -#endif + BlockwiseGemmBackend *backend = BlockwiseGemmBackend::get(); if (backend == nullptr) { throw std::runtime_error( "kittens_blockwise_fp8_gemm: not implemented for this GPU arch (sm_arch=" + - std::to_string(arch) + "). Only gfx942 and gfx950 are supported."); + std::to_string(transformer_engine::cuda::sm_arch()) + "). Only gfx942 and gfx950 are supported."); } BlockwiseGemmArgs args{ @@ -109,3 +100,136 @@ inline void kittens_blockwise_fp8_gemm( workspace, workspace_size, stream}; backend->run(args); } + +struct MXFP8GemmArgs { + const void *A; + const void *B; + void *C; + const void *scale_A; + const void *scale_B; + int M, N, K; + bool transa, transb; + int a_dtype, b_dtype; + const void *bias; + int bias_dtype; + void *aux_gelu; + int out_dtype, aux_dtype; + float beta; + void *workspace; + size_t workspace_size; + hipStream_t stream; +}; + +struct MXFP8GroupedGemmArgs { + const void *const *A_array; + const void *const *B_array; + void *const *C_array; + const void *const *scale_A_array; + const void *const *scale_B_array; + int M; + const int *N_array; + int K; + int num_experts; + bool transa, transb; + int a_dtype, b_dtype, out_dtype; + void *workspace; + size_t workspace_size; + hipStream_t stream; +}; + +struct MXFP8WgradArgs { + const void *const *A_array; + const void *const *B_array; + void *const *D_array; + const void *const *scale_A_array; + const void *const *scale_B_array; + int N, K; + const int *M_array; + int num_experts; + int a_dtype, b_dtype, out_dtype; + bool accumulate; + void *workspace; + size_t workspace_size; + hipStream_t stream; +}; + +class MXFP8GemmBackend { + public: + virtual ~MXFP8GemmBackend() = default; + virtual bool gemm(const MXFP8GemmArgs &args) = 0; + virtual bool grouped_gemm(const MXFP8GroupedGemmArgs &args) = 0; + virtual bool grouped_wgrad(const MXFP8WgradArgs &args) = 0; + + static MXFP8GemmBackend *get() { +#ifdef KITTENS_HAVE_CDNA4 + if (transformer_engine::cuda::sm_arch() == 95) { + return get_cdna4(); + } +#endif + return nullptr; + } + + private: + static MXFP8GemmBackend *get_cdna4(); +}; + +inline bool kittens_mxfp8_supported() { return MXFP8GemmBackend::get() != nullptr; } + +inline bool kittens_mxfp8_gemm( + const void *A, const void *B, void *C, + const void *scale_A, const void *scale_B, + int M, int N, int K, + bool transa, bool transb, + int a_dtype, int b_dtype, + const void *bias, int bias_dtype, + void *aux_gelu, int out_dtype, int aux_dtype, + float beta, void *workspace, size_t workspace_size, + hipStream_t stream) { + MXFP8GemmBackend *backend = MXFP8GemmBackend::get(); + if (backend == nullptr) { + return false; + } + + MXFP8GemmArgs args{ + A, B, C, scale_A, scale_B, M, N, K, transa, transb, + a_dtype, b_dtype, bias, bias_dtype, aux_gelu, out_dtype, aux_dtype, + beta, workspace, workspace_size, stream}; + return backend->gemm(args); +} + +inline bool kittens_grouped_mxfp8_gemm( + const void *const *A_array, const void *const *B_array, void *const *C_array, + const void *const *scale_A_array, const void *const *scale_B_array, + int M, const int *N_array, int K, int num_experts, + bool transa, bool transb, int a_dtype, int b_dtype, int out_dtype, + void *workspace, size_t workspace_size, hipStream_t stream) { + MXFP8GemmBackend *backend = MXFP8GemmBackend::get(); + if (backend == nullptr) { + return false; + } + + MXFP8GroupedGemmArgs args{ + A_array, B_array, C_array, scale_A_array, scale_B_array, + M, N_array, K, num_experts, transa, transb, + a_dtype, b_dtype, out_dtype, workspace, workspace_size, stream}; + return backend->grouped_gemm(args); +} + +inline bool kittens_grouped_mxfp8_wgrad( + const void *const *A_array, const void *const *B_array, void *const *D_array, + const void *const *scale_A_array, const void *const *scale_B_array, + int N, int K, const int *M_array, int num_experts, + int a_dtype, int b_dtype, int out_dtype, + bool accumulate, + void *workspace, size_t workspace_size, hipStream_t stream) { + MXFP8GemmBackend *backend = MXFP8GemmBackend::get(); + if (backend == nullptr) { + return false; + } + + MXFP8WgradArgs args{ + A_array, B_array, D_array, scale_A_array, scale_B_array, + N, K, M_array, num_experts, a_dtype, b_dtype, out_dtype, accumulate, + workspace, workspace_size, stream}; + return backend->grouped_wgrad(args); +} diff --git a/transformer_engine/common/gemm/kittens/kittens_kernel_common.cuh b/transformer_engine/common/gemm/kittens/kittens_kernel_common.cuh new file mode 100644 index 000000000..744f69d9c --- /dev/null +++ b/transformer_engine/common/gemm/kittens/kittens_kernel_common.cuh @@ -0,0 +1,53 @@ +/************************************************************************* + * Copyright (c) 2026, Advanced Micro Devices, Inc. All rights reserved. + * License for AMD contributions = MIT. See LICENSE for more information +*************************************************************************/ + +#pragma once + +#include "hip/hip_runtime.h" + +#include + +#define KITTENS_BOOL_SWITCH(val, NAME, ...) \ + if (val) { constexpr bool NAME = true; __VA_ARGS__ } \ + else { constexpr bool NAME = false; __VA_ARGS__ } + +static inline size_t kittens_align_up(size_t x, size_t a) { return (x + a - 1) & ~(a - 1); } + +namespace te_kittens::blockwise { + +// dtype codes are NVTEDType: 6 = bfloat16, 5 = float16, anything else = float32 +__device__ inline float read_elem(const void *p, int dtype, int idx) { + if (dtype == 6) return __bfloat162float(reinterpret_cast(p)[idx]); + if (dtype == 5) return __half2float(reinterpret_cast(p)[idx]); + return reinterpret_cast(p)[idx]; +} + +enum struct GemmEpilogue { + DEFAULT, + BIAS, + GELU_AUX, + BETA, + BIAS_BETA, + GELU_AUX_BETA, +}; + +__host__ __device__ inline constexpr bool epilogue_has_bias(GemmEpilogue e) { + return e == GemmEpilogue::BIAS || e == GemmEpilogue::BIAS_BETA; +} +__host__ __device__ inline constexpr bool epilogue_has_gelu(GemmEpilogue e) { + return e == GemmEpilogue::GELU_AUX || e == GemmEpilogue::GELU_AUX_BETA; +} +__host__ __device__ inline constexpr bool epilogue_has_beta(GemmEpilogue e) { + return e == GemmEpilogue::BETA || e == GemmEpilogue::BIAS_BETA + || e == GemmEpilogue::GELU_AUX_BETA; +} + +inline GemmEpilogue select_epilogue(bool has_bias, bool has_gelu, bool has_beta) { + if (has_gelu) return has_beta ? GemmEpilogue::GELU_AUX_BETA : GemmEpilogue::GELU_AUX; + if (has_bias) return has_beta ? GemmEpilogue::BIAS_BETA : GemmEpilogue::BIAS; + return has_beta ? GemmEpilogue::BETA : GemmEpilogue::DEFAULT; +} + +} // namespace te_kittens::blockwise diff --git a/transformer_engine/common/gemm/rocm_gemm.cu b/transformer_engine/common/gemm/rocm_gemm.cu index 082c96ee7..ce8a35b9e 100644 --- a/transformer_engine/common/gemm/rocm_gemm.cu +++ b/transformer_engine/common/gemm/rocm_gemm.cu @@ -34,9 +34,15 @@ #ifdef USE_HIPKITTENS_GEMM #include "kittens/kittens_common.h" -#ifdef KITTENS_HAVE_CDNA4 -#include "kittens/cdna4/mxfp8_gemm.h" -#endif + +// Kittens enums mirror NVTE enums by value because kittens lib is built without TE. +static_assert(KITTENS_FLOAT32 == static_cast(transformer_engine::DType::kFloat32), "KittensDType out of sync with NVTEDType"); +static_assert(KITTENS_FLOAT16 == static_cast(transformer_engine::DType::kFloat16), "KittensDType out of sync with NVTEDType"); +static_assert(KITTENS_BFLOAT16 == static_cast(transformer_engine::DType::kBFloat16), "KittensDType out of sync with NVTEDType"); +static_assert(KITTENS_FP8E4M3 == static_cast(transformer_engine::DType::kFloat8E4M3), "KittensDType out of sync with NVTEDType"); +static_assert(KITTENS_FP8E5M2 == static_cast(transformer_engine::DType::kFloat8E5M2), "KittensDType out of sync with NVTEDType"); +static_assert(KITTENS_BLOCK_SCALING_1D == NVTE_BLOCK_SCALING_1D, "KittensScalingMode out of sync with NVTEScalingMode"); +static_assert(KITTENS_BLOCK_SCALING_2D == NVTE_BLOCK_SCALING_2D, "KittensScalingMode out of sync with NVTEScalingMode"); #endif namespace transformer_engine { @@ -2071,17 +2077,13 @@ void cublas_gemm(const Tensor *inputA, const Tensor *inputB, Tensor *outputD, || inputB->scaling_mode == NVTE_MXFP8_1D_SCALING; #ifdef USE_HIPKITTENS_GEMM - bool use_hipkittens = false; -#ifdef KITTENS_HAVE_CDNA4 - if (is_mxfp8) { - bool is_gfx950 = (cuda::sm_arch() == 95); + if (is_mxfp8 && kittens_mxfp8_supported()) { bool force_hipblaslt = false; if (const char *env_p = std::getenv("NVTE_ROCM_USE_HIPBLASLT_MXFP8")) { force_hipblaslt = (strcmp(env_p, "1") == 0); } - use_hipkittens = is_gfx950 && !force_hipblaslt - && m % 256 == 0 && n % 256 == 0 && k % 128 == 0 && k >= 256; + use_hipkittens = !force_hipblaslt; } if (use_hipkittens) { @@ -2100,7 +2102,7 @@ void cublas_gemm(const Tensor *inputA, const Tensor *inputB, Tensor *outputD, beta, workspace, workspaceSize, gemm_stream); } -#endif + if (!use_hipkittens) { if (is_mxfp8) { NVTE_CHECK(inputBias->data.dptr == nullptr, @@ -2131,29 +2133,38 @@ void cublas_gemm(const Tensor *inputA, const Tensor *inputB, Tensor *outputD, #pragma GCC diagnostic pop #ifdef USE_HIPKITTENS_GEMM -bool try_kittens_grouped_mxfp8_gemm(const NVTETensor *A, const NVTETensor *B, NVTETensor *D, - int num_gemms, bool transa, bool transb, NVTETensor *workspace, - bool accumulate, cudaStream_t stream) { - if (accumulate || num_gemms <= 1) return false; +static bool kittens_env_set(const char *name) { + const char *v = std::getenv(name); + return v && v[0] == '1'; +} - auto env_set = [](const char *name) { - const char *v = std::getenv(name); - return v && v[0] == '1'; - }; - static bool enabled = [&] { - if (cuda::sm_arch() != 95) return false; - bool use = env_set("NVTE_USE_CUTLASS_GROUPED_GEMM"); - bool use_hk = env_set("NVTE_USE_HIPKITTENS_GROUPED_GEMM"); - bool use_ck = env_set("NVTE_USE_CK_GROUPED_GEMM"); +static bool kittens_grouped_mxfp8_enabled() { + static bool enabled = [] { + if (!kittens_mxfp8_supported()) return false; + bool use = kittens_env_set("NVTE_USE_CUTLASS_GROUPED_GEMM"); + bool use_hk = kittens_env_set("NVTE_USE_HIPKITTENS_GROUPED_GEMM"); + bool use_ck = kittens_env_set("NVTE_USE_CK_GROUPED_GEMM"); if (use_hk && use_ck) { fprintf(stderr, "[HK-grouped] both NVTE_USE_HIPKITTENS_GROUPED_GEMM and " "NVTE_USE_CK_GROUPED_GEMM set; defaulting to HipKittens\n"); } return use_hk || (use && !use_ck); }(); - if (!enabled) return false; + return enabled; +} - static bool warn_fallback = env_set("NVTE_CUTLASS_GROUPED_GEMM_WARN_FALLBACK"); +static bool kittens_grouped_warn_fallback() { + static bool warn = kittens_env_set("NVTE_CUTLASS_GROUPED_GEMM_WARN_FALLBACK"); + return warn; +} + +bool try_kittens_grouped_mxfp8_gemm(const NVTETensor *A, const NVTETensor *B, NVTETensor *D, + int num_gemms, bool transa, bool transb, NVTETensor *workspace, + bool accumulate, cudaStream_t stream) { + if (accumulate || num_gemms <= 1) return false; + if (!kittens_grouped_mxfp8_enabled()) return false; + + const bool warn_fallback = kittens_grouped_warn_fallback(); std::vector a_ptrs(num_gemms), b_ptrs(num_gemms), c_ptrs(num_gemms); std::vector sa_ptrs(num_gemms), sb_ptrs(num_gemms); @@ -2249,20 +2260,9 @@ bool try_kittens_grouped_mxfp8_wgrad(const NVTETensor *A, const NVTETensor *B, N if (transa || !transb) return false; if (num_gemms <= 1) return false; - auto env_set = [](const char *name) { - const char *v = std::getenv(name); - return v && v[0] == '1'; - }; - static bool enabled = [&] { - if (cuda::sm_arch() != 95) return false; - bool use = env_set("NVTE_USE_CUTLASS_GROUPED_GEMM"); - bool use_hk = env_set("NVTE_USE_HIPKITTENS_GROUPED_GEMM"); - bool use_ck = env_set("NVTE_USE_CK_GROUPED_GEMM"); - return use_hk || (use && !use_ck); - }(); - if (!enabled) return false; + if (!kittens_grouped_mxfp8_enabled()) return false; - static bool warn_fallback = env_set("NVTE_CUTLASS_GROUPED_GEMM_WARN_FALLBACK"); + const bool warn_fallback = kittens_grouped_warn_fallback(); std::vector a_ptrs(num_gemms), b_ptrs(num_gemms); std::vector d_ptrs(num_gemms);