Generate qmm implementaions with cmake (#3424)

This commit is contained in:
Cheng
2026-04-22 13:11:55 +09:00
committed by GitHub
parent 68cf2fddd8
commit b9b1bfb9a5
19 changed files with 176 additions and 263 deletions
+33 -17
View File
@@ -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()
+47 -28
View File
@@ -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 <typename TileShape, typename ClusterShape>
void qmm_impl_sm90(
// Defined in qmm_sm90.cu.
template <int TileN>
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 = [&]<int tile_m, int tile_n, int cluster_m>() {
using cute::Int;
using TileShapeMN = cute::Shape<Int<tile_m>, Int<tile_n>>;
using ClusterShape = cute::Shape<Int<cluster_m>, Int<1>, Int<1>>;
qmm_impl_sm90<TileShapeMN, ClusterShape>(
auto dispatch = [&]<int TileN>() {
qmm_sm90_impl<TileN>(
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 <int TileM>
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 = [&]<int TileM>() {
qmm_impl_sm80<TileM>(
qmm_sm80_impl<TileM>(
x,
w,
scales,
@@ -197,9 +194,9 @@ void qmm_sm80(
}
}
// Defined in qmm_impl_naive_xxx.cu files.
template <int TileM, bool KMajor>
void qmm_impl_naive(
// Defined in qmm_naive.cu.
template <int TileM, bool KMajor, bool HasKResidue, bool SM80>
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 = [&]<int TileM, bool KMajor>() {
qmm_impl_naive<TileM, KMajor>(
auto dispatch = [&]<int TileM, bool KMajor, bool HasKResidue, bool SM80>() {
qmm_naive_impl<TileM, KMajor, HasKResidue, SM80>(
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()<false>();
} else {
dispatch.template operator()<64, k_major.value>();
dispatch_bool(has_k_residue, [&](auto has_k_residue) {
f.template operator()<has_k_residue.value>();
});
}
};
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, [&]<bool HasKResidue>() {
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>();
}
});
});
});
}
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh"
QMM_NAIVE_GPU(16, true)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh"
QMM_NAIVE_GPU(16, false)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh"
QMM_NAIVE_GPU(32, true)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh"
QMM_NAIVE_GPU(32, false)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh"
QMM_NAIVE_GPU(64, true)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_naive.cuh"
QMM_NAIVE_GPU(64, false)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh"
QMM_SM80_GPU(16)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh"
QMM_SM80_GPU(32)
@@ -1,5 +0,0 @@
// Copyright © 2026 Apple Inc.
#include "mlx/backend/cuda/quantized/qmm/qmm_impl_sm80.cuh"
QMM_SM80_GPU(64)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -316,7 +316,7 @@ inline constexpr auto make_scales_layout(auto n, auto k, auto l, auto group_size
}
}
template <int TileM = 16, bool KMajor = true, bool SM80 = true, bool HasKResidue = false,
template <int TileM = 16, bool KMajor = true, bool HasKResidue = false, bool SM80 = true,
typename Element, typename Quant, typename Scale>
void qmm_naive(
const Element* A,
@@ -396,21 +396,6 @@ void qmm_naive(
namespace mlx::core {
template <bool KMajor, typename F>
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()<false>();
} else {
dispatch_bool(has_k_residue, [&](auto has_k_residue) {
f.template operator()<has_k_residue.value>();
});
}
}
template <typename F>
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 <int TileM, bool KMajor>
void qmm_impl_naive(
template <int TileM, bool KMajor, bool HasKResidue, bool SM80>
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<KMajor>(k % group_size != 0, tag, [&]<bool has_k_residue>() {
dispatch_element_types(out.dtype(), tag, [&]<typename Element>() {
dispatch_quant_types<Element>(
bits,
group_size,
mode,
tag,
[&]<typename Quant, typename Scale, int group_size>() {
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<TileM, KMajor, sm80.value, has_k_residue>(
gpu_ptr<Element>(x),
gpu_ptr<Quant>(w),
gpu_ptr<Scale>(scales),
biases ? gpu_ptr<Element>(*biases) : nullptr,
lhs_indices ? gpu_ptr<uint32_t>(*lhs_indices) : nullptr,
rhs_indices ? gpu_ptr<uint32_t>(*rhs_indices) : nullptr,
gpu_ptr<Element>(out),
m,
n,
k,
l,
broadcast_b,
cute::Int<group_size>{},
[&](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, [&]<typename Element>() {
dispatch_quant_types<Element>(
bits,
group_size,
mode,
tag,
[&]<typename Quant, typename Scale, int group_size>() {
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<TileM, KMajor, HasKResidue, SM80>(
gpu_ptr<Element>(x),
gpu_ptr<Quant>(w),
gpu_ptr<Scale>(scales),
biases ? gpu_ptr<Element>(*biases) : nullptr,
lhs_indices ? gpu_ptr<uint32_t>(*lhs_indices) : nullptr,
rhs_indices ? gpu_ptr<uint32_t>(*rhs_indices) : nullptr,
gpu_ptr<Element>(out),
m,
n,
k,
l,
broadcast_b,
cute::Int<group_size>{},
[&](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<array>& biases,
const std::optional<array>& lhs_indices,
const std::optional<array>& 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<TileM, KMajor>( \
const array& x, \
const array& w, \
const array& scales, \
const std::optional<array>& biases, \
const std::optional<array>& lhs_indices, \
const std::optional<array>& rhs_indices, \
array& out, \
int bits, \
int group_size, \
QuantizationMode mode, \
cu::CommandEncoder& encoder); \
}
} // namespace mlx::core
@@ -434,7 +434,7 @@ inline void dispatch_quant_types(
}
template <int TileM>
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<array>& biases,
const std::optional<array>& lhs_indices,
const std::optional<array>& 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<TileM>( \
const array& x, \
const array& w, \
const array& scales, \
const std::optional<array>& biases, \
const std::optional<array>& lhs_indices, \
const std::optional<array>& rhs_indices, \
array& out, \
int bits, \
int group_size, \
QuantizationMode mode, \
cu::CommandEncoder& encoder); \
}
} // namespace mlx::core
@@ -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<kTileShapeK>{}));
using TileShape = Shape<_128, Int<TileN>, Int<kTileShapeK>>;
using ClusterShape = Shape<Int<(TileN <= 32) ? 1 : 2>, _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 <typename TileShapeMN, typename ClusterShape>
void qmm_impl_sm90(
template <int TileN>
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<TileN>(
gpu_ptr<Element>(x),
gpu_ptr<Quant>(w),
gpu_ptr<Element>(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<TileShapeMN, ClusterShape>( \
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)