[CUDA] Fix half type matmul in cutlass kernels (#3469)
This commit is contained in:
@@ -43,6 +43,12 @@ class GatherGemm {
|
||||
using ElementD = typename CollectiveEpilogue::ElementD;
|
||||
using StrideD = typename CollectiveEpilogue::StrideD;
|
||||
|
||||
static_assert(
|
||||
cute::is_same_v<
|
||||
ElementAccumulator,
|
||||
typename CollectiveEpilogue::ElementAccumulator>,
|
||||
"Mainloop and epilogue do not agree on accumulator value type.");
|
||||
|
||||
static constexpr int SharedStorageSize = static_cast<int>(cute::max(
|
||||
sizeof(typename CollectiveMainloop::SharedStorage),
|
||||
sizeof(typename CollectiveEpilogue::SharedStorage)));
|
||||
@@ -98,7 +104,9 @@ class GatherGemm {
|
||||
|
||||
CUTLASS_DEVICE void operator()(const Params& params, char* smem_buf) {
|
||||
int thread_idx = int(threadIdx.x);
|
||||
auto [m_coord, n_coord, l_coord] = uint3(blockIdx);
|
||||
int m_coord = int(blockIdx.x);
|
||||
int n_coord = int(blockIdx.y);
|
||||
int l_coord = int(blockIdx.z);
|
||||
|
||||
auto shape_MNKL = append<4>(params.problem_shape, Int<1>{});
|
||||
auto cta_tile = TileShape{};
|
||||
@@ -220,7 +228,7 @@ void gather_mm(
|
||||
using TileShape = Shape<_128, _128, _8>;
|
||||
using DispatchPolicy = cutlass::gemm::MainloopSm70TwoStage;
|
||||
using TiledMma = TiledMMA<
|
||||
MMA_Atom<UniversalFMA<Accumulator, Element, Element, Element>>,
|
||||
MMA_Atom<UniversalFMA<Accumulator, Element, Element, Accumulator>>,
|
||||
Layout<Shape<_16, _16, _1>>>;
|
||||
|
||||
using CopyTraitsA = SimtCopyTraits<Element, k_major_a.value>;
|
||||
@@ -296,9 +304,6 @@ void cutlass_gather_mm(
|
||||
int n = out.shape(-1);
|
||||
int k = a.shape(-1);
|
||||
int l = out.size() / (m * n);
|
||||
if (m < 16 || n < 16) {
|
||||
throw std::invalid_argument("[gather_mm] M/N is too small.");
|
||||
}
|
||||
|
||||
encoder.set_input_array(a);
|
||||
encoder.set_input_array(b);
|
||||
|
||||
@@ -245,7 +245,7 @@ void grouped_gemm_v2(
|
||||
LayoutB,
|
||||
cutlass::ComplexTransform::kNone,
|
||||
GemmConfiguration::kAlignmentAB,
|
||||
typename GemmConfiguration::Element,
|
||||
typename GemmConfiguration::Accumulator,
|
||||
cutlass::layout::RowMajor,
|
||||
typename GemmConfiguration::Accumulator,
|
||||
typename GemmConfiguration::OpClass,
|
||||
|
||||
@@ -52,7 +52,7 @@ bool supports_qmm_sm90(
|
||||
if (!biases) {
|
||||
return false;
|
||||
}
|
||||
if (!x.flags().row_contiguous || !is_last_2_dims_row_contiguous(w) ||
|
||||
if (!is_last_2_dims_row_contiguous(w) ||
|
||||
!is_last_2_dims_row_contiguous(scales) ||
|
||||
!is_last_2_dims_row_contiguous(*biases)) {
|
||||
return false;
|
||||
@@ -139,7 +139,7 @@ bool supports_qmm_sm80(
|
||||
if ((n % 128 != 0) || (k % std::max(64, group_size) != 0)) {
|
||||
return false;
|
||||
}
|
||||
if (!x.flags().row_contiguous || !is_last_2_dims_row_contiguous(w) ||
|
||||
if (!is_last_2_dims_row_contiguous(w) ||
|
||||
!is_last_2_dims_row_contiguous(scales)) {
|
||||
return false;
|
||||
}
|
||||
@@ -224,7 +224,7 @@ bool supports_qmm_naive(
|
||||
if (transpose && (k % std::max(64, group_size) != 0)) {
|
||||
return false;
|
||||
}
|
||||
if (!x.flags().row_contiguous || !is_last_2_dims_row_contiguous(w) ||
|
||||
if (!is_last_2_dims_row_contiguous(w) ||
|
||||
!is_last_2_dims_row_contiguous(scales)) {
|
||||
return false;
|
||||
}
|
||||
@@ -343,7 +343,7 @@ bool supports_qmv(
|
||||
if (k % 8 != 0) {
|
||||
return false;
|
||||
}
|
||||
if (!x.flags().row_contiguous || !is_last_2_dims_row_contiguous(w) ||
|
||||
if (!is_last_2_dims_row_contiguous(w) ||
|
||||
!is_last_2_dims_row_contiguous(scales)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -60,7 +60,9 @@ __global__ void qmm_naive_kernel(
|
||||
CUTE_STATIC_ASSERT_V(congruent(select<0,1,3>(shape_MNKL), dC));
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
auto [m_coord, n_coord, l_coord] = static_cast<uint3>(blockIdx);
|
||||
int m_coord = int(blockIdx.x);
|
||||
int n_coord = int(blockIdx.y);
|
||||
int l_coord = int(blockIdx.z);
|
||||
|
||||
auto m_max_coord = size<0>(shape_MNKL) - size<0>(cta_tiler) * m_coord; // M - BLK_M * m_coord
|
||||
auto n_max_coord = size<1>(shape_MNKL) - size<1>(cta_tiler) * n_coord; // N - BLK_N * n_coord
|
||||
|
||||
@@ -48,7 +48,9 @@ __global__ void qmm_sm80_kernel(
|
||||
CUTE_STATIC_ASSERT_V(congruent(select<0,1,3>(shape_MNKL), dC));
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
auto [m_coord, n_coord, l_coord] = static_cast<uint3>(blockIdx);
|
||||
int m_coord = int(blockIdx.x);
|
||||
int n_coord = int(blockIdx.y);
|
||||
int l_coord = int(blockIdx.z);
|
||||
|
||||
// For gather, use index lookup for input batch slicing.
|
||||
uint32_t a_batch = lhs_indices ? lhs_indices[l_coord] : l_coord;
|
||||
|
||||
@@ -17,7 +17,7 @@ void QuantizedMatmul::eval_gpu(const std::vector<array>& inputs, array& out) {
|
||||
auto& s = stream();
|
||||
auto& encoder = cu::get_command_encoder(s);
|
||||
|
||||
const array& x = inputs[0];
|
||||
array x = ensure_row_contiguous(inputs[0], encoder, s);
|
||||
const array& w = inputs[1];
|
||||
const array& scales = inputs[2];
|
||||
std::optional<array> biases;
|
||||
@@ -146,15 +146,17 @@ void GatherQMM::eval_gpu(const std::vector<array>& inputs, array& out) {
|
||||
auto& s = stream();
|
||||
auto& encoder = cu::get_command_encoder(s);
|
||||
|
||||
const array& x = inputs[0];
|
||||
array x = ensure_row_contiguous(inputs[0], encoder, s);
|
||||
const array& w = inputs[1];
|
||||
const array& scales = inputs[2];
|
||||
std::optional<array> biases;
|
||||
if (inputs.size() == 6) {
|
||||
biases = inputs[3];
|
||||
}
|
||||
array lhs_indices = ensure_contiguous(inputs[inputs.size() - 2], encoder, s);
|
||||
array rhs_indices = ensure_contiguous(inputs[inputs.size() - 1], encoder, s);
|
||||
array lhs_indices =
|
||||
ensure_row_contiguous(inputs[inputs.size() - 2], encoder, s);
|
||||
array rhs_indices =
|
||||
ensure_row_contiguous(inputs[inputs.size() - 1], encoder, s);
|
||||
|
||||
int M = out.ndim() > 1 ? out.shape(-2) : 1;
|
||||
int N = out.shape(-1);
|
||||
|
||||
@@ -1,22 +1,5 @@
|
||||
cuda_skip = {
|
||||
# Lapack ops NYI
|
||||
"TestLinalg.test_cholesky",
|
||||
"TestLinalg.test_cholesky_inv",
|
||||
"TestLinalg.test_eig",
|
||||
"TestLinalg.test_eigh",
|
||||
"TestLinalg.test_inverse",
|
||||
"TestVmap.test_vmap_inverse",
|
||||
"TestLinalg.test_lu",
|
||||
"TestLinalg.test_lu_factor",
|
||||
"TestLinalg.test_pseudo_inverse",
|
||||
"TestLinalg.test_qr_factorization",
|
||||
"TestInit.test_orthogonal",
|
||||
"TestLinalg.test_svd_decomposition",
|
||||
"TestVmap.test_vmap_svd",
|
||||
"TestLinalg.test_tri_inverse",
|
||||
# Quantization NYI
|
||||
"TestQuantized.test_gather_matmul_grad",
|
||||
"TestQuantized.test_gather_qmm",
|
||||
"TestQuantized.test_gather_qmm_sorted",
|
||||
"TestQuantized.test_gather_qmm_grad",
|
||||
}
|
||||
|
||||
@@ -1046,7 +1046,7 @@ class TestQuantized(mlx_tests.MLXTestCase):
|
||||
y3 = scatter_unsort(y3, inv_order, indices.shape)
|
||||
y4 = scatter_unsort(y4, inv_order, indices.shape)
|
||||
|
||||
tol = 1.5e-5 if (dtype == mx.float32) else 2.5e-4
|
||||
tol = 1.5e-5 if (dtype == mx.float32) else 1e-3
|
||||
|
||||
self.assertLess((y1 - y2).abs().max(), tol)
|
||||
self.assertLess((y1 - y3).abs().max(), tol)
|
||||
|
||||
Reference in New Issue
Block a user