Generate qmm implementaions with cmake (#3424)
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
+61
-82
@@ -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
|
||||
+16
-17
@@ -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
|
||||
+19
-24
@@ -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)
|
||||
Reference in New Issue
Block a user