diff --git a/mlx/backend/cuda/quantized/qmm/CMakeLists.txt b/mlx/backend/cuda/quantized/qmm/CMakeLists.txt index 3b88403e..0d682ead 100644 --- a/mlx/backend/cuda/quantized/qmm/CMakeLists.txt +++ b/mlx/backend/cuda/quantized/qmm/CMakeLists.txt @@ -1,19 +1,35 @@ target_sources( mlx - PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}/qmm.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmv.cu - ${CMAKE_CURRENT_SOURCE_DIR}/fp_qmv.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_naive_m16_k.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_naive_m16_n.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_naive_m32_k.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_naive_m32_n.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_naive_m64_k.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_naive_m64_n.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm80_m16.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm80_m32.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm80_m64.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm90_m128_n16_m1.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm90_m128_n32_m1.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm90_m128_n64_m2.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm90_m128_n128_m2.cu - ${CMAKE_CURRENT_SOURCE_DIR}/qmm_impl_sm90_m128_n256_m2.cu) + PRIVATE ${CMAKE_CURRENT_SOURCE_DIR}/qmm.cu ${CMAKE_CURRENT_SOURCE_DIR}/qmv.cu + ${CMAKE_CURRENT_SOURCE_DIR}/fp_qmv.cu) + +foreach(TileN 16 32 64 128 256) + set(OUTPUT_FILE "qmm_sm90_impl_n${TileN}.cu") + configure_file("${CMAKE_CURRENT_SOURCE_DIR}/qmm_sm90.cu" + "${CMAKE_CURRENT_BINARY_DIR}/${OUTPUT_FILE}" @ONLY) + target_sources(mlx PRIVATE ${CMAKE_CURRENT_BINARY_DIR}/${OUTPUT_FILE}) +endforeach() + +foreach(TileM 16 32 64) + set(OUTPUT_FILE "qmm_sm80_impl_m${TileM}.cu") + configure_file("${CMAKE_CURRENT_SOURCE_DIR}/qmm_sm80.cu" + "${CMAKE_CURRENT_BINARY_DIR}/${OUTPUT_FILE}" @ONLY) + target_sources(mlx PRIVATE ${CMAKE_CURRENT_BINARY_DIR}/${OUTPUT_FILE}) +endforeach() + +foreach(TileM 16 32 64) + foreach(KMajor true false) + foreach(HasKResidue true false) + foreach(SM80 true false) + if(${KMajor} AND ${HasKResidue}) + continue() + endif() + set(OUTPUT_FILE + "qmm_naive_impl_m${TileM}_${KMajor}_${HasKResidue}_${SM80}.cu") + configure_file("${CMAKE_CURRENT_SOURCE_DIR}/qmm_naive.cu" + "${CMAKE_CURRENT_BINARY_DIR}/${OUTPUT_FILE}" @ONLY) + target_sources(mlx PRIVATE ${CMAKE_CURRENT_BINARY_DIR}/${OUTPUT_FILE}) + endforeach() + endforeach() + endforeach() +endforeach() diff --git a/mlx/backend/cuda/quantized/qmm/qmm.cu b/mlx/backend/cuda/quantized/qmm/qmm.cu index 93982b08..fe3b791f 100644 --- a/mlx/backend/cuda/quantized/qmm/qmm.cu +++ b/mlx/backend/cuda/quantized/qmm/qmm.cu @@ -17,9 +17,9 @@ inline bool is_last_2_dims_row_contiguous(const array& x) { } // namespace #if defined(MLX_CUDA_SM90A_ENABLED) -// Defined in qmm_impl_sm90_xxx.cu files. -template -void qmm_impl_sm90( +// Defined in qmm_sm90.cu. +template +void qmm_sm90_impl( const array& x, const array& w, const array& scales, @@ -83,24 +83,21 @@ void qmm_sm90( cu::CommandEncoder& encoder, Stream s) { #if defined(MLX_CUDA_SM90A_ENABLED) - auto dispatch = [&]() { - using cute::Int; - using TileShapeMN = cute::Shape, Int>; - using ClusterShape = cute::Shape, Int<1>, Int<1>>; - qmm_impl_sm90( + auto dispatch = [&]() { + qmm_sm90_impl( x, w, scales, biases, out, bits, group_size, encoder, s); }; int m = out.ndim() > 1 ? out.shape(-2) : 1; if (m <= 16) { - dispatch.template operator()<128, 16, 1>(); + dispatch.template operator()<16>(); } else if (m <= 32) { - dispatch.template operator()<128, 32, 1>(); + dispatch.template operator()<32>(); } else if (m <= 64) { - dispatch.template operator()<128, 64, 2>(); + dispatch.template operator()<64>(); } else if (m <= 128) { - dispatch.template operator()<128, 128, 2>(); + dispatch.template operator()<128>(); } else { - dispatch.template operator()<128, 256, 2>(); + dispatch.template operator()<256>(); } #else throw std::runtime_error( @@ -108,9 +105,9 @@ void qmm_sm90( #endif // defined(MLX_CUDA_SM90A_ENABLED) } -// Defined in qmm_impl_sm80_xxx.cu files. +// Defined in qmm_sm80.cu. template -void qmm_impl_sm80( +void qmm_sm80_impl( const array& x, const array& w, const array& scales, @@ -174,7 +171,7 @@ void qmm_sm80( QuantizationMode mode, cu::CommandEncoder& encoder) { auto dispatch = [&]() { - qmm_impl_sm80( + qmm_sm80_impl( x, w, scales, @@ -197,9 +194,9 @@ void qmm_sm80( } } -// Defined in qmm_impl_naive_xxx.cu files. -template -void qmm_impl_naive( +// Defined in qmm_naive.cu. +template +void qmm_naive_impl( const array& x, const array& w, const array& scales, @@ -250,8 +247,8 @@ void qmm_naive( int group_size, QuantizationMode mode, cu::CommandEncoder& encoder) { - auto dispatch = [&]() { - qmm_impl_naive( + auto dispatch = [&]() { + qmm_naive_impl( x, w, scales, @@ -264,15 +261,37 @@ void qmm_naive( mode, encoder); }; - dispatch_bool(transpose, [&](auto k_major) { - int m = out.ndim() > 1 ? out.shape(-2) : 1; - if (m <= 16) { - dispatch.template operator()<16, k_major.value>(); - } else if (m <= 32) { - dispatch.template operator()<32, k_major.value>(); + auto dispatch_k = [&](auto k_major, bool has_k_residue, auto&& f) { + if constexpr (k_major.value) { + if (has_k_residue) { + throw std::invalid_argument( + "[quantized_matmul] K must be multiples of group_size."); + } + f.template operator()(); } else { - dispatch.template operator()<64, k_major.value>(); + dispatch_bool(has_k_residue, [&](auto has_k_residue) { + f.template operator()(); + }); } + }; + int m = out.ndim() > 1 ? out.shape(-2) : 1; + int k = x.shape(-1); + bool has_k_residue = k % group_size != 0; + bool sm80 = encoder.device().compute_capability_major() >= 8; + dispatch_bool(transpose, [&](auto k_major) { + dispatch_k(k_major, has_k_residue, [&]() { + dispatch_bool(sm80, [&](auto sm80) { + constexpr bool KMajor = k_major.value; + constexpr bool SM80 = sm80.value; + if (m <= 16) { + dispatch.template operator()<16, KMajor, HasKResidue, SM80>(); + } else if (m <= 32) { + dispatch.template operator()<32, KMajor, HasKResidue, SM80>(); + } else { + dispatch.template operator()<64, KMajor, HasKResidue, SM80>(); + } + }); + }); }); } diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m16_k.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m16_k.cu deleted file mode 100644 index 4bead82a..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m16_k.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh" - -QMM_NAIVE_GPU(16, true) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m16_n.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m16_n.cu deleted file mode 100644 index 993243d9..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m16_n.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh" - -QMM_NAIVE_GPU(16, false) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m32_k.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m32_k.cu deleted file mode 100644 index def1b4e7..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m32_k.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh" - -QMM_NAIVE_GPU(32, true) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m32_n.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m32_n.cu deleted file mode 100644 index bf1a500c..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m32_n.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh" - -QMM_NAIVE_GPU(32, false) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m64_k.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m64_k.cu deleted file mode 100644 index 92f03c78..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m64_k.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh" - -QMM_NAIVE_GPU(64, true) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m64_n.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m64_n.cu deleted file mode 100644 index 1d1f0400..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive_m64_n.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh" - -QMM_NAIVE_GPU(64, false) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m16.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m16.cu deleted file mode 100644 index cd682be3..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m16.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh" - -QMM_SM80_GPU(16) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m32.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m32.cu deleted file mode 100644 index 1f79364b..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m32.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh" - -QMM_SM80_GPU(32) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m64.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m64.cu deleted file mode 100644 index 9af44f8e..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80_m64.cu +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh" - -QMM_SM80_GPU(64) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n128_m2.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n128_m2.cu deleted file mode 100644 index 8db29f2a..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n128_m2.cu +++ /dev/null @@ -1,10 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh" - -using namespace cute; - -using TileShapeMN = Shape<_128, _128>; -using ClusterShape = Shape<_2, _1, _1>; - -QMM_SM90_GPU(TileShapeMN, ClusterShape) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n16_m1.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n16_m1.cu deleted file mode 100644 index ba1a62b1..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n16_m1.cu +++ /dev/null @@ -1,10 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh" - -using namespace cute; - -using TileShapeMN = Shape<_128, _16>; -using ClusterShape = Shape<_1, _1, _1>; - -QMM_SM90_GPU(TileShapeMN, ClusterShape) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n256_m2.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n256_m2.cu deleted file mode 100644 index 81f82c72..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n256_m2.cu +++ /dev/null @@ -1,10 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh" - -using namespace cute; - -using TileShapeMN = Shape<_128, _256>; -using ClusterShape = Shape<_2, _1, _1>; - -QMM_SM90_GPU(TileShapeMN, ClusterShape) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n32_m1.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n32_m1.cu deleted file mode 100644 index 0955296c..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n32_m1.cu +++ /dev/null @@ -1,10 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh" - -using namespace cute; - -using TileShapeMN = Shape<_128, _32>; -using ClusterShape = Shape<_1, _1, _1>; - -QMM_SM90_GPU(TileShapeMN, ClusterShape) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n64_m2.cu b/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n64_m2.cu deleted file mode 100644 index 89a2643b..00000000 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90_m128_n64_m2.cu +++ /dev/null @@ -1,10 +0,0 @@ -// Copyright © 2026 Apple Inc. - -#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh" - -using namespace cute; - -using TileShapeMN = Shape<_128, _64>; -using ClusterShape = Shape<_2, _1, _1>; - -QMM_SM90_GPU(TileShapeMN, ClusterShape) diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh b/mlx/backend/cuda/quantized/qmm/qmm_naive.cu similarity index 84% rename from mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh rename to mlx/backend/cuda/quantized/qmm/qmm_naive.cu index 9207171c..5be75bd0 100644 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh +++ b/mlx/backend/cuda/quantized/qmm/qmm_naive.cu @@ -316,7 +316,7 @@ inline constexpr auto make_scales_layout(auto n, auto k, auto l, auto group_size } } -template void qmm_naive( const Element* A, @@ -396,21 +396,6 @@ void qmm_naive( namespace mlx::core { -template -inline void dispatch_k(bool has_k_residue, const char* tag, F&& f) { - if constexpr (KMajor) { - if (has_k_residue) { - throw std::invalid_argument( - fmt::format("{} K must be multiples of group_size.", tag)); - } - f.template operator()(); - } else { - dispatch_bool(has_k_residue, [&](auto has_k_residue) { - f.template operator()(); - }); - } -} - template inline void dispatch_element_types(Dtype dtype, const char* tag, F&& f) { if (dtype == float32) { @@ -474,8 +459,8 @@ inline void dispatch_quant_types( } } -template -void qmm_impl_naive( +template +void qmm_naive_impl( const array& x, const array& w, const array& scales, @@ -494,71 +479,65 @@ void qmm_impl_naive( int l = out.size() / (m * n); bool broadcast_b = (w.ndim() <= 2) || (w.size() != w.data_size()); - bool is_sm80 = encoder.device().compute_capability_major() >= 8; - dispatch_bool(is_sm80, [&](auto sm80) { - dispatch_k(k % group_size != 0, tag, [&]() { - dispatch_element_types(out.dtype(), tag, [&]() { - dispatch_quant_types( - bits, - group_size, - mode, - tag, - [&]() { - encoder.set_input_array(x); - encoder.set_input_array(w); - encoder.set_input_array(scales); - if (biases) { - encoder.set_input_array(*biases); - } - if (lhs_indices) { - encoder.set_input_array(*lhs_indices); - } - if (rhs_indices) { - encoder.set_input_array(*rhs_indices); - } - encoder.set_output_array(out); - cutlass_gemm::qmm_naive( - gpu_ptr(x), - gpu_ptr(w), - gpu_ptr(scales), - biases ? gpu_ptr(*biases) : nullptr, - lhs_indices ? gpu_ptr(*lhs_indices) : nullptr, - rhs_indices ? gpu_ptr(*rhs_indices) : nullptr, - gpu_ptr(out), - m, - n, - k, - l, - broadcast_b, - cute::Int{}, - [&](auto* kernel, - dim3 num_blocks, - dim3 block_dims, - uint32_t smem_bytes, - void** args) { - encoder.add_kernel_node_raw( - kernel, num_blocks, block_dims, {}, smem_bytes, args); - }); - }); - }); - }); + dispatch_element_types(out.dtype(), tag, [&]() { + dispatch_quant_types( + bits, + group_size, + mode, + tag, + [&]() { + encoder.set_input_array(x); + encoder.set_input_array(w); + encoder.set_input_array(scales); + if (biases) { + encoder.set_input_array(*biases); + } + if (lhs_indices) { + encoder.set_input_array(*lhs_indices); + } + if (rhs_indices) { + encoder.set_input_array(*rhs_indices); + } + encoder.set_output_array(out); + cutlass_gemm::qmm_naive( + gpu_ptr(x), + gpu_ptr(w), + gpu_ptr(scales), + biases ? gpu_ptr(*biases) : nullptr, + lhs_indices ? gpu_ptr(*lhs_indices) : nullptr, + rhs_indices ? gpu_ptr(*rhs_indices) : nullptr, + gpu_ptr(out), + m, + n, + k, + l, + broadcast_b, + cute::Int{}, + [&](auto* kernel, + dim3 num_blocks, + dim3 block_dims, + uint32_t smem_bytes, + void** args) { + encoder.add_kernel_node_raw( + kernel, num_blocks, block_dims, {}, smem_bytes, args); + }); + }); }); } -} // namespace mlx::core +// clang-format off +template void qmm_naive_impl<@TileM@, @KMajor@, @HasKResidue@, @SM80@>( + const array& x, + const array& w, + const array& scales, + const std::optional& biases, + const std::optional& lhs_indices, + const std::optional& rhs_indices, + array& out, + int bits, + int group_size, + QuantizationMode mode, + cu::CommandEncoder& encoder); +// clang-format on -#define QMM_NAIVE_GPU(TileM, KMajor) \ - namespace mlx::core { \ - template void qmm_impl_naive( \ - const array& x, \ - const array& w, \ - const array& scales, \ - const std::optional& biases, \ - const std::optional& lhs_indices, \ - const std::optional& rhs_indices, \ - array& out, \ - int bits, \ - int group_size, \ - QuantizationMode mode, \ - cu::CommandEncoder& encoder); \ - } +} // namespace mlx::core diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh b/mlx/backend/cuda/quantized/qmm/qmm_sm80.cu similarity index 96% rename from mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh rename to mlx/backend/cuda/quantized/qmm/qmm_sm80.cu index 302d5bd9..028f18cd 100644 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh +++ b/mlx/backend/cuda/quantized/qmm/qmm_sm80.cu @@ -434,7 +434,7 @@ inline void dispatch_quant_types( } template -void qmm_impl_sm80( +void qmm_sm80_impl( const array& x, const array& w, const array& scales, @@ -499,20 +499,19 @@ void qmm_impl_sm80( }); } -} // namespace mlx::core +// clang-format off +template void qmm_sm80_impl<@TileM@>( + const array& x, + const array& w, + const array& scales, + const std::optional& biases, + const std::optional& lhs_indices, + const std::optional& rhs_indices, + array& out, + int bits, + int group_size, + QuantizationMode mode, + cu::CommandEncoder& encoder); +// clang-format on -#define QMM_SM80_GPU(TileM) \ - namespace mlx::core { \ - template void qmm_impl_sm80( \ - const array& x, \ - const array& w, \ - const array& scales, \ - const std::optional& biases, \ - const std::optional& lhs_indices, \ - const std::optional& rhs_indices, \ - array& out, \ - int bits, \ - int group_size, \ - QuantizationMode mode, \ - cu::CommandEncoder& encoder); \ - } +} // namespace mlx::core diff --git a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh b/mlx/backend/cuda/quantized/qmm/qmm_sm90.cu similarity index 87% rename from mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh rename to mlx/backend/cuda/quantized/qmm/qmm_sm90.cu index be552a66..e9425d83 100644 --- a/mlx/backend/cuda/quantized/qmm/qmm_impl_sm90.cuh +++ b/mlx/backend/cuda/quantized/qmm/qmm_sm90.cu @@ -20,8 +20,7 @@ namespace cutlass_gemm { using namespace cute; template < - typename TileShapeMN = Shape<_128, _16>, - typename ClusterShape = Shape<_1, _1, _1>, + int TileN = 16, typename Element, typename Quant, typename GroupSize, @@ -47,7 +46,8 @@ void qmm_sm90( using Arch = cutlass::arch::Sm90; using Accumulator = float; - using TileShape = decltype(append(TileShapeMN{}, Int{})); + using TileShape = Shape<_128, Int, Int>; + using ClusterShape = Shape, _1, _1>; using Epilogue = typename cutlass::epilogue::collective::CollectiveBuilder< Arch, @@ -177,8 +177,8 @@ inline void dispatch_groups(int group_size, const char* tag, F&& f) { } } -template -void qmm_impl_sm90( +template +void qmm_sm90_impl( const array& x, const array& w, const array& scales_, @@ -207,7 +207,7 @@ void qmm_impl_sm90( encoder.set_input_array(scales); encoder.set_input_array(biases); encoder.set_output_array(out); - cutlass_gemm::qmm_sm90( + cutlass_gemm::qmm_sm90( gpu_ptr(x), gpu_ptr(w), gpu_ptr(scales), @@ -238,24 +238,19 @@ void qmm_impl_sm90( }); } +// clang-format off +template void qmm_sm90_impl<@TileN@>( + const array& x, + const array& w, + const array& scales, + const array& biases, + array& out, + int bits, + int group_size, + cu::CommandEncoder& encoder, + Stream s); +// clang-format on + } // namespace mlx::core -#define QMM_SM90_GPU(TileShapeMN, ClusterShape) \ - namespace mlx::core { \ - template void qmm_impl_sm90( \ - const array& x, \ - const array& w, \ - const array& scales, \ - const array& biases, \ - array& out, \ - int bits, \ - int group_size, \ - cu::CommandEncoder& encoder, \ - Stream s); \ - } - -#else - -#define QMM_SM90_GPU(TileShapeMN, ClusterShape) - #endif // defined(MLX_CUDA_SM90A_ENABLED)