diff --git a/mlx/backend/cuda/cutlass_utils.cuh b/mlx/backend/cuda/cutlass_utils.cuh new file mode 100644 index 00000000..45e1c166 --- /dev/null +++ b/mlx/backend/cuda/cutlass_utils.cuh @@ -0,0 +1,45 @@ +// Copyright © 2025 Apple Inc. + +#pragma once + +#include "mlx/dtype.h" + +#include +#include +#include + +namespace mlx::core { + +// Throw exception if the cutlass API does not succeed. +inline void check_cutlass_error(const char* name, cutlass::Status status) { + if (status != cutlass::Status::kSuccess) { + throw std::runtime_error(fmt::format( + "{} failed with code: {}.", + name, + cutlass::cutlassGetStatusString(status))); + } +} + +// The macro version that prints the command that failed. +#define CHECK_CUTLASS_ERROR(cmd) check_cutlass_error(#cmd, (cmd)) + +// Maps CPU types to CUTLASS types. +template +struct CTypeToCutlassType { + using type = T; +}; + +template <> +struct CTypeToCutlassType { + using type = cutlass::half_t; +}; + +template <> +struct CTypeToCutlassType { + using type = cutlass::bfloat16_t; +}; + +template +using cutlass_type_t = typename CTypeToCutlassType::type; + +} // namespace mlx::core diff --git a/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu b/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu index 0fe60dbd..b00f9b2e 100644 --- a/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu +++ b/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu @@ -1,5 +1,7 @@ // Copyright © 2025 Apple Inc. +#include "mlx/backend/cuda/cublas_utils.h" +#include "mlx/backend/cuda/cutlass_utils.cuh" #include "mlx/backend/cuda/device.h" #include "mlx/backend/cuda/gemms/grouped_gemm.h" #include "mlx/backend/cuda/kernel_utils.cuh" @@ -9,7 +11,6 @@ #include #include #include -#include #include namespace mlx::core { @@ -96,7 +97,78 @@ __global__ void prepare_grouped_mm_data( namespace { -template +// Shared GEMM configuration for every type and arch. +template +struct CommonGemmConfiguration { + using Element = T; + using Arch = ArchTag; + using Accumulator = std::conditional_t<(sizeof(T) < 4), float, T>; + using EpilogueOutputOp = cutlass::epilogue::thread:: + LinearCombination; +}; + +// Slow GEMM configuration as fallback. +template < + typename T, + typename Arch, + int kAlignmentC = 1, + bool kEnableTF32 = false, + typename Enable = void> +struct GemmConfiguration : public CommonGemmConfiguration { + using OpClass = cutlass::arch::OpClassSimt; + using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 8>; + using WarpShape = cutlass::gemm::GemmShape<32, 64, 8>; + using InstructionShape = cutlass::gemm::GemmShape<1, 1, 1>; + static const int kAlignmentAB = 1; + static const int kStages = 2; +}; + +// Specialized GEMM configuration for sm80 and later. +template +struct GemmConfiguration< + T, + Arch, + kAlignmentC, + kEnableTF32, + std::enable_if_t= 80 && sizeof(T) <= 4>> + : public CommonGemmConfiguration { + using OpClass = cutlass::arch::OpClassTensorOp; + using ThreadblockShape = cutlass::gemm::GemmShape<256, 128, 32>; + using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>; + using InstructionShape = cutlass::gemm::GemmShape<16, 8, 32 / sizeof(T)>; + static const int kAlignmentAB = 1; + static const int kStages = 2; +}; + +// Specialized GEMM configuration for tf32 on sm80. +template +struct GemmConfiguration + : public CommonGemmConfiguration { + using OpClass = cutlass::arch::OpClassTensorOp; + using ThreadblockShape = cutlass::gemm::GemmShape<256, 128, 32>; + using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>; + using InstructionShape = cutlass::gemm::GemmShape<16, 8, 8>; + static const int kAlignmentAB = 1; + static const int kStages = 3; // use SM80_CP_ASYNC +}; + +// Get direct access to kernel. +template +class GemmGroupedEncoder + : public cutlass::gemm::device::GemmGrouped { + public: + void encode(cu::CommandEncoder& encoder) { + encoder.add_kernel_node( + cutlass::Kernel, + {static_cast(this->params_.threadblock_count), 1, 1}, + {GemmKernel::kThreadCount, 1, 1}, + sizeof(typename GemmKernel::SharedStorage), + this->params_); + } +}; + +// Invoke the grouped GEMM of CUTLASS 2.x API, which supports small alignments. +template void grouped_gemm_v2( bool a_transposed, bool b_transposed, @@ -109,11 +181,6 @@ void grouped_gemm_v2( void* b_ptrs, void* out_ptrs, cu::CommandEncoder& encoder) { - using ElementAccumulator = float; - using GemmConfiguration = typename cutlass::gemm::device:: - DefaultGemmConfiguration; - using EpilogueOutputOp = typename GemmConfiguration::EpilogueOutputOp; - dispatch_bool(a_transposed, [&](auto a_transposed_tag) { dispatch_bool(b_transposed, [&](auto b_transposed_tag) { using LayoutA = std::conditional_t< @@ -125,63 +192,81 @@ void grouped_gemm_v2( cutlass::layout::ColumnMajor, cutlass::layout::RowMajor>; using GemmKernel = typename cutlass::gemm::kernel::DefaultGemmGrouped< - T, + typename GemmConfiguration::Element, LayoutA, cutlass::ComplexTransform::kNone, - kAlignment, - T, + GemmConfiguration::kAlignmentAB, + typename GemmConfiguration::Element, LayoutB, cutlass::ComplexTransform::kNone, - kAlignment, - T, + GemmConfiguration::kAlignmentAB, + typename GemmConfiguration::Element, cutlass::layout::RowMajor, - ElementAccumulator, - OpClass, - Arch, + typename GemmConfiguration::Accumulator, + typename GemmConfiguration::OpClass, + typename GemmConfiguration::Arch, typename GemmConfiguration::ThreadblockShape, typename GemmConfiguration::WarpShape, typename GemmConfiguration::InstructionShape, - EpilogueOutputOp, + typename GemmConfiguration::EpilogueOutputOp, cutlass::gemm::threadblock::GemmBatchedIdentityThreadblockSwizzle, GemmConfiguration::kStages>::GemmKernel; - using GemmGrouped = - typename cutlass::gemm::device::GemmGrouped; + using GemmGrouped = GemmGroupedEncoder; - typename EpilogueOutputOp::Params epilogue_op( - /* alpha */ 1, /* beta */ 0); + static int threadblock_count = GemmGrouped::sufficient(); typename GemmGrouped::Arguments args( problem_sizes, group_count, - GemmGrouped::sufficient(), - epilogue_op, - reinterpret_cast(a_ptrs), - reinterpret_cast(b_ptrs), - reinterpret_cast(out_ptrs), - reinterpret_cast(out_ptrs), + threadblock_count, + {/* alpha */ 1, /* beta */ 0}, + reinterpret_cast(a_ptrs), + reinterpret_cast(b_ptrs), + reinterpret_cast(out_ptrs), + reinterpret_cast(out_ptrs), a_lds, b_lds, out_lds, out_lds); GemmGrouped gemm; - cutlass::Status status = gemm.initialize(args, nullptr, encoder.stream()); - if (status != cutlass::Status::kSuccess) { - throw std::runtime_error(fmt::format( - "Failed to initialize GemmGrouped: {}", - cutlass::cutlassGetStatusString(status))); - } - - auto capture = encoder.capture_context(); - status = gemm.run(encoder.stream()); - if (status != cutlass::Status::kSuccess) { - throw std::runtime_error(fmt::format( - "Failed to run GemmGrouped: {}", - cutlass::cutlassGetStatusString(status))); - } + CHECK_CUTLASS_ERROR(gemm.initialize( + args, + allocate_workspace(encoder, gemm.get_workspace_size(args)), + encoder.stream())); + gemm.encode(encoder); }); }); } +template +void dispatch_cutlass_arch(cu::Device& device, F&& f) { + if (device.compute_capability_major() < 8) { + f(type_identity{}); + } else if (device.compute_capability_major() == 8) { + f(type_identity{}); + } else { + f(type_identity{}); + } +} + +auto* get_grouped_mm_funcion(Dtype dtype, int N, cu::Device& device) { + auto* fun = grouped_gemm_v2>; + dispatch_float_types(dtype, "grouped_gemm_v2", [&](auto type_tag) { + using DataType = cutlass_type_t; + dispatch_cutlass_arch(device, [&](auto arch_tag) { + using Arch = MLX_GET_TYPE(arch_tag); + dispatch_bool(N % 8 == 0, [&](auto is_out_aligned) { + constexpr int kAlignmentC = is_out_aligned ? 8 : 1; + dispatch_bool(env::enable_tf32(), [&](auto kEnableTF32) { + fun = grouped_gemm_v2< + GemmConfiguration>; + }); + }); + }); + }); + return fun; +} + } // namespace void cutlass_grouped_gemm_unaligned( @@ -195,6 +280,9 @@ void cutlass_grouped_gemm_unaligned( const array& indices, array& out, cu::CommandEncoder& encoder) { + int K = a.shape(-1); + int N = b.shape(-1); + // Prepare device pointers for matmul. int problem_sizes_nbytes = group_count * cuda::ceil_div(sizeof(ProblemSize), 8) * 8; @@ -229,8 +317,8 @@ void cutlass_grouped_gemm_unaligned( gpu_ptr(indices), indices.size(), group_count, - a.shape(-1), // K - b.shape(-1), // N, + K, + N, lda, ldb, out.itemsize(), @@ -249,29 +337,11 @@ void cutlass_grouped_gemm_unaligned( out_ptrs); // Invoke grouped GEMM. - constexpr int kAlignment = 1; - using Arch = cutlass::arch::Sm75; - using OpClass = cutlass::arch::OpClassSimt; - auto* fun = grouped_gemm_v2; - switch (a.dtype()) { - case float32: - break; - case float16: - fun = grouped_gemm_v2; - break; - case bfloat16: - fun = grouped_gemm_v2; - break; - default: - throw std::runtime_error(fmt::format( - "Unsupported dtype in cutlass_grouped_gemm_sm75: {}.", - dtype_to_string(a.dtype()))); - } - encoder.set_input_array(a); encoder.set_input_array(b); encoder.set_input_array(gemm_args); encoder.set_output_array(out); + auto* fun = get_grouped_mm_funcion(a.dtype(), N, encoder.device()); fun(a_transposed, b_transposed, group_count,