Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
165 changes: 165 additions & 0 deletions ggml/rocmfp4/rocmfp4_hip_scale.cuh
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
// ROCmFPx quant formats, hand-ported into this fork.
//
// Origin: https://github.com/charlie12345/ROCmFPX - creator of the ROCmFP4 format.
// Ported from: https://github.com/ciru-ai/ROCmFPX - a fork of the above.
//
// Both upstream projects are MIT licensed and based on llama.cpp; upstream authors
// retain their authorship and MIT license credit. See LICENSE.

#pragma once

#include <cstdint>
#include <cstring>
#include <cmath>

#ifndef GGML_ROCMFP4_USE_SCALE_LUT
#define GGML_ROCMFP4_USE_SCALE_LUT 0
#endif

#if defined(GGML_USE_HIP) && GGML_ROCMFP4_USE_SCALE_LUT
#define ROCMFP4_SCALE_SUB(M) ((M) * 0x1p-10f)
#define ROCMFP4_SCALE_E1(M) ((8 + (M)) * 0x1p-10f)
#define ROCMFP4_SCALE_E2(M) ((8 + (M)) * 0x1p-9f)
#define ROCMFP4_SCALE_E3(M) ((8 + (M)) * 0x1p-8f)
#define ROCMFP4_SCALE_E4(M) ((8 + (M)) * 0x1p-7f)
#define ROCMFP4_SCALE_E5(M) ((8 + (M)) * 0x1p-6f)
#define ROCMFP4_SCALE_E6(M) ((8 + (M)) * 0x1p-5f)
#define ROCMFP4_SCALE_E7(M) ((8 + (M)) * 0x1p-4f)
#define ROCMFP4_SCALE_E8(M) ((8 + (M)) * 0x1p-3f)
#define ROCMFP4_SCALE_E9(M) ((8 + (M)) * 0x1p-2f)
#define ROCMFP4_SCALE_E10(M) ((8 + (M)) * 0x1p-1f)
#define ROCMFP4_SCALE_E11(M) ((8 + (M)) * 0x1p0f)
#define ROCMFP4_SCALE_E12(M) ((8 + (M)) * 0x1p1f)
#define ROCMFP4_SCALE_E13(M) ((8 + (M)) * 0x1p2f)
#define ROCMFP4_SCALE_E14(M) ((8 + (M)) * 0x1p3f)
#define ROCMFP4_SCALE_E15(M) ((8 + (M)) * 0x1p4f)

static __device__ __constant__ const float rocmfp4_scale_ue4m3_half_lut[127] = {
ROCMFP4_SCALE_SUB(0), ROCMFP4_SCALE_SUB(1), ROCMFP4_SCALE_SUB(2), ROCMFP4_SCALE_SUB(3),
ROCMFP4_SCALE_SUB(4), ROCMFP4_SCALE_SUB(5), ROCMFP4_SCALE_SUB(6), ROCMFP4_SCALE_SUB(7),
ROCMFP4_SCALE_E1(0), ROCMFP4_SCALE_E1(1), ROCMFP4_SCALE_E1(2), ROCMFP4_SCALE_E1(3),
ROCMFP4_SCALE_E1(4), ROCMFP4_SCALE_E1(5), ROCMFP4_SCALE_E1(6), ROCMFP4_SCALE_E1(7),
ROCMFP4_SCALE_E2(0), ROCMFP4_SCALE_E2(1), ROCMFP4_SCALE_E2(2), ROCMFP4_SCALE_E2(3),
ROCMFP4_SCALE_E2(4), ROCMFP4_SCALE_E2(5), ROCMFP4_SCALE_E2(6), ROCMFP4_SCALE_E2(7),
ROCMFP4_SCALE_E3(0), ROCMFP4_SCALE_E3(1), ROCMFP4_SCALE_E3(2), ROCMFP4_SCALE_E3(3),
ROCMFP4_SCALE_E3(4), ROCMFP4_SCALE_E3(5), ROCMFP4_SCALE_E3(6), ROCMFP4_SCALE_E3(7),
ROCMFP4_SCALE_E4(0), ROCMFP4_SCALE_E4(1), ROCMFP4_SCALE_E4(2), ROCMFP4_SCALE_E4(3),
ROCMFP4_SCALE_E4(4), ROCMFP4_SCALE_E4(5), ROCMFP4_SCALE_E4(6), ROCMFP4_SCALE_E4(7),
ROCMFP4_SCALE_E5(0), ROCMFP4_SCALE_E5(1), ROCMFP4_SCALE_E5(2), ROCMFP4_SCALE_E5(3),
ROCMFP4_SCALE_E5(4), ROCMFP4_SCALE_E5(5), ROCMFP4_SCALE_E5(6), ROCMFP4_SCALE_E5(7),
ROCMFP4_SCALE_E6(0), ROCMFP4_SCALE_E6(1), ROCMFP4_SCALE_E6(2), ROCMFP4_SCALE_E6(3),
ROCMFP4_SCALE_E6(4), ROCMFP4_SCALE_E6(5), ROCMFP4_SCALE_E6(6), ROCMFP4_SCALE_E6(7),
ROCMFP4_SCALE_E7(0), ROCMFP4_SCALE_E7(1), ROCMFP4_SCALE_E7(2), ROCMFP4_SCALE_E7(3),
ROCMFP4_SCALE_E7(4), ROCMFP4_SCALE_E7(5), ROCMFP4_SCALE_E7(6), ROCMFP4_SCALE_E7(7),
ROCMFP4_SCALE_E8(0), ROCMFP4_SCALE_E8(1), ROCMFP4_SCALE_E8(2), ROCMFP4_SCALE_E8(3),
ROCMFP4_SCALE_E8(4), ROCMFP4_SCALE_E8(5), ROCMFP4_SCALE_E8(6), ROCMFP4_SCALE_E8(7),
ROCMFP4_SCALE_E9(0), ROCMFP4_SCALE_E9(1), ROCMFP4_SCALE_E9(2), ROCMFP4_SCALE_E9(3),
ROCMFP4_SCALE_E9(4), ROCMFP4_SCALE_E9(5), ROCMFP4_SCALE_E9(6), ROCMFP4_SCALE_E9(7),
ROCMFP4_SCALE_E10(0), ROCMFP4_SCALE_E10(1), ROCMFP4_SCALE_E10(2), ROCMFP4_SCALE_E10(3),
ROCMFP4_SCALE_E10(4), ROCMFP4_SCALE_E10(5), ROCMFP4_SCALE_E10(6), ROCMFP4_SCALE_E10(7),
ROCMFP4_SCALE_E11(0), ROCMFP4_SCALE_E11(1), ROCMFP4_SCALE_E11(2), ROCMFP4_SCALE_E11(3),
ROCMFP4_SCALE_E11(4), ROCMFP4_SCALE_E11(5), ROCMFP4_SCALE_E11(6), ROCMFP4_SCALE_E11(7),
ROCMFP4_SCALE_E12(0), ROCMFP4_SCALE_E12(1), ROCMFP4_SCALE_E12(2), ROCMFP4_SCALE_E12(3),
ROCMFP4_SCALE_E12(4), ROCMFP4_SCALE_E12(5), ROCMFP4_SCALE_E12(6), ROCMFP4_SCALE_E12(7),
ROCMFP4_SCALE_E13(0), ROCMFP4_SCALE_E13(1), ROCMFP4_SCALE_E13(2), ROCMFP4_SCALE_E13(3),
ROCMFP4_SCALE_E13(4), ROCMFP4_SCALE_E13(5), ROCMFP4_SCALE_E13(6), ROCMFP4_SCALE_E13(7),
ROCMFP4_SCALE_E14(0), ROCMFP4_SCALE_E14(1), ROCMFP4_SCALE_E14(2), ROCMFP4_SCALE_E14(3),
ROCMFP4_SCALE_E14(4), ROCMFP4_SCALE_E14(5), ROCMFP4_SCALE_E14(6), ROCMFP4_SCALE_E14(7),
ROCMFP4_SCALE_E15(0), ROCMFP4_SCALE_E15(1), ROCMFP4_SCALE_E15(2), ROCMFP4_SCALE_E15(3),
ROCMFP4_SCALE_E15(4), ROCMFP4_SCALE_E15(5), ROCMFP4_SCALE_E15(6),
};

#undef ROCMFP4_SCALE_SUB
#undef ROCMFP4_SCALE_E1
#undef ROCMFP4_SCALE_E2
#undef ROCMFP4_SCALE_E3
#undef ROCMFP4_SCALE_E4
#undef ROCMFP4_SCALE_E5
#undef ROCMFP4_SCALE_E6
#undef ROCMFP4_SCALE_E7
#undef ROCMFP4_SCALE_E8
#undef ROCMFP4_SCALE_E9
#undef ROCMFP4_SCALE_E10
#undef ROCMFP4_SCALE_E11
#undef ROCMFP4_SCALE_E12
#undef ROCMFP4_SCALE_E13
#undef ROCMFP4_SCALE_E14
#undef ROCMFP4_SCALE_E15
#endif

static __device__ __forceinline__ float rocmfp4_u32_as_f32(uint32_t bits) {
#if defined(GGML_USE_HIP)
return __uint_as_float(bits);
#else
float result;
memcpy(&result, &bits, sizeof(float));
return result;
#endif
}

// ROCmFP4 validates scale bytes before backend execution, so HIP/ROCm hot
// paths can decode finite unsigned E4M3 half-scales directly without the
// generic FP8 NaN handling used by other formats.
static __device__ __forceinline__ float rocmfp4_ue4m3_to_fp32_half_finite(uint8_t x) {
#if defined(GGML_USE_HIP) && GGML_ROCMFP4_USE_SCALE_LUT
return x <= 0x7e ? rocmfp4_scale_ue4m3_half_lut[x] : 0.0f;
#else
const int exp = (x >> 3) & 0xF;
const int man = x & 0x7;

if (exp == 0) {
return (float) man * (1.0f / 1024.0f);
}

const uint32_t bits = ((uint32_t) exp + 119u) << 23 | ((uint32_t) man << 20);
return rocmfp4_u32_as_f32(bits);
#endif
}

static __device__ __forceinline__ float rocmfpx_ue4m3_to_fp32_finite(uint8_t x) {
if (x > 0x7e) {
return 0.0f;
}

const int exp = (x >> 3) & 0xF;
const int man = x & 0x7;

if (exp == 0) {
return (float) man * (1.0f / 1024.0f);
}

const uint32_t bits = ((uint32_t) exp + 119u) << 23 | ((uint32_t) man << 20);
return rocmfp4_u32_as_f32(bits);
}

static __device__ __forceinline__ uint8_t rocmfpx_nearest_scale_ue4m3_cuda(float target_scale) {
if (!(target_scale > 0.0f) || !isfinite(target_scale)) {
return 0;
}

uint8_t lo = 1;
uint8_t hi = 0x7e;
while (lo < hi) {
const uint8_t mid = lo + (hi - lo) / 2;
if (rocmfpx_ue4m3_to_fp32_finite(mid) < target_scale) {
lo = mid + 1;
} else {
hi = mid;
}
}

if (lo == 1) {
return 1;
}

const float hi_scale = rocmfpx_ue4m3_to_fp32_finite(lo);
const float lo_scale = rocmfpx_ue4m3_to_fp32_finite((uint8_t) (lo - 1));
return (target_scale - lo_scale <= hi_scale - target_scale) ? (uint8_t) (lo - 1) : lo;
}

static __device__ __forceinline__ int8_t rocmfp4_decode_i8(uint8_t q) {
q &= 0x0f;
const int mag3 = q & 0x07;
const int mag = mag3 <= 4 ? mag3 : 2*mag3 - 4;
return (q & 0x08) ? -mag : mag;
}
18 changes: 18 additions & 0 deletions ggml/src/ggml-cuda/common.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
#endif
#endif
#include "ggml-common.h"
#include "../../rocmfp4/rocmfp4.h"

#include <array>
#include <algorithm>
Expand All @@ -40,6 +41,9 @@
#include "vendors/cuda.h"
#endif // defined(GGML_USE_HIP)

// device-side UE4M3 scale decoder for ROCmFP4 (needs the runtime header above)
#include "../../rocmfp4/rocmfp4_hip_scale.cuh"

#define STRINGIZE_IMPL(...) #__VA_ARGS__
#define STRINGIZE(...) STRINGIZE_IMPL(__VA_ARGS__)

Expand Down Expand Up @@ -1086,6 +1090,20 @@ struct ggml_cuda_type_traits<GGML_TYPE_NVFP4> {
static constexpr int bs = sizeof(block_nvfp4);
};

template<>
struct ggml_cuda_type_traits<GGML_TYPE_Q4_0_ROCMFP4> {
static constexpr int qk = QK_ROCMFP4;
static constexpr int qr = QR_ROCMFP4;
static constexpr int qi = QI_ROCMFP4;
};

template<>
struct ggml_cuda_type_traits<GGML_TYPE_Q4_0_ROCMFP4_FAST> {
static constexpr int qk = QK_ROCMFP4;
static constexpr int qr = QR_ROCMFP4;
static constexpr int qi = QI_ROCMFP4;
};

template<>
struct ggml_cuda_type_traits<GGML_TYPE_Q2_K> {
static constexpr int qk = QK_K;
Expand Down
38 changes: 38 additions & 0 deletions ggml/src/ggml-cuda/convert.cu
Original file line number Diff line number Diff line change
Expand Up @@ -243,6 +243,20 @@ static __global__ void dequantize_block_mxfp4(const void * __restrict__ vx, dst_
dequantize_mxfp4(vx, i, yy + i*QK_K, threadIdx.x);
}

template<typename dst_t>
static __global__ void dequantize_block_rocmfp4(const void * __restrict__ vx, dst_t * __restrict__ yy) {
const int64_t i = blockIdx.x;

dequantize_rocmfp4(vx, i, yy + i*QK_K, threadIdx.x);
}

template<typename dst_t>
static __global__ void dequantize_block_rocmfp4_fast(const void * __restrict__ vx, dst_t * __restrict__ yy) {
const int64_t i = blockIdx.x;

dequantize_rocmfp4_fast(vx, i, yy + i*QK_K, threadIdx.x);
}

template <int qk, int qr, dequantize_kernel_t dequantize_kernel, typename dst_t>
static void dequantize_block_cuda(const void * vx, dst_t * y,
const int64_t ne00, const int64_t ne01, const int64_t ne02, const int64_t ne03,
Expand Down Expand Up @@ -374,6 +388,18 @@ static void dequantize_row_mxfp4_cuda(const void * vx, dst_t * y, const int64_t
dequantize_block_mxfp4<<<nb, 32, 0, stream>>>(vx, y);
}

template<typename dst_t>
static void dequantize_row_rocmfp4_cuda(const void * vx, dst_t * y, const int64_t k, cudaStream_t stream) {
const int nb = (k + QK_K - 1) / QK_K;
dequantize_block_rocmfp4<<<nb, 32, 0, stream>>>(vx, y);
}

template<typename dst_t>
static void dequantize_row_rocmfp4_fast_cuda(const void * vx, dst_t * y, const int64_t k, cudaStream_t stream) {
const int nb = (k + QK_K - 1) / QK_K;
dequantize_block_rocmfp4_fast<<<nb, 32, 0, stream>>>(vx, y);
}

template <typename dst_t>
static __global__ void dequantize_block_nvfp4(
const void * __restrict__ vx,
Expand Down Expand Up @@ -533,6 +559,10 @@ to_bf16_cuda_t ggml_get_to_bf16_cuda(ggml_type type) {
return dequantize_row_iq3_s_cuda;
case GGML_TYPE_MXFP4:
return dequantize_row_mxfp4_cuda;
case GGML_TYPE_Q4_0_ROCMFP4:
return dequantize_row_rocmfp4_cuda;
case GGML_TYPE_Q4_0_ROCMFP4_FAST:
return dequantize_row_rocmfp4_fast_cuda;
case GGML_TYPE_NVFP4:
return dequantize_row_nvfp4_cuda;
case GGML_TYPE_F32:
Expand Down Expand Up @@ -593,6 +623,10 @@ to_fp16_cuda_t ggml_get_to_fp16_cuda(ggml_type type) {
return dequantize_row_iq3_s_cuda;
case GGML_TYPE_MXFP4:
return dequantize_row_mxfp4_cuda;
case GGML_TYPE_Q4_0_ROCMFP4:
return dequantize_row_rocmfp4_cuda;
case GGML_TYPE_Q4_0_ROCMFP4_FAST:
return dequantize_row_rocmfp4_fast_cuda;
case GGML_TYPE_NVFP4:
return dequantize_row_nvfp4_cuda;
case GGML_TYPE_F32:
Expand Down Expand Up @@ -650,6 +684,10 @@ to_fp32_cuda_t ggml_get_to_fp32_cuda(ggml_type type) {
return dequantize_row_iq3_s_cuda;
case GGML_TYPE_MXFP4:
return dequantize_row_mxfp4_cuda;
case GGML_TYPE_Q4_0_ROCMFP4:
return dequantize_row_rocmfp4_cuda;
case GGML_TYPE_Q4_0_ROCMFP4_FAST:
return dequantize_row_rocmfp4_fast_cuda;
case GGML_TYPE_NVFP4:
return dequantize_row_nvfp4_cuda;
case GGML_TYPE_F16:
Expand Down
34 changes: 34 additions & 0 deletions ggml/src/ggml-cuda/dequantize.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -525,3 +525,37 @@ template<typename dst_t>
static __device__ __forceinline__ void dequantize_mxfp4(const void * vx, const int64_t ib, dst_t * yy, const int tid) {
dequantize_mxfp4<dst_t, dst_t *>(vx, ib, yy, tid);
}

template<typename dst_t>
static __device__ __forceinline__ void dequantize_rocmfp4(const void * vx, const int64_t ibs, dst_t * yy, const int tid) {

const block_rocmfp4 * x = (const block_rocmfp4 *) vx + ibs*(QK_K/QK_ROCMFP4);

const int64_t il = tid/8; // 0...3
const int64_t ib = tid%8; // 0...7
dst_t * y = yy + 32*ib + 4*il;
const uint8_t * q4 = x[ib].qs + 4*il;
// dual UE4M3 scales: low nibbles are weights j (e[0]), high nibbles are weights j+16 (e[1])
const float d0 = rocmfp4_ue4m3_to_fp32_half_finite(x[ib].e[0]);
const float d1 = rocmfp4_ue4m3_to_fp32_half_finite(x[ib].e[1]);
for (int j = 0; j < 4; ++j) {
y[j+ 0] = ggml_cuda_cast<dst_t>(d0 * kvalues_rocmfp4[q4[j] & 0xf]);
y[j+16] = ggml_cuda_cast<dst_t>(d1 * kvalues_rocmfp4[q4[j] >> 4]);
}
}

template<typename dst_t>
static __device__ __forceinline__ void dequantize_rocmfp4_fast(const void * vx, const int64_t ibs, dst_t * yy, const int tid) {

const block_rocmfp4_fast * x = (const block_rocmfp4_fast *) vx + ibs*(QK_K/QK_ROCMFP4);

const int64_t il = tid/8; // 0...3
const int64_t ib = tid%8; // 0...7
dst_t * y = yy + 32*ib + 4*il;
const uint8_t * q4 = x[ib].qs + 4*il;
const float d = rocmfp4_ue4m3_to_fp32_half_finite(x[ib].e);
for (int j = 0; j < 4; ++j) {
y[j+ 0] = ggml_cuda_cast<dst_t>(d * kvalues_rocmfp4[q4[j] & 0xf]);
y[j+16] = ggml_cuda_cast<dst_t>(d * kvalues_rocmfp4[q4[j] >> 4]);
}
}
8 changes: 8 additions & 0 deletions ggml/src/ggml-cuda/getrows.cu
Original file line number Diff line number Diff line change
Expand Up @@ -470,6 +470,14 @@ static void ggml_cuda_get_rows_switch_src0_type(
get_rows_cuda_kq<32, dst_t, dequantize_mxfp4<dst_t>>(src0_d, src1_d, dst_d,
ne00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q4_0_ROCMFP4:
get_rows_cuda_kq<32, dst_t, dequantize_rocmfp4<dst_t>>(src0_d, src1_d, dst_d,
ne00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q4_0_ROCMFP4_FAST:
get_rows_cuda_kq<32, dst_t, dequantize_rocmfp4_fast<dst_t>>(src0_d, src1_d, dst_d,
ne00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb1, nb2, nb3, stream);
break;
default:
GGML_ABORT("%s: unsupported src0 type: %s\n", __func__, ggml_type_name(src0_type));
break;
Expand Down
4 changes: 4 additions & 0 deletions ggml/src/ggml-cuda/ggml-cuda.cu
Original file line number Diff line number Diff line change
Expand Up @@ -7152,6 +7152,8 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
case GGML_TYPE_Q8_0:
case GGML_TYPE_MXFP4:
case GGML_TYPE_NVFP4:
case GGML_TYPE_Q4_0_ROCMFP4:
case GGML_TYPE_Q4_0_ROCMFP4_FAST:
case GGML_TYPE_Q2_K:
case GGML_TYPE_Q3_K:
case GGML_TYPE_Q4_K:
Expand Down Expand Up @@ -7205,6 +7207,8 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
return true;
case GGML_TYPE_IQ4_NL:
case GGML_TYPE_MXFP4:
case GGML_TYPE_Q4_0_ROCMFP4:
case GGML_TYPE_Q4_0_ROCMFP4_FAST:
// 32-value sub-blocks, the row size does not guarantee
// the QK_K super-blocks the get_rows kernel iterates on
return op->src[0]->ne[0] % QK_K == 0;
Expand Down
16 changes: 16 additions & 0 deletions ggml/src/ggml-cuda/mmq-config-ampere.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -346,21 +346,37 @@ static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_conf
// ---------------------------------------------------------------------------------------------

CASE(GGML_TYPE_MXFP4, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_MXFP4, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);
CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 1, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, true, false);

CASE(GGML_TYPE_NVFP4, 256, 1, 128, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, true, true);
CASE(GGML_TYPE_NVFP4, 256, 1, 128, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, true, true);
Expand Down
Loading