diff --git a/ggml/rocmfp4/rocmfp4.c b/ggml/rocmfp4/rocmfp4.c deleted file mode 100644 index b731c4cedd1..00000000000 --- a/ggml/rocmfp4/rocmfp4.c +++ /dev/null @@ -1,679 +0,0 @@ -#define GGML_COMMON_DECL_C -#include "../src/ggml-common.h" - -#include "rocmfp4.h" - -#include -#include -#include - -// ROCmFP4 stores a signed integer FP4-like codebook at half-scale. It is -// E2M1-derived, but the largest magnitude is retuned from 12 to 10 after -// sampling Qwen3 dense tensors; this reduces outlier pull without changing the -// packed 4-bit layout or integer dot-product path. -static const int8_t rocmfp4_codebook[16] = { - 0, 1, 2, 3, 4, 6, 8, 10, - 0, -1, -2, -3, -4, -6, -8,-10, -}; - -static inline int8_t rocmfp4_decode(uint8_t q) { - q &= 0x0f; - const int mag3 = q & 0x07; - const int mag = mag3 <= 4 ? mag3 : 2*mag3 - 4; - return (q & 0x08) ? -mag : mag; -} - -static inline int8_t rocmfp4_decode_table(uint8_t q) { - return rocmfp4_codebook[q & 0x0f]; -} - -// Finite unsigned E4M3 scale bytes decoded to the half-scale values used by -// ROCmFP4. Keeping this as a table avoids rebuilding identical FP32 values for -// every candidate during exhaustive scale search. -#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 const float rocmfp4_scale_ue4m3_half[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 - -static inline float rocmfp4_ue4m3_to_fp32_half(uint8_t e) { - return e <= 0x7e ? rocmfp4_scale_ue4m3_half[e] : 0.0f; -} - -static inline uint8_t rocmfp4_best_index_scaled_finite(float x, float inv_scale_half) { - // Exact nearest-neighbor thresholds for Codebook10: - // 0, +/-1, +/-2, +/-3, +/-4, +/-6, +/-8, +/-10 - // Ties intentionally choose the lower-magnitude code, matching the former - // linear scan because the positive codes and zero appear first. - const float a = fabsf(x * inv_scale_half); - if (a <= 0.5f) { - return 0; - } - - const bool neg = x < 0.0f; - if (a <= 1.5f) { - return neg ? 9 : 1; - } - if (a <= 2.5f) { - return neg ? 10 : 2; - } - if (a <= 3.5f) { - return neg ? 11 : 3; - } - if (a <= 5.0f) { - return neg ? 12 : 4; - } - if (a <= 7.0f) { - return neg ? 13 : 5; - } - if (a <= 9.0f) { - return neg ? 14 : 6; - } - - return neg ? 15 : 7; -} - -static inline uint8_t rocmfp4_best_index_scaled(float x, float inv_scale_half) { - if (!isfinite(x)) { - return 0; - } - - return rocmfp4_best_index_scaled_finite(x, inv_scale_half); -} - -static inline bool rocmfp4_scale_is_valid(uint8_t e) { - // ROCmFP4 scale bytes are unsigned finite E4M3 values. 0x7f is NaN in the - // unsigned encoding and values with the sign bit set are not valid scales. - return e <= 0x7e; -} - -static float rocmfp4_block_mse_for_scale_unweighted( - const float * x, int n, int e, float best_err) { - const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); - const float inv_scale_half = 1.0f / scale_half; - float err = 0.0f; - - for (int i = 0; i < n; ++i) { - const uint8_t q = rocmfp4_best_index_scaled(x[i], inv_scale_half); - const float y = (float) rocmfp4_decode(q) * scale_half; - const float d = x[i] - y; - - err += d*d; - if (err > best_err) { - return err; - } - } - - return err; -} - -static float rocmfp4_block_mse_for_scale_unweighted_finite( - const float * x, int n, int e, float best_err) { - const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); - const float inv_scale_half = 1.0f / scale_half; - float err = 0.0f; - - for (int i = 0; i < n; ++i) { - const uint8_t q = rocmfp4_best_index_scaled_finite(x[i], inv_scale_half); - const float y = (float) rocmfp4_decode(q) * scale_half; - const float d = x[i] - y; - - err += d*d; - if (err > best_err) { - return err; - } - } - - return err; -} - -static float rocmfp4_block_mse_for_scale_weighted( - const float * x, int n, const float * mse_weights, int e, float best_err) { - const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); - const float inv_scale_half = 1.0f / scale_half; - float err = 0.0f; - - for (int i = 0; i < n; ++i) { - const uint8_t q = rocmfp4_best_index_scaled(x[i], inv_scale_half); - const float y = (float) rocmfp4_decode(q) * scale_half; - const float d = x[i] - y; - - err += mse_weights[i]*d*d; - if (err > best_err) { - return err; - } - } - - return err; -} - -static float rocmfp4_block_mse_for_scale_weighted_finite( - const float * x, int n, const float * mse_weights, int e, float best_err) { - const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); - const float inv_scale_half = 1.0f / scale_half; - float err = 0.0f; - - for (int i = 0; i < n; ++i) { - const uint8_t q = rocmfp4_best_index_scaled_finite(x[i], inv_scale_half); - const float y = (float) rocmfp4_decode(q) * scale_half; - const float d = x[i] - y; - - err += mse_weights[i]*d*d; - if (err > best_err) { - return err; - } - } - - return err; -} - -static void rocmfp4_prepare_mse_weights( - float * dst, const float * x, int n, const float * quant_weights, float sigma2, - float * max_abs, float * max_abs_weight, bool * all_finite) { - *max_abs = 0.0f; - *max_abs_weight = 0.0f; - *all_finite = true; - - for (int i = 0; i < n; ++i) { - const float qw = quant_weights[i]; - const float ax = fabsf(x[i]); - const float weight = isfinite(qw) && qw > 0.0f ? qw * sqrtf(sigma2 + x[i]*x[i]) : 0.0f; - *all_finite = *all_finite && isfinite(x[i]); - - if (ax > *max_abs) { - *max_abs = ax; - *max_abs_weight = weight; - } else if (ax == *max_abs && weight > *max_abs_weight) { - *max_abs_weight = weight; - } - - // Match llama.cpp's imatrix weighting style for Q4_0: calibration - // importance is scaled by row energy so large activations remain protected. - dst[i] = weight; - } -} - -static int rocmfp4_nearest_scale_ue4m3(float target_scale_half) { - if (!(target_scale_half > 0.0f) || !isfinite(target_scale_half)) { - return 1; - } - - int lo = 1; - int hi = 126; - while (lo < hi) { - const int mid = lo + (hi - lo) / 2; - if (rocmfp4_ue4m3_to_fp32_half((uint8_t) mid) < target_scale_half) { - lo = mid + 1; - } else { - hi = mid; - } - } - - if (lo == 1) { - return 1; - } - - const float hi_scale = rocmfp4_ue4m3_to_fp32_half((uint8_t) lo); - const float lo_scale = rocmfp4_ue4m3_to_fp32_half((uint8_t) (lo - 1)); - - // Match the former ascending nearest scan: exact midpoint ties keep the - // lower scale byte. - return (target_scale_half - lo_scale <= hi_scale - target_scale_half) ? lo - 1 : lo; -} - -static uint8_t rocmfp4_choose_scale_ue4m3_exhaustive_unweighted( - const float * x, int n, float max_abs, bool all_finite) { - const int start_e = rocmfp4_nearest_scale_ue4m3(max_abs / 10.0f); - - int best_e = 0; - float best_err = FLT_MAX; - bool lower_done = false; - - for (int delta = 0; delta <= 125; ++delta) { - const int e0 = start_e - delta; - if (!lower_done && e0 >= 1 && e0 <= 126) { - const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e0); - const float clip_delta = max_abs - 10.0f*scale_half; - if (clip_delta > 0.0f && clip_delta*clip_delta > best_err) { - lower_done = true; - } else { - const float err = all_finite ? - rocmfp4_block_mse_for_scale_unweighted_finite(x, n, e0, best_err) : - rocmfp4_block_mse_for_scale_unweighted(x, n, e0, best_err); - if (err < best_err || (err == best_err && e0 < best_e)) { - best_err = err; - best_e = e0; - } - } - } - - const int e1 = start_e + delta; - if (delta != 0 && e1 >= 1 && e1 <= 126) { - const float err = all_finite ? - rocmfp4_block_mse_for_scale_unweighted_finite(x, n, e1, best_err) : - rocmfp4_block_mse_for_scale_unweighted(x, n, e1, best_err); - if (err < best_err || (err == best_err && e1 < best_e)) { - best_err = err; - best_e = e1; - } - } - - if ((lower_done || e0 <= 1) && e1 >= 126) { - break; - } - } - - return (uint8_t) best_e; -} - -static uint8_t rocmfp4_choose_scale_ue4m3_exhaustive_weighted( - const float * x, int n, const float * mse_weights, float max_abs, float max_abs_weight, bool all_finite) { - const int start_e = rocmfp4_nearest_scale_ue4m3(max_abs / 10.0f); - - int best_e = 0; - float best_err = FLT_MAX; - bool lower_done = false; - - for (int delta = 0; delta <= 125; ++delta) { - const int e0 = start_e - delta; - if (!lower_done && e0 >= 1 && e0 <= 126) { - const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e0); - const float clip_delta = max_abs - 10.0f*scale_half; - if (max_abs_weight > 0.0f && clip_delta > 0.0f && max_abs_weight*clip_delta*clip_delta > best_err) { - lower_done = true; - } else { - const float err = all_finite ? - rocmfp4_block_mse_for_scale_weighted_finite(x, n, mse_weights, e0, best_err) : - rocmfp4_block_mse_for_scale_weighted(x, n, mse_weights, e0, best_err); - if (err < best_err || (err == best_err && e0 < best_e)) { - best_err = err; - best_e = e0; - } - } - } - - const int e1 = start_e + delta; - if (delta != 0 && e1 >= 1 && e1 <= 126) { - const float err = all_finite ? - rocmfp4_block_mse_for_scale_weighted_finite(x, n, mse_weights, e1, best_err) : - rocmfp4_block_mse_for_scale_weighted(x, n, mse_weights, e1, best_err); - if (err < best_err || (err == best_err && e1 < best_e)) { - best_err = err; - best_e = e1; - } - } - - if ((lower_done || e0 <= 1) && e1 >= 126) { - break; - } - } - - return (uint8_t) best_e; -} - -static uint8_t rocmfp4_choose_scale_ue4m3(const float * x, int n, const float * quant_weights, float sigma2) { - if (quant_weights) { - assert(n <= QK_ROCMFP4); - float mse_weights_buf[QK_ROCMFP4]; - float weighted_max_abs; - float max_abs_weight; - bool all_finite; - rocmfp4_prepare_mse_weights(mse_weights_buf, x, n, quant_weights, sigma2, &weighted_max_abs, &max_abs_weight, &all_finite); - if (!(weighted_max_abs > 0.0f) || !isfinite(weighted_max_abs)) { - return 0; - } - return rocmfp4_choose_scale_ue4m3_exhaustive_weighted(x, n, mse_weights_buf, weighted_max_abs, max_abs_weight, all_finite); - } - - float max_abs = 0.0f; - bool all_finite = true; - for (int i = 0; i < n; ++i) { - all_finite = all_finite && isfinite(x[i]); - const float ax = fabsf(x[i]); - if (ax > max_abs) { - max_abs = ax; - } - } - - if (!(max_abs > 0.0f) || !isfinite(max_abs)) { - return 0; - } - - return rocmfp4_choose_scale_ue4m3_exhaustive_unweighted(x, n, max_abs, all_finite); -} - -static void rocmfp4_quantize_row_q4_0_weighted( - const float * GGML_RESTRICT x, block_rocmfp4 * GGML_RESTRICT y, int64_t k, const float * GGML_RESTRICT quant_weights) { - assert(k % QK_ROCMFP4 == 0); - - float sum_x2 = 0.0f; - for (int64_t i = 0; i < k; ++i) { - sum_x2 += x[i]*x[i]; - } - const float sigma2 = sum_x2 / (float) k; - - const int64_t nb = k / QK_ROCMFP4; - for (int64_t ib = 0; ib < nb; ++ib) { - const float * xb = x + ib*QK_ROCMFP4; - const float * qw = quant_weights ? quant_weights + ib*QK_ROCMFP4 : NULL; - const uint8_t e0 = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4/2, qw, sigma2); - const uint8_t e1 = rocmfp4_choose_scale_ue4m3(xb + QK_ROCMFP4/2, QK_ROCMFP4/2, qw ? qw + QK_ROCMFP4/2 : NULL, sigma2); - const float scale_half0 = rocmfp4_ue4m3_to_fp32_half(e0); - const float scale_half1 = rocmfp4_ue4m3_to_fp32_half(e1); - const float inv_scale_half0 = scale_half0 > 0.0f ? 1.0f / scale_half0 : 0.0f; - const float inv_scale_half1 = scale_half1 > 0.0f ? 1.0f / scale_half1 : 0.0f; - - y[ib].e[0] = e0; - y[ib].e[1] = e1; - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half0); - const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half1); - y[ib].qs[j] = q0 | (q1 << 4); - } - } -} - -static void rocmfp4_quantize_row_q4_0_fast_weighted( - const float * GGML_RESTRICT x, block_rocmfp4_fast * GGML_RESTRICT y, int64_t k, const float * GGML_RESTRICT quant_weights) { - assert(k % QK_ROCMFP4 == 0); - - float sum_x2 = 0.0f; - for (int64_t i = 0; i < k; ++i) { - sum_x2 += x[i]*x[i]; - } - const float sigma2 = sum_x2 / (float) k; - - const int64_t nb = k / QK_ROCMFP4; - for (int64_t ib = 0; ib < nb; ++ib) { - const float * xb = x + ib*QK_ROCMFP4; - const float * qw = quant_weights ? quant_weights + ib*QK_ROCMFP4 : NULL; - const uint8_t e = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4, qw, sigma2); - const float scale_half = rocmfp4_ue4m3_to_fp32_half(e); - const float inv_scale_half = scale_half > 0.0f ? 1.0f / scale_half : 0.0f; - - y[ib].e = e; - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half); - const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half); - y[ib].qs[j] = q0 | (q1 << 4); - } - } -} - -void rocmfp4_quantize_row_q4_0_ref(const float * GGML_RESTRICT x, block_rocmfp4 * GGML_RESTRICT y, int64_t k) { - assert(k % QK_ROCMFP4 == 0); - - const int64_t nb = k / QK_ROCMFP4; - for (int64_t ib = 0; ib < nb; ++ib) { - const float * xb = x + ib*QK_ROCMFP4; - const uint8_t e0 = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4/2, NULL, 0.0f); - const uint8_t e1 = rocmfp4_choose_scale_ue4m3(xb + QK_ROCMFP4/2, QK_ROCMFP4/2, NULL, 0.0f); - const float scale_half0 = rocmfp4_ue4m3_to_fp32_half(e0); - const float scale_half1 = rocmfp4_ue4m3_to_fp32_half(e1); - const float inv_scale_half0 = scale_half0 > 0.0f ? 1.0f / scale_half0 : 0.0f; - const float inv_scale_half1 = scale_half1 > 0.0f ? 1.0f / scale_half1 : 0.0f; - - y[ib].e[0] = e0; - y[ib].e[1] = e1; - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half0); - const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half1); - y[ib].qs[j] = q0 | (q1 << 4); - } - } -} - -void rocmfp4_quantize_row_q4_0_fast_ref(const float * GGML_RESTRICT x, block_rocmfp4_fast * GGML_RESTRICT y, int64_t k) { - assert(k % QK_ROCMFP4 == 0); - - const int64_t nb = k / QK_ROCMFP4; - for (int64_t ib = 0; ib < nb; ++ib) { - const float * xb = x + ib*QK_ROCMFP4; - const uint8_t e = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4, NULL, 0.0f); - const float scale_half = rocmfp4_ue4m3_to_fp32_half(e); - const float inv_scale_half = scale_half > 0.0f ? 1.0f / scale_half : 0.0f; - - y[ib].e = e; - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half); - const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half); - y[ib].qs[j] = q0 | (q1 << 4); - } - } -} - -void rocmfp4_dequantize_row_q4_0(const block_rocmfp4 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { - assert(k % QK_ROCMFP4 == 0); - - const int64_t nb = k / QK_ROCMFP4; - for (int64_t ib = 0; ib < nb; ++ib) { - const float d0 = rocmfp4_ue4m3_to_fp32_half(x[ib].e[0]); - const float d1 = rocmfp4_ue4m3_to_fp32_half(x[ib].e[1]); - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - y[ib*QK_ROCMFP4 + j] = (float) rocmfp4_decode(x[ib].qs[j] & 0x0f) * d0; - y[ib*QK_ROCMFP4 + j + QK_ROCMFP4/2] = (float) rocmfp4_decode(x[ib].qs[j] >> 4) * d1; - } - } -} - -void rocmfp4_dequantize_row_q4_0_fast(const block_rocmfp4_fast * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { - assert(k % QK_ROCMFP4 == 0); - - const int64_t nb = k / QK_ROCMFP4; - for (int64_t ib = 0; ib < nb; ++ib) { - const float d = rocmfp4_ue4m3_to_fp32_half(x[ib].e); - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - y[ib*QK_ROCMFP4 + j] = (float) rocmfp4_decode(x[ib].qs[j] & 0x0f) * d; - y[ib*QK_ROCMFP4 + j + QK_ROCMFP4/2] = (float) rocmfp4_decode(x[ib].qs[j] >> 4) * d; - } - } -} - -void rocmfp4_quantize_row_q4_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { - rocmfp4_quantize_row_q4_0_ref(x, (block_rocmfp4 *) y, k); -} - -void rocmfp4_quantize_row_q4_0_fast(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { - rocmfp4_quantize_row_q4_0_fast_ref(x, (block_rocmfp4_fast *) y, k); -} - -size_t rocmfp4_quantize_q4_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix) { - const size_t row_size = ggml_row_size(GGML_TYPE_Q4_0_ROCMFP4, n_per_row); - - if (!imatrix) { - rocmfp4_quantize_row_q4_0_ref(src, (block_rocmfp4 *) dst, nrows*n_per_row); - return nrows * row_size; - } - - char * qrow = (char *) dst; - for (int64_t row = 0; row < nrows; ++row) { - rocmfp4_quantize_row_q4_0_weighted(src, (block_rocmfp4 *) qrow, n_per_row, imatrix); - src += n_per_row; - qrow += row_size; - } - - return nrows * row_size; -} - -size_t rocmfp4_quantize_q4_0_fast(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix) { - const size_t row_size = ggml_row_size(GGML_TYPE_Q4_0_ROCMFP4_FAST, n_per_row); - - if (!imatrix) { - rocmfp4_quantize_row_q4_0_fast_ref(src, (block_rocmfp4_fast *) dst, nrows*n_per_row); - return nrows * row_size; - } - - char * qrow = (char *) dst; - for (int64_t row = 0; row < nrows; ++row) { - rocmfp4_quantize_row_q4_0_fast_weighted(src, (block_rocmfp4_fast *) qrow, n_per_row, imatrix); - src += n_per_row; - qrow += row_size; - } - - return nrows * row_size; -} - -bool rocmfp4_validate_row_data(const void * data, size_t nbytes) { - if (nbytes % sizeof(block_rocmfp4) != 0) { - return false; - } - - const block_rocmfp4 * blocks = (const block_rocmfp4 *) data; - const size_t nblocks = nbytes / sizeof(block_rocmfp4); - for (size_t i = 0; i < nblocks; ++i) { - if (!rocmfp4_scale_is_valid(blocks[i].e[0]) || !rocmfp4_scale_is_valid(blocks[i].e[1])) { - return false; - } - } - - return true; -} - -bool rocmfp4_validate_row_data_fast(const void * data, size_t nbytes) { - if (nbytes % sizeof(block_rocmfp4_fast) != 0) { - return false; - } - - const block_rocmfp4_fast * blocks = (const block_rocmfp4_fast *) data; - const size_t nblocks = nbytes / sizeof(block_rocmfp4_fast); - for (size_t i = 0; i < nblocks; ++i) { - if (!rocmfp4_scale_is_valid(blocks[i].e)) { - return false; - } - } - - return true; -} - -void rocmfp4_vec_dot_q4_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { - GGML_UNUSED(bs); - GGML_UNUSED(bx); - GGML_UNUSED(by); - assert(nrc == 1); - GGML_UNUSED(nrc); - assert(n % QK_ROCMFP4 == 0); - assert(QK_ROCMFP4 == QK8_0); - - const block_rocmfp4 * GGML_RESTRICT x = (const block_rocmfp4 *) vx; - const block_q8_0 * GGML_RESTRICT y = (const block_q8_0 *) vy; - - const int nb = n / QK_ROCMFP4; - float sumf = 0.0f; - - for (int ib = 0; ib < nb; ++ib) { - const float d0 = rocmfp4_ue4m3_to_fp32_half(x[ib].e[0]) * ggml_fp16_to_fp32(y[ib].d); - const float d1 = rocmfp4_ue4m3_to_fp32_half(x[ib].e[1]) * ggml_fp16_to_fp32(y[ib].d); - int sumi0 = 0; - int sumi1 = 0; - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - const uint8_t q = x[ib].qs[j]; - sumi0 += rocmfp4_decode_table(q) * y[ib].qs[j]; - sumi1 += rocmfp4_decode_table(q >> 4) * y[ib].qs[j + QK_ROCMFP4/2]; - } - - sumf += d0 * (float) sumi0 + d1 * (float) sumi1; - } - - *s = sumf; -} - -void rocmfp4_vec_dot_q4_0_fast_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { - GGML_UNUSED(bs); - GGML_UNUSED(bx); - GGML_UNUSED(by); - assert(nrc == 1); - GGML_UNUSED(nrc); - assert(n % QK_ROCMFP4 == 0); - assert(QK_ROCMFP4 == QK8_0); - - const block_rocmfp4_fast * GGML_RESTRICT x = (const block_rocmfp4_fast *) vx; - const block_q8_0 * GGML_RESTRICT y = (const block_q8_0 *) vy; - - const int nb = n / QK_ROCMFP4; - float sumf = 0.0f; - - for (int ib = 0; ib < nb; ++ib) { - const float d = rocmfp4_ue4m3_to_fp32_half(x[ib].e) * ggml_fp16_to_fp32(y[ib].d); - int sumi = 0; - - for (int j = 0; j < QK_ROCMFP4/2; ++j) { - const uint8_t q = x[ib].qs[j]; - sumi += rocmfp4_decode_table(q) * y[ib].qs[j]; - sumi += rocmfp4_decode_table(q >> 4) * y[ib].qs[j + QK_ROCMFP4/2]; - } - - sumf += d * (float) sumi; - } - - *s = sumf; -} diff --git a/ggml/rocmfp4/rocmfp4.h b/ggml/rocmfp4/rocmfp4.h deleted file mode 100644 index 9756f6ad4c3..00000000000 --- a/ggml/rocmfp4/rocmfp4.h +++ /dev/null @@ -1,58 +0,0 @@ -#pragma once - -#include -#include -#include - -#include "ggml.h" - -#ifdef __cplusplus -extern "C" { -#endif - -#define QK_ROCMFP4 32 -#define QR_ROCMFP4 2 -#define QI_ROCMFP4 (QK_ROCMFP4 / (4 * QR_ROCMFP4)) -#define QS_ROCMFP4 32 - -// AMD-tuned compact layout: 16 bytes of packed E2M1-derived 4-bit codes, then -// one unsigned E4M3 scale byte per 16-weight half block. -typedef struct { - uint8_t qs[QK_ROCMFP4/2]; - uint8_t e[2]; -} block_rocmfp4; - -// Speed-focused layout: same 32 packed ROCmFP4 nibbles, but one UE4M3 scale -// for the whole block. This is a separate GGUF type so fast 4.25 BPW artifacts -// never alias the safer dual-scale format above. -typedef struct { - uint8_t qs[QK_ROCMFP4/2]; - uint8_t e; -} block_rocmfp4_fast; - -#if defined(__cplusplus) -static_assert(sizeof(block_rocmfp4) == QK_ROCMFP4/2 + 2*sizeof(uint8_t), "wrong rocmfp4 block size/padding"); -static_assert(sizeof(block_rocmfp4_fast) == QK_ROCMFP4/2 + sizeof(uint8_t), "wrong rocmfp4 fast block size/padding"); -#else -_Static_assert(sizeof(block_rocmfp4) == QK_ROCMFP4/2 + 2*sizeof(uint8_t), "wrong rocmfp4 block size/padding"); -_Static_assert(sizeof(block_rocmfp4_fast) == QK_ROCMFP4/2 + sizeof(uint8_t), "wrong rocmfp4 fast block size/padding"); -#endif - -GGML_API void rocmfp4_quantize_row_q4_0_ref(const float * GGML_RESTRICT x, block_rocmfp4 * GGML_RESTRICT y, int64_t k); -GGML_API void rocmfp4_dequantize_row_q4_0(const block_rocmfp4 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); -GGML_API void rocmfp4_quantize_row_q4_0_fast_ref(const float * GGML_RESTRICT x, block_rocmfp4_fast * GGML_RESTRICT y, int64_t k); -GGML_API void rocmfp4_dequantize_row_q4_0_fast(const block_rocmfp4_fast * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); - -GGML_API void rocmfp4_quantize_row_q4_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); -GGML_API size_t rocmfp4_quantize_q4_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); -GGML_API void rocmfp4_quantize_row_q4_0_fast(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); -GGML_API size_t rocmfp4_quantize_q4_0_fast(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); -GGML_API bool rocmfp4_validate_row_data(const void * data, size_t nbytes); -GGML_API bool rocmfp4_validate_row_data_fast(const void * data, size_t nbytes); - -GGML_API void rocmfp4_vec_dot_q4_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); -GGML_API void rocmfp4_vec_dot_q4_0_fast_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); - -#ifdef __cplusplus -} -#endif diff --git a/ggml/rocmfp4/rocmfp4_hip.cu b/ggml/rocmfp4/rocmfp4_hip.cu deleted file mode 100644 index d2c9048c187..00000000000 --- a/ggml/rocmfp4/rocmfp4_hip.cu +++ /dev/null @@ -1,85 +0,0 @@ -#include "rocmfp4.h" - -#include - -#include "rocmfp4_hip_scale.cuh" - -// Standalone ROCm/HIP dequant kernel for integration tests and future fused -// paths. One lane owns one packed byte and writes the matching low/high -// half-block values, so each byte is read once. -extern "C" __global__ void rocmfp4_dequantize_q4_0_f32_kernel( - const block_rocmfp4 * __restrict__ x, - float * __restrict__ y, - int64_t k) { - const int64_t packed_idx = (int64_t) blockIdx.x*blockDim.x + threadIdx.x; - const int64_t nblocks = (k + QK_ROCMFP4 - 1) / QK_ROCMFP4; - const int64_t packed_count = nblocks * (QK_ROCMFP4/2); - - if (packed_idx >= packed_count) { - return; - } - - const int64_t ib = packed_idx / (QK_ROCMFP4/2); - const int tid = packed_idx - ib*(QK_ROCMFP4/2); - const int64_t base = ib*QK_ROCMFP4; - const uint8_t packed = x[ib].qs[tid]; - 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]); - - if (base + tid < k) { - y[base + tid] = (float) rocmfp4_decode_i8(packed & 0x0f) * d0; - } - if (base + tid + QK_ROCMFP4/2 < k) { - y[base + tid + QK_ROCMFP4/2] = (float) rocmfp4_decode_i8(packed >> 4) * d1; - } -} - -extern "C" __global__ void rocmfp4_dequantize_q4_0_fast_f32_kernel( - const block_rocmfp4_fast * __restrict__ x, - float * __restrict__ y, - int64_t k) { - const int64_t packed_idx = (int64_t) blockIdx.x*blockDim.x + threadIdx.x; - const int64_t nblocks = (k + QK_ROCMFP4 - 1) / QK_ROCMFP4; - const int64_t packed_count = nblocks * (QK_ROCMFP4/2); - - if (packed_idx >= packed_count) { - return; - } - - const int64_t ib = packed_idx / (QK_ROCMFP4/2); - const int tid = packed_idx - ib*(QK_ROCMFP4/2); - const int64_t base = ib*QK_ROCMFP4; - const uint8_t packed = x[ib].qs[tid]; - const float d = rocmfp4_ue4m3_to_fp32_half_finite(x[ib].e); - - if (base + tid < k) { - y[base + tid] = (float) rocmfp4_decode_i8(packed & 0x0f) * d; - } - if (base + tid + QK_ROCMFP4/2 < k) { - y[base + tid + QK_ROCMFP4/2] = (float) rocmfp4_decode_i8(packed >> 4) * d; - } -} - -extern "C" void rocmfp4_hip_dequantize_q4_0_to_f32( - const void * src, - float * dst, - int64_t k, - hipStream_t stream) { - const int64_t nblocks = (k + QK_ROCMFP4 - 1) / QK_ROCMFP4; - const int64_t packed_count = nblocks * (QK_ROCMFP4/2); - const dim3 block(256); - const dim3 grid((unsigned int) ((packed_count + block.x - 1) / block.x)); - rocmfp4_dequantize_q4_0_f32_kernel<<>>((const block_rocmfp4 *) src, dst, k); -} - -extern "C" void rocmfp4_hip_dequantize_q4_0_fast_to_f32( - const void * src, - float * dst, - int64_t k, - hipStream_t stream) { - const int64_t nblocks = (k + QK_ROCMFP4 - 1) / QK_ROCMFP4; - const int64_t packed_count = nblocks * (QK_ROCMFP4/2); - const dim3 block(256); - const dim3 grid((unsigned int) ((packed_count + block.x - 1) / block.x)); - rocmfp4_dequantize_q4_0_fast_f32_kernel<<>>((const block_rocmfp4_fast *) src, dst, k); -} diff --git a/ggml/rocmfp4/rocmfp4_hip_codebook.cuh b/ggml/rocmfp4/rocmfp4_hip_codebook.cuh deleted file mode 100644 index 1d91061129a..00000000000 --- a/ggml/rocmfp4/rocmfp4_hip_codebook.cuh +++ /dev/null @@ -1,79 +0,0 @@ -#pragma once - -#include "rocmfp4_hip_scale.cuh" - -#include -#include - -#ifndef GGML_ROCMFP4_UNALIGNED_QS_DWORD_LOAD -#define GGML_ROCMFP4_UNALIGNED_QS_DWORD_LOAD 1 -#endif - -static __device__ __forceinline__ int rocmfp4_get_qs_i32(const void * x, const int & i32) { -#if defined(GGML_USE_HIP) && GGML_ROCMFP4_UNALIGNED_QS_DWORD_LOAD - return *((const int *) ((const uint8_t *) x + 4*i32)); -#else - const uint8_t * x8 = (const uint8_t *) x; - - int x32 = x8[4*i32 + 0] << 0; - x32 |= x8[4*i32 + 1] << 8; - x32 |= x8[4*i32 + 2] << 16; - x32 |= x8[4*i32 + 3] << 24; - - return x32; -#endif -} - -// AMD-specific fast path for expanding eight packed ROCmFP4 nibbles into two -// int32 DP4A operands. This encodes the Codebook10 table directly as four -// 32-bit constants: -// [0, 1, 2, 3], [4, 6, 8, 10], [0, -1, -2, -3], [-4, -6, -8, -10] -// Avoiding the table pointer keeps the ROCm/HIP MMVQ/MMQ hot path fully local -// to this format. Non-HIP builds still use llama.cpp's generic table expander. -static __device__ __forceinline__ int2 rocmfp4_get_int_from_codebook_16(const int & q4, const int8_t * fallback_table) { -#if defined(GGML_USE_HIP) - constexpr uint32_t values0 = 0x03020100u; - constexpr uint32_t values1 = 0x0a080604u; - constexpr uint32_t values2 = 0xfdfeff00u; - constexpr uint32_t values3 = 0xf6f8fafcu; - - const uint32_t q_even = q4; - const uint32_t q_odd = q4 >> 4; - - const uint32_t v_even_low = __builtin_amdgcn_perm(values1, values0, q_even & 0x07070707u); - const uint32_t v_odd_low = __builtin_amdgcn_perm(values1, values0, q_odd & 0x07070707u); - const uint32_t v_even_high = __builtin_amdgcn_perm(values3, values2, q_even & 0x07070707u); - const uint32_t v_odd_high = __builtin_amdgcn_perm(values3, values2, q_odd & 0x07070707u); - - const uint32_t mask_even = 0x03020100u | ((q_even & 0x08080808u) >> 1); - const uint32_t mask_odd = 0x03020100u | ((q_odd & 0x08080808u) >> 1); - - return make_int2( - __builtin_amdgcn_perm(v_even_high, v_even_low, mask_even), - __builtin_amdgcn_perm(v_odd_high, v_odd_low, mask_odd)); -#else - return get_int_from_table_16(q4, fallback_table); -#endif -} - -// Variant for call sites that already selected either the low or high nibble -// stream and only need one DP4A operand. This avoids the extra odd/even table -// expansion work in ROCmFP4 FlashAttention K/V decode. -static __device__ __forceinline__ int rocmfp4_get_low_int_from_codebook_16(const int & q4, const int8_t * fallback_table) { -#if defined(GGML_USE_HIP) - constexpr uint32_t values0 = 0x03020100u; - constexpr uint32_t values1 = 0x0a080604u; - constexpr uint32_t values2 = 0xfdfeff00u; - constexpr uint32_t values3 = 0xf6f8fafcu; - - const uint32_t q = q4; - - const uint32_t v_low = __builtin_amdgcn_perm(values1, values0, q & 0x07070707u); - const uint32_t v_high = __builtin_amdgcn_perm(values3, values2, q & 0x07070707u); - const uint32_t mask = 0x03020100u | ((q & 0x08080808u) >> 1); - - return __builtin_amdgcn_perm(v_high, v_low, mask); -#else - return get_int_from_table_16(q4, fallback_table).x; -#endif -} diff --git a/ggml/rocmfp4/rocmfp4_hip_scale.cuh b/ggml/rocmfp4/rocmfp4_hip_scale.cuh deleted file mode 100644 index 39fc80d4c11..00000000000 --- a/ggml/rocmfp4/rocmfp4_hip_scale.cuh +++ /dev/null @@ -1,115 +0,0 @@ -#pragma once - -#include -#include - -#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__ 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; -} diff --git a/ggml/src/CMakeLists.txt b/ggml/src/CMakeLists.txt index dc83c0a2621..82e9480c2f2 100644 --- a/ggml/src/CMakeLists.txt +++ b/ggml/src/CMakeLists.txt @@ -206,8 +206,6 @@ add_library(ggml-base ggml-threading.h ggml-quants.c ggml-quants.h - ../rocmfp4/rocmfp4.c - ../rocmfp4/rocmfp4.h gguf.cpp) set_target_properties(ggml-base PROPERTIES diff --git a/ggml/src/ggml-common.h b/ggml/src/ggml-common.h index 998195d1c87..b3360ea4ecd 100644 --- a/ggml/src/ggml-common.h +++ b/ggml/src/ggml-common.h @@ -112,6 +112,9 @@ typedef sycl::half2 ggml_half2; #define QI_NVFP4 (QK_NVFP4 / (4 * QR_NVFP4)) #define QR_NVFP4 2 +#define QI_ROCMFP4 (QK_ROCMFP4 / (4 * QR_ROCMFP4)) +#define QR_ROCMFP4 2 + #define QI5_0 (QK5_0 / (4 * QR5_0)) #define QR5_0 2 @@ -226,6 +229,23 @@ typedef struct { } block_nvfp4; static_assert(sizeof(block_nvfp4) == sizeof(uint8_t)*(QK_NVFP4/QK_NVFP4_SUB) + QK_NVFP4/2, "wrong nvfp4 block size/padding"); +#define QK_ROCMFP4 32 +// AMD-tuned compact layout: 16 bytes of packed E2M1-derived 4-bit codes, then +// one unsigned E4M3 scale byte per 16-weight half block. +typedef struct { + uint8_t qs[QK_ROCMFP4/2]; + uint8_t e[2]; +} block_rocmfp4; +static_assert(sizeof(block_rocmfp4) == QK_ROCMFP4/2 + 2*sizeof(uint8_t), "wrong rocmfp4 block size/padding"); + +// Speed-focused layout: same 32 packed ROCmFP4 nibbles, but one UE4M3 scale +// for the whole block. +typedef struct { + uint8_t qs[QK_ROCMFP4/2]; + uint8_t e; +} block_rocmfp4_fast; +static_assert(sizeof(block_rocmfp4_fast) == QK_ROCMFP4/2 + sizeof(uint8_t), "wrong rocmfp4 fast block size/padding"); + #define QK5_0 32 typedef struct { ggml_half d; // delta @@ -1136,6 +1156,47 @@ GGML_TABLE_BEGIN(int8_t, kvalues_rocmfp4, 16) 0, 1, 2, 3, 4, 6, 8, 10, 0, -1, -2, -3, -4, -6, -8, -10, GGML_TABLE_END() +// ROCmFP4 UE4M3 "half-scale" values for the finite scale bytes 0x00..0x7e (127 +// entries): the subnormal run (byte>>3 == 0, value M*2^-10) followed by 15 +// normal exponent groups (value (8+M)*2^(e-10), e=1..15). This is the single +// source of truth shared by both materializations of the table: the CPU +// quantizer scale-search table (ggml-quants.c) and the opt-in GPU +// constant-memory LUT (ggml-cuda/common.cuh). Each backend stamps its own +// storage-qualified array from this list so the two can never drift. +#define GGML_ROCMFP4_SCALE_UE4M3_HALF_LIST \ + (0) * 0x1p-10f, (1) * 0x1p-10f, (2) * 0x1p-10f, (3) * 0x1p-10f, \ + (4) * 0x1p-10f, (5) * 0x1p-10f, (6) * 0x1p-10f, (7) * 0x1p-10f, \ + (8 + 0) * 0x1p-10f, (8 + 1) * 0x1p-10f, (8 + 2) * 0x1p-10f, (8 + 3) * 0x1p-10f, \ + (8 + 4) * 0x1p-10f, (8 + 5) * 0x1p-10f, (8 + 6) * 0x1p-10f, (8 + 7) * 0x1p-10f, \ + (8 + 0) * 0x1p-9f, (8 + 1) * 0x1p-9f, (8 + 2) * 0x1p-9f, (8 + 3) * 0x1p-9f, \ + (8 + 4) * 0x1p-9f, (8 + 5) * 0x1p-9f, (8 + 6) * 0x1p-9f, (8 + 7) * 0x1p-9f, \ + (8 + 0) * 0x1p-8f, (8 + 1) * 0x1p-8f, (8 + 2) * 0x1p-8f, (8 + 3) * 0x1p-8f, \ + (8 + 4) * 0x1p-8f, (8 + 5) * 0x1p-8f, (8 + 6) * 0x1p-8f, (8 + 7) * 0x1p-8f, \ + (8 + 0) * 0x1p-7f, (8 + 1) * 0x1p-7f, (8 + 2) * 0x1p-7f, (8 + 3) * 0x1p-7f, \ + (8 + 4) * 0x1p-7f, (8 + 5) * 0x1p-7f, (8 + 6) * 0x1p-7f, (8 + 7) * 0x1p-7f, \ + (8 + 0) * 0x1p-6f, (8 + 1) * 0x1p-6f, (8 + 2) * 0x1p-6f, (8 + 3) * 0x1p-6f, \ + (8 + 4) * 0x1p-6f, (8 + 5) * 0x1p-6f, (8 + 6) * 0x1p-6f, (8 + 7) * 0x1p-6f, \ + (8 + 0) * 0x1p-5f, (8 + 1) * 0x1p-5f, (8 + 2) * 0x1p-5f, (8 + 3) * 0x1p-5f, \ + (8 + 4) * 0x1p-5f, (8 + 5) * 0x1p-5f, (8 + 6) * 0x1p-5f, (8 + 7) * 0x1p-5f, \ + (8 + 0) * 0x1p-4f, (8 + 1) * 0x1p-4f, (8 + 2) * 0x1p-4f, (8 + 3) * 0x1p-4f, \ + (8 + 4) * 0x1p-4f, (8 + 5) * 0x1p-4f, (8 + 6) * 0x1p-4f, (8 + 7) * 0x1p-4f, \ + (8 + 0) * 0x1p-3f, (8 + 1) * 0x1p-3f, (8 + 2) * 0x1p-3f, (8 + 3) * 0x1p-3f, \ + (8 + 4) * 0x1p-3f, (8 + 5) * 0x1p-3f, (8 + 6) * 0x1p-3f, (8 + 7) * 0x1p-3f, \ + (8 + 0) * 0x1p-2f, (8 + 1) * 0x1p-2f, (8 + 2) * 0x1p-2f, (8 + 3) * 0x1p-2f, \ + (8 + 4) * 0x1p-2f, (8 + 5) * 0x1p-2f, (8 + 6) * 0x1p-2f, (8 + 7) * 0x1p-2f, \ + (8 + 0) * 0x1p-1f, (8 + 1) * 0x1p-1f, (8 + 2) * 0x1p-1f, (8 + 3) * 0x1p-1f, \ + (8 + 4) * 0x1p-1f, (8 + 5) * 0x1p-1f, (8 + 6) * 0x1p-1f, (8 + 7) * 0x1p-1f, \ + (8 + 0) * 0x1p+0f, (8 + 1) * 0x1p+0f, (8 + 2) * 0x1p+0f, (8 + 3) * 0x1p+0f, \ + (8 + 4) * 0x1p+0f, (8 + 5) * 0x1p+0f, (8 + 6) * 0x1p+0f, (8 + 7) * 0x1p+0f, \ + (8 + 0) * 0x1p+1f, (8 + 1) * 0x1p+1f, (8 + 2) * 0x1p+1f, (8 + 3) * 0x1p+1f, \ + (8 + 4) * 0x1p+1f, (8 + 5) * 0x1p+1f, (8 + 6) * 0x1p+1f, (8 + 7) * 0x1p+1f, \ + (8 + 0) * 0x1p+2f, (8 + 1) * 0x1p+2f, (8 + 2) * 0x1p+2f, (8 + 3) * 0x1p+2f, \ + (8 + 4) * 0x1p+2f, (8 + 5) * 0x1p+2f, (8 + 6) * 0x1p+2f, (8 + 7) * 0x1p+2f, \ + (8 + 0) * 0x1p+3f, (8 + 1) * 0x1p+3f, (8 + 2) * 0x1p+3f, (8 + 3) * 0x1p+3f, \ + (8 + 4) * 0x1p+3f, (8 + 5) * 0x1p+3f, (8 + 6) * 0x1p+3f, (8 + 7) * 0x1p+3f, \ + (8 + 0) * 0x1p+4f, (8 + 1) * 0x1p+4f, (8 + 2) * 0x1p+4f, (8 + 3) * 0x1p+4f, \ + (8 + 4) * 0x1p+4f, (8 + 5) * 0x1p+4f, (8 + 6) * 0x1p+4f + #define NGRID_IQ1S 2048 #define IQ1S_DELTA 0.125f #define IQ1M_DELTA 0.125f diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 03552e7afa5..863eaf2248c 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -14,7 +14,6 @@ #include "ops.h" #include "ggml.h" #include "common.h" -#include "../../rocmfp4/rocmfp4.h" #if defined(_MSC_VER) || defined(__MINGW32__) #include // using malloc.h with MSC/MINGW diff --git a/ggml/src/ggml-cpu/quants.c b/ggml/src/ggml-cpu/quants.c index 5e36459f8cb..8ce16c03f93 100644 --- a/ggml/src/ggml-cpu/quants.c +++ b/ggml/src/ggml-cpu/quants.c @@ -62,6 +62,14 @@ void quantize_row_nvfp4(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, i quantize_row_nvfp4_ref(x, y, k); } +void rocmfp4_quantize_row_q4_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { + rocmfp4_quantize_row_q4_0_ref(x, (block_rocmfp4 *) y, k); +} + +void rocmfp4_quantize_row_q4_0_fast(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { + rocmfp4_quantize_row_q4_0_fast_ref(x, (block_rocmfp4_fast *) y, k); +} + // // 2-6 bit quantization in super-blocks // @@ -362,6 +370,71 @@ void ggml_vec_dot_nvfp4_q8_0_generic(int n, float * GGML_RESTRICT s, size_t bs, *s = sumf; } +// ROCmFP4: Q4_0-layout FP4 (QK_ROCMFP4 == QK8_0), two UE4M3 half-scales per block. +void rocmfp4_vec_dot_q4_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { + UNUSED(bs); + UNUSED(bx); + UNUSED(by); + assert(nrc == 1); + UNUSED(nrc); + assert(n % QK_ROCMFP4 == 0); + static_assert(QK_ROCMFP4 == QK8_0, "QK_ROCMFP4 and QK8_0 must be the same"); + + const block_rocmfp4 * GGML_RESTRICT x = vx; + const block_q8_0 * GGML_RESTRICT y = vy; + + const int nb = n / QK_ROCMFP4; + float sumf = 0; + + for (int ib = 0; ib < nb; ++ib) { + const float d0 = ggml_ue4m3_to_fp32(x[ib].e[0]) * GGML_CPU_FP16_TO_FP32(y[ib].d); + const float d1 = ggml_ue4m3_to_fp32(x[ib].e[1]) * GGML_CPU_FP16_TO_FP32(y[ib].d); + + int sumi0 = 0; + int sumi1 = 0; + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + const uint8_t q = x[ib].qs[j]; + sumi0 += kvalues_rocmfp4[q & 0x0f] * y[ib].qs[j]; + sumi1 += kvalues_rocmfp4[q >> 4] * y[ib].qs[j + QK_ROCMFP4/2]; + } + + sumf += d0 * (float) sumi0 + d1 * (float) sumi1; + } + + *s = sumf; +} + +void rocmfp4_vec_dot_q4_0_fast_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { + UNUSED(bs); + UNUSED(bx); + UNUSED(by); + assert(nrc == 1); + UNUSED(nrc); + assert(n % QK_ROCMFP4 == 0); + static_assert(QK_ROCMFP4 == QK8_0, "QK_ROCMFP4 and QK8_0 must be the same"); + + const block_rocmfp4_fast * GGML_RESTRICT x = vx; + const block_q8_0 * GGML_RESTRICT y = vy; + + const int nb = n / QK_ROCMFP4; + float sumf = 0; + + for (int ib = 0; ib < nb; ++ib) { + const float d = ggml_ue4m3_to_fp32(x[ib].e) * GGML_CPU_FP16_TO_FP32(y[ib].d); + int sumi = 0; + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + const uint8_t q = x[ib].qs[j]; + sumi += kvalues_rocmfp4[q & 0x0f] * y[ib].qs[j]; + sumi += kvalues_rocmfp4[q >> 4] * y[ib].qs[j + QK_ROCMFP4/2]; + } + + sumf += d * (float) sumi; + } + + *s = sumf; +} + void ggml_vec_dot_q5_0_q8_0_generic(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { const int qk = QK8_0; const int nb = n / qk; diff --git a/ggml/src/ggml-cpu/quants.h b/ggml/src/ggml-cpu/quants.h index 93ea7eeffe5..c89f9f69b8a 100644 --- a/ggml/src/ggml-cpu/quants.h +++ b/ggml/src/ggml-cpu/quants.h @@ -24,6 +24,10 @@ void quantize_row_q8_1(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, in void quantize_row_mxfp4(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void quantize_row_nvfp4(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); +// CPU from_float wrappers; defined in ggml-cpu/quants.c over the _ref quantizers in ggml-quants.c. +GGML_API void rocmfp4_quantize_row_q4_0 (const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfp4_quantize_row_q4_0_fast(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); + void quantize_row_q2_K(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void quantize_row_q3_K(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void quantize_row_q4_K(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); @@ -49,6 +53,9 @@ void ggml_vec_dot_q8_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const voi void ggml_vec_dot_mxfp4_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); void ggml_vec_dot_nvfp4_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); +GGML_API void rocmfp4_vec_dot_q4_0_q8_0 (int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); +GGML_API void rocmfp4_vec_dot_q4_0_fast_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); + void ggml_vec_dot_q2_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); void ggml_vec_dot_q3_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc); diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index 44b00a6ccd2..0e9349dc5e1 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -21,7 +21,6 @@ #endif #endif #include "ggml-common.h" -#include "../../rocmfp4/rocmfp4.h" #include #include @@ -867,6 +866,36 @@ static __device__ __forceinline__ float ggml_cuda_ue4m3_to_fp32(uint8_t x) { #endif // defined(GGML_USE_HIP) && defined(CDNA3) && defined(FP8_AVAILABLE) && HIP_VERSION >= 60200000 } +// ============================================================================ +// ROCmFP4 (AMD gfx1151) device scale + FP4 codebook decode helpers. +// GGML_ROCMFP4_USE_SCALE_LUT is an opt-in AMD-specific escape hatch (constant-memory table) kept for profiling. +// ============================================================================ + +#ifndef GGML_ROCMFP4_USE_SCALE_LUT +#define GGML_ROCMFP4_USE_SCALE_LUT 0 +#endif + +#if defined(GGML_USE_HIP) && GGML_ROCMFP4_USE_SCALE_LUT +// Values come from the shared GGML_ROCMFP4_SCALE_UE4M3_HALF_LIST in ggml-common.h +// (single source of truth, also used by the CPU quantizer table in ggml-quants.c). +static __device__ __constant__ const float rocmfp4_scale_ue4m3_half_lut[127] = { GGML_ROCMFP4_SCALE_UE4M3_HALF_LIST }; +#endif + +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; // opt-in fast table +#else + return ggml_cuda_ue4m3_to_fp32(x); // default: shared decoder +#endif +} + +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; +} + static __device__ __forceinline__ uint8_t ggml_cuda_fp32_to_ue4m3(float x) { #if defined(BLACKWELL_MMA_AVAILABLE) // This is used for NVFP4 subblock scale quantizations only if (!(x > 0.0f)) { diff --git a/ggml/src/ggml-cuda/convert.cu b/ggml/src/ggml-cuda/convert.cu index 8c61b826f4d..c4ea602f27c 100644 --- a/ggml/src/ggml-cuda/convert.cu +++ b/ggml/src/ggml-cuda/convert.cu @@ -1,6 +1,5 @@ #include "convert.cuh" #include "dequantize.cuh" -#include "../../rocmfp4/rocmfp4_hip_scale.cuh" #include diff --git a/ggml/src/ggml-cuda/dequantize.cuh b/ggml/src/ggml-cuda/dequantize.cuh index 05dac742f7c..f8601903ccf 100644 --- a/ggml/src/ggml-cuda/dequantize.cuh +++ b/ggml/src/ggml-cuda/dequantize.cuh @@ -1,6 +1,5 @@ #include "common.cuh" #include "convert.cuh" -#include "../../rocmfp4/rocmfp4_hip_scale.cuh" static __device__ __forceinline__ void dequantize_q1_0(const void * vx, const int64_t ib, const int iqs, float2 & v){ const block_q1_0 * x = (const block_q1_0 *) vx; diff --git a/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh b/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh deleted file mode 100644 index 180b2d9370d..00000000000 --- a/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh +++ /dev/null @@ -1,290 +0,0 @@ -static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna3_5(ggml_type type, int J, bool fallback) { - CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q1_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q1_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q4_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q4_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q5_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q5_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_1, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_1, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q8_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q8_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q8_0, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q8_0, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - -// --------------------------------------------------------------------------------------------- - - CASE(GGML_TYPE_Q2_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q2_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q2_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q3_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q3_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q3_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q3_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q4_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q4_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q4_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q5_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q5_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q5_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_Q6_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q6_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_Q6_K, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_Q6_K, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); - -// --------------------------------------------------------------------------------------------- - - CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ1_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ2_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_XXS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ3_S, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_XS, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_IQ4_NL, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); - -// --------------------------------------------------------------------------------------------- - - CASE(GGML_TYPE_MXFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_MXFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_MXFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_MXFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); - - CASE(GGML_TYPE_NVFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_NVFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); - CASE(GGML_TYPE_NVFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - CASE(GGML_TYPE_NVFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); - - return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); -} diff --git a/ggml/src/ggml-cuda/mmq-config-rdna35.cuh b/ggml/src/ggml-cuda/mmq-config-rdna35.cuh new file mode 100644 index 00000000000..1fbeb5dd646 --- /dev/null +++ b/ggml/src/ggml-cuda/mmq-config-rdna35.cuh @@ -0,0 +1,41 @@ +// RDNA3.5 (gfx1151) MMQ config overlay. +// +// This file holds ONLY the rocmfp4 CASE rows. rocmfp4 (Q4_0_ROCMFP4 dual-scale and +// _FAST single-scale) is an AMD-only, gfx1151-targeted quantization, so its MMQ config +// lives here rather than polluting the shared rdna4 table. Every other type falls +// through to ggml_cuda_mmq_get_config_rdna4() below. +// +// The dual type mirrors NVFP4 (SRAM layout NVFP4, Q8_0_16 vec_dot); the fast type mirrors +// MXFP4 (SRAM layout Q8_1). nthreads/I here are placeholders — ggml_cuda_mmq_get_config_rdna35 +// in mmq.cuh force-overrides them to 128/64 for the fork's RDNA3.5 kernels. +static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna3_5(ggml_type type, int J, bool fallback) { + // Q4_0_ROCMFP4 (dual-scale) — modeled on NVFP4. + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + + // Q4_0_ROCMFP4_FAST (single-scale) — modeled on MXFP4. + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 2, 128, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 2, 128, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 2, 128, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST, 256, 2, 128, 128, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + return ggml_cuda_mmq_get_config_rdna4(type, J, fallback); +} diff --git a/ggml/src/ggml-cuda/mmq-load-tiles.cuh b/ggml/src/ggml-cuda/mmq-load-tiles.cuh index acf62501e07..4ce59b46dc8 100644 --- a/ggml/src/ggml-cuda/mmq-load-tiles.cuh +++ b/ggml/src/ggml-cuda/mmq-load-tiles.cuh @@ -1915,3 +1915,140 @@ template static __device__ __forceinline_ x_u32_scale[i*sram_stride] = get_int_b4(bxi->d, 0); } } + +template static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_rocmfp4( + const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) { + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); + constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size; + constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); + constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + int * x_qs = (int *) x_tile; + float * x_df = (float *) (x_qs + MMQ_TILE_NE_K*2); +#else + constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_0_ROCMFP4, I); + int * x_qs = (int *) x_tile; + float * x_df = (float *) (x_qs + txs.qs); +#endif // defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + + constexpr int threads_per_row = MMQ_ITER_K / (4 * QR_ROCMFP4); + constexpr int nrows = warp_size / threads_per_row; + const int txi = warp_size > threads_per_row ? threadIdx.x % threads_per_row : threadIdx.x; + const int kbx = txi / QI_ROCMFP4; + const int kqsx = txi % QI_ROCMFP4; + +#pragma unroll + for (int i0 = 0; i0 < I; i0 += nrows*nwarps) { + int i = i0 + (nrows == 1 ? threadIdx.y : threadIdx.y*nrows + threadIdx.x/threads_per_row); + + if constexpr (fallback) { + i = min(i, i_max); + } + + const block_rocmfp4 * bxi = (const block_rocmfp4 *) x + kbx0 + i*stride + kbx; + + const int aux_q4 = rocmfp4_get_qs_i32(bxi->qs, kqsx); + const int2 v = rocmfp4_get_int_from_codebook_16(aux_q4, kvalues_rocmfp4); + const int k0 = kbx * (2 * QI_ROCMFP4) + kqsx; + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + x_qs[i*sram_stride + k0 + 0] = v.x; + x_qs[i*sram_stride + k0 + QI_ROCMFP4] = v.y; +#else + x_qs[i*(2*MMQ_TILE_NE_K + 1) + k0 + 0] = v.x; + x_qs[i*(2*MMQ_TILE_NE_K + 1) + k0 + QI_ROCMFP4] = v.y; +#endif // defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + } + + constexpr int blocks_per_tile_x_row = MMQ_TILE_NE_K / QI_ROCMFP4; + constexpr int rows_per_warp = warp_size / blocks_per_tile_x_row; + const int kbxd = threadIdx.x % blocks_per_tile_x_row; + +#pragma unroll + for (int i0 = 0; i0 < I; i0 += nwarps * rows_per_warp) { + int i = i0 + threadIdx.y * rows_per_warp + threadIdx.x / blocks_per_tile_x_row; + + if constexpr (fallback) { + i = min(i, i_max); + } + + const block_rocmfp4 * bxi = (const block_rocmfp4 *) x + kbx0 + i*stride + kbxd; + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + x_df[i*sram_stride + 2*kbxd + 0] = rocmfp4_ue4m3_to_fp32_half_finite(bxi->e[0]); + x_df[i*sram_stride + 2*kbxd + 1] = rocmfp4_ue4m3_to_fp32_half_finite(bxi->e[1]); +#else + x_df[i*(2*MMQ_TILE_NE_K*2/QI8_0) + i/(QI8_0/4) + 2*kbxd + 0] = rocmfp4_ue4m3_to_fp32_half_finite(bxi->e[0]); + x_df[i*(2*MMQ_TILE_NE_K*2/QI8_0) + i/(QI8_0/4) + 2*kbxd + 1] = rocmfp4_ue4m3_to_fp32_half_finite(bxi->e[1]); +#endif // defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + } +} + +template static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_rocmfp4_fast( + const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) { + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); + constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size; + constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback); + constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback); + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + int * x_qs = (int *) x_tile; + float * x_df = (float *) (x_qs + MMQ_TILE_NE_K*2); +#else + constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_0_ROCMFP4_FAST, I); + int * x_qs = (int *) x_tile; + float * x_df = (float *) (x_qs + txs.qs); +#endif // defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + + constexpr int threads_per_row = MMQ_ITER_K / (4 * QR_ROCMFP4); + constexpr int nrows = warp_size / threads_per_row; + const int txi = warp_size > threads_per_row ? threadIdx.x % threads_per_row : threadIdx.x; + const int kbx = txi / QI_ROCMFP4; + const int kqsx = txi % QI_ROCMFP4; + +#pragma unroll + for (int i0 = 0; i0 < I; i0 += nrows*nwarps) { + int i = i0 + (nrows == 1 ? threadIdx.y : threadIdx.y*nrows + threadIdx.x/threads_per_row); + + if constexpr (fallback) { + i = min(i, i_max); + } + + const block_rocmfp4_fast * bxi = (const block_rocmfp4_fast *) x + kbx0 + i*stride + kbx; + + const int aux_q4 = rocmfp4_get_qs_i32(bxi->qs, kqsx); + const int2 v = rocmfp4_get_int_from_codebook_16(aux_q4, kvalues_rocmfp4); + const int k0 = kbx * (2 * QI_ROCMFP4) + kqsx; + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + x_qs[i*sram_stride + k0 + 0] = v.x; + x_qs[i*sram_stride + k0 + QI_ROCMFP4] = v.y; +#else + x_qs[i*(2*MMQ_TILE_NE_K + 1) + k0 + 0] = v.x; + x_qs[i*(2*MMQ_TILE_NE_K + 1) + k0 + QI_ROCMFP4] = v.y; +#endif // defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + } + + constexpr int blocks_per_tile_x_row = MMQ_TILE_NE_K / QI_ROCMFP4; + constexpr int rows_per_warp = warp_size / blocks_per_tile_x_row; + const int kbxd = threadIdx.x % blocks_per_tile_x_row; + +#pragma unroll + for (int i0 = 0; i0 < I; i0 += nwarps * rows_per_warp) { + int i = i0 + threadIdx.y * rows_per_warp + threadIdx.x / blocks_per_tile_x_row; + + if constexpr (fallback) { + i = min(i, i_max); + } + + const block_rocmfp4_fast * bxi = (const block_rocmfp4_fast *) x + kbx0 + i*stride + kbxd; + const float d = rocmfp4_ue4m3_to_fp32_half_finite(bxi->e); + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + x_df[i*sram_stride + kbxd] = d; +#else + x_df[i*(2*MMQ_TILE_NE_K/QI8_0) + i/(QI8_0/2) + kbxd] = d; +#endif // defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + } +} diff --git a/ggml/src/ggml-cuda/mmq.cu b/ggml/src/ggml-cuda/mmq.cu index 707437ea3e5..478088b1777 100644 --- a/ggml/src/ggml-cuda/mmq.cu +++ b/ggml/src/ggml-cuda/mmq.cu @@ -76,6 +76,12 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con case GGML_TYPE_NVFP4: mul_mat_q_case(ctx, args, stream); break; + case GGML_TYPE_Q4_0_ROCMFP4: + mul_mat_q_case(ctx, args, stream); + break; + case GGML_TYPE_Q4_0_ROCMFP4_FAST: + mul_mat_q_case(ctx, args, stream); + break; default: GGML_ABORT("fatal error"); break; @@ -291,6 +297,11 @@ bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t case GGML_TYPE_NVFP4: mmq_supported = true; break; + case GGML_TYPE_Q4_0_ROCMFP4: + case GGML_TYPE_Q4_0_ROCMFP4_FAST: + // guard on RDNA3.5 (gfx1151) for now. + mmq_supported = GGML_CUDA_CC_IS_RDNA3_5(cc); + break; default: mmq_supported = false; break; diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh index 42c0ef5542f..389c227bd3b 100644 --- a/ggml/src/ggml-cuda/mmq.cuh +++ b/ggml/src/ggml-cuda/mmq.cuh @@ -76,6 +76,9 @@ static mmq_q8_1_ds_layout mmq_get_q8_1_ds_layout(const ggml_type type_x) { return MMQ_Q8_1_DS_LAYOUT_D4; case GGML_TYPE_NVFP4: return MMQ_Q8_1_DS_LAYOUT_D4; + case GGML_TYPE_Q4_0_ROCMFP4: + case GGML_TYPE_Q4_0_ROCMFP4_FAST: + return MMQ_Q8_1_DS_LAYOUT_D4; case GGML_TYPE_Q2_K: return MMQ_Q8_1_DS_LAYOUT_D2S6; case GGML_TYPE_Q3_K: @@ -228,19 +231,19 @@ struct ggml_cuda_mmq_config { #include "mmq-config-cdna.cuh" #include "mmq-config-rdna2.cuh" #include "mmq-config-rdna3.cuh" -#include "mmq-config-rdna3-5.cuh" #include "mmq-config-rdna4.cuh" +#include "mmq-config-rdna35.cuh" // after rdna4: falls through to ggml_cuda_mmq_get_config_rdna4 #undef CASE // RDNA3.5 (gfx1151) uses fork-specific MMQ kernels (ggml_cuda_mmq_load_tiles_q4_K_rdna35, // ggml_cuda_mmq_vec_dot_q6_K_q8_1_mma_rdna35) that static_assert nthreads == 128 && I == 64, -// so its config must force that shape. Upstream's dedicated get_config_rdna3_5 (256/128 at -// wide J) is incompatible with those kernels; keep the fork override for RDNA3.5 while the -// other AMD arches use upstream's per-arch tables. +// so its config must force that shape. ggml_cuda_mmq_get_config_rdna3_5 (mmq-config-rdna35.cuh) +// supplies the gfx1151-only rocmfp4 CASEs and falls through to the rdna4 table for every other +// type; we then force the 128/64 shape the fork's RDNA3.5 kernels require. static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_rdna35( ggml_type type, int J, bool fallback) { - ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config_rdna4(type, J, fallback); + ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config_rdna3_5(type, J, fallback); config.nthreads = 128; config.I = 64; return config; @@ -415,6 +418,8 @@ static constexpr __host__ __device__ tile_x_sizes mmq_get_dp4a_tile_x_sizes(ggml case GGML_TYPE_Q8_0: return MMQ_DP4A_TXS_Q8_0; case GGML_TYPE_MXFP4: return MMQ_DP4A_TXS_Q8_1; case GGML_TYPE_NVFP4: return MMQ_DP4A_TXS_Q8_0_16; + case GGML_TYPE_Q4_0_ROCMFP4: return MMQ_DP4A_TXS_Q8_0_16; + case GGML_TYPE_Q4_0_ROCMFP4_FAST: return MMQ_DP4A_TXS_Q8_0; case GGML_TYPE_Q2_K: return MMQ_DP4A_TXS_Q2_K; case GGML_TYPE_Q3_K: return MMQ_DP4A_TXS_Q3_K; case GGML_TYPE_Q4_K: return MMQ_DP4A_TXS_Q4_K; @@ -769,6 +774,18 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func ggml_cuda_mmq_load_tiles_nvfp4, ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a, ggml_cuda_mmq_write_back_dp4a); + case GGML_TYPE_Q4_0_ROCMFP4: + return ggml_cuda_mmq_util_funcs( + VDR_ROCMFP4_Q8_1_MMQ, + ggml_cuda_mmq_load_tiles_rocmfp4, + ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a, + ggml_cuda_mmq_write_back_dp4a); + case GGML_TYPE_Q4_0_ROCMFP4_FAST: + return ggml_cuda_mmq_util_funcs( + VDR_ROCMFP4_FAST_Q8_1_MMQ, + ggml_cuda_mmq_load_tiles_rocmfp4_fast, + ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a, + ggml_cuda_mmq_write_back_dp4a); default: return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr); } @@ -937,6 +954,18 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func ggml_cuda_mmq_load_tiles_nvfp4, ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma, ggml_cuda_mmq_write_back_mma); + case GGML_TYPE_Q4_0_ROCMFP4: + return ggml_cuda_mmq_util_funcs( + -1, + ggml_cuda_mmq_load_tiles_rocmfp4, + ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma, + ggml_cuda_mmq_write_back_mma); + case GGML_TYPE_Q4_0_ROCMFP4_FAST: + return ggml_cuda_mmq_util_funcs( + -1, + ggml_cuda_mmq_load_tiles_rocmfp4_fast, + ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma, + ggml_cuda_mmq_write_back_mma); default: return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr); } @@ -2034,6 +2063,8 @@ extern DECL_MMQ_CASE(GGML_TYPE_IQ4_XS); // ----------------------------------------- extern DECL_MMQ_CASE(GGML_TYPE_MXFP4); extern DECL_MMQ_CASE(GGML_TYPE_NVFP4); +extern DECL_MMQ_CASE(GGML_TYPE_Q4_0_ROCMFP4); +extern DECL_MMQ_CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST); // ------------------------------------------------------------------------------------------------------------------------- diff --git a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py index d7cd271675e..d9a05868d08 100755 --- a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +++ b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py @@ -40,7 +40,8 @@ "GGML_TYPE_Q4_0", "GGML_TYPE_Q4_1", "GGML_TYPE_Q5_0", "GGML_TYPE_Q5_1", "GGML_TYPE_Q8_0", "GGML_TYPE_Q2_K", "GGML_TYPE_Q3_K", "GGML_TYPE_Q4_K", "GGML_TYPE_Q5_K", "GGML_TYPE_Q6_K", "GGML_TYPE_IQ2_XXS", "GGML_TYPE_IQ2_XS", "GGML_TYPE_IQ2_S", "GGML_TYPE_IQ3_XXS", "GGML_TYPE_IQ3_S", - "GGML_TYPE_IQ1_S", "GGML_TYPE_IQ4_NL", "GGML_TYPE_IQ4_XS", "GGML_TYPE_MXFP4", "GGML_TYPE_NVFP4" + "GGML_TYPE_IQ1_S", "GGML_TYPE_IQ4_NL", "GGML_TYPE_IQ4_XS", "GGML_TYPE_MXFP4", "GGML_TYPE_NVFP4", + "GGML_TYPE_Q4_0_ROCMFP4", "GGML_TYPE_Q4_0_ROCMFP4_FAST" ] SOURCE_MMQ = """// This file has been autogenerated by generate_cu_files.py, do not edit manually. diff --git a/ggml/src/ggml-cuda/vecdotq.cuh b/ggml/src/ggml-cuda/vecdotq.cuh index 4decde2eecf..47b9fa3ff4f 100644 --- a/ggml/src/ggml-cuda/vecdotq.cuh +++ b/ggml/src/ggml-cuda/vecdotq.cuh @@ -4,8 +4,6 @@ #include -#include "../../rocmfp4/rocmfp4_hip_codebook.cuh" - static __device__ __forceinline__ int get_int_b1(const void * x, const int & i32) { const uint8_t * x8 = (const uint8_t *) x; @@ -386,6 +384,56 @@ static __device__ __forceinline__ float vec_dot_nvfp4_q8_1( #define VDR_ROCMFP4_FAST_Q8_1_MMVQ GGML_ROCMFP4_FAST_Q8_1_MMVQ_VDR #define VDR_ROCMFP4_FAST_Q8_1_MMQ GGML_ROCMFP4_FAST_Q8_1_MMQ_VDR +#ifndef GGML_ROCMFP4_UNALIGNED_QS_DWORD_LOAD +#define GGML_ROCMFP4_UNALIGNED_QS_DWORD_LOAD 1 +#endif + +static __device__ __forceinline__ int rocmfp4_get_qs_i32(const void * x, const int & i32) { +#if defined(GGML_USE_HIP) && GGML_ROCMFP4_UNALIGNED_QS_DWORD_LOAD + return *((const int *) ((const uint8_t *) x + 4*i32)); +#else + const uint8_t * x8 = (const uint8_t *) x; + + int x32 = x8[4*i32 + 0] << 0; + x32 |= x8[4*i32 + 1] << 8; + x32 |= x8[4*i32 + 2] << 16; + x32 |= x8[4*i32 + 3] << 24; + + return x32; +#endif +} + +// AMD-specific fast path for expanding eight packed ROCmFP4 nibbles into two int32 DP4A operands. +// This encodes the Codebook10 table directly as four 32-bit constants: +// [0, 1, 2, 3], [4, 6, 8, 10], [0, -1, -2, -3], [-4, -6, -8, -10] +// Avoiding the table pointer keeps the ROCm/HIP MMVQ/MMQ hot path fully local to this format. +// Non-HIP builds still use llama.cpp's generic table expander. +static __device__ __forceinline__ int2 rocmfp4_get_int_from_codebook_16(const int & q4, const int8_t * fallback_table) { +#if defined(GGML_USE_HIP) + constexpr uint32_t values0 = 0x03020100u; + constexpr uint32_t values1 = 0x0a080604u; + constexpr uint32_t values2 = 0xfdfeff00u; + constexpr uint32_t values3 = 0xf6f8fafcu; + + const uint32_t q_even = q4; + const uint32_t q_odd = q4 >> 4; + + const uint32_t v_even_low = __builtin_amdgcn_perm(values1, values0, q_even & 0x07070707u); + const uint32_t v_odd_low = __builtin_amdgcn_perm(values1, values0, q_odd & 0x07070707u); + const uint32_t v_even_high = __builtin_amdgcn_perm(values3, values2, q_even & 0x07070707u); + const uint32_t v_odd_high = __builtin_amdgcn_perm(values3, values2, q_odd & 0x07070707u); + + const uint32_t mask_even = 0x03020100u | ((q_even & 0x08080808u) >> 1); + const uint32_t mask_odd = 0x03020100u | ((q_odd & 0x08080808u) >> 1); + + return make_int2( + __builtin_amdgcn_perm(v_even_high, v_even_low, mask_even), + __builtin_amdgcn_perm(v_odd_high, v_odd_low, mask_odd)); +#else + return get_int_from_table_16(q4, fallback_table); +#endif +} + static __device__ __forceinline__ float vec_dot_rocmfp4_q8_1( const void * __restrict__ vbq, const block_q8_1 * __restrict__ bq8_1, const int & kbx, const int & iqs) { diff --git a/ggml/src/ggml-hip/CMakeLists.txt b/ggml/src/ggml-hip/CMakeLists.txt index 624eb7db06a..a17b980aca0 100644 --- a/ggml/src/ggml-hip/CMakeLists.txt +++ b/ggml/src/ggml-hip/CMakeLists.txt @@ -61,7 +61,6 @@ file(GLOB GGML_HEADERS_ROCM "../ggml-cuda/*.cuh") list(APPEND GGML_HEADERS_ROCM "../../include/ggml-cuda.h") file(GLOB GGML_SOURCES_ROCM "../ggml-cuda/*.cu") -list(APPEND GGML_SOURCES_ROCM "../../rocmfp4/rocmfp4_hip.cu") file(GLOB SRCS "../ggml-cuda/template-instances/fattn-tile*.cu") list(APPEND GGML_SOURCES_ROCM ${SRCS}) file(GLOB SRCS "../ggml-cuda/template-instances/fattn-mma*.cu") diff --git a/ggml/src/ggml-quants.c b/ggml/src/ggml-quants.c index b69ed324c0a..9b301095a7f 100644 --- a/ggml/src/ggml-quants.c +++ b/ggml/src/ggml-quants.c @@ -5,7 +5,6 @@ #include "ggml-impl.h" #include "ggml-cpu/ggml-cpu-impl.h" #include "ggml-cpu.h" -#include "../rocmfp4/rocmfp4.h" #include #include @@ -612,6 +611,535 @@ void dequantize_row_nvfp4(const block_nvfp4 * GGML_RESTRICT x, float * GGML_REST } } +// ============================================================================ +// ROCmFP4 (AMD gfx1151): Q4_0-layout FP4 with per-half UE4M3 scales. +// The codebook is E2M1-derived at half-scale, with the top magnitude retuned to 12 ->10. +// Decoded values live in kvalues_rocmfp4 (== 2*E2M1 with that retune), and UE4M3 scale bytes decode via ggml_ue4m3_to_fp32 (raw*0.5, "half" scale). +// The scale-search/quantize path below keeps a local half-scale table and an arithmetic decoder so the exhaustive MSE search avoids per-candidate rebuilds. +// ============================================================================ + +static inline int8_t rocmfp4_decode(uint8_t q) { + q &= 0x0f; + const int mag3 = q & 0x07; + const int mag = mag3 <= 4 ? mag3 : 2*mag3 - 4; + return (q & 0x08) ? -mag : mag; +} + +// Finite unsigned E4M3 scale bytes decoded to the half-scale values used by +// ROCmFP4. Keeping this as a table avoids rebuilding identical FP32 values for +// every candidate during exhaustive scale search. Values come from the shared +// GGML_ROCMFP4_SCALE_UE4M3_HALF_LIST in ggml-common.h (single source of truth). +static const float rocmfp4_scale_ue4m3_half[127] = { GGML_ROCMFP4_SCALE_UE4M3_HALF_LIST }; + +static inline float rocmfp4_ue4m3_to_fp32_half(uint8_t e) { + return e <= 0x7e ? rocmfp4_scale_ue4m3_half[e] : 0.0f; +} + +static inline uint8_t rocmfp4_best_index_scaled_finite(float x, float inv_scale_half) { + // Exact nearest-neighbor thresholds for Codebook10: + // 0, +/-1, +/-2, +/-3, +/-4, +/-6, +/-8, +/-10 + // Ties intentionally choose the lower-magnitude code, matching the former linear scan because the positive codes and zero appear first. + const float a = fabsf(x * inv_scale_half); + if (a <= 0.5f) { + return 0; + } + + const bool neg = x < 0.0f; + if (a <= 1.5f) { + return neg ? 9 : 1; + } + if (a <= 2.5f) { + return neg ? 10 : 2; + } + if (a <= 3.5f) { + return neg ? 11 : 3; + } + if (a <= 5.0f) { + return neg ? 12 : 4; + } + if (a <= 7.0f) { + return neg ? 13 : 5; + } + if (a <= 9.0f) { + return neg ? 14 : 6; + } + + return neg ? 15 : 7; +} + +static inline uint8_t rocmfp4_best_index_scaled(float x, float inv_scale_half) { + if (!isfinite(x)) { + return 0; + } + + return rocmfp4_best_index_scaled_finite(x, inv_scale_half); +} + +static inline bool rocmfp4_scale_is_valid(uint8_t e) { + // ROCmFP4 scale bytes are unsigned finite E4M3 values. 0x7f is NaN in the + // unsigned encoding and values with the sign bit set are not valid scales. + return e <= 0x7e; +} + +static float rocmfp4_block_mse_for_scale_unweighted( + const float * x, int n, int e, float best_err) { + const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); + const float inv_scale_half = 1.0f / scale_half; + float err = 0.0f; + + for (int i = 0; i < n; ++i) { + const uint8_t q = rocmfp4_best_index_scaled(x[i], inv_scale_half); + const float y = (float) rocmfp4_decode(q) * scale_half; + const float d = x[i] - y; + + err += d*d; + if (err > best_err) { + return err; + } + } + + return err; +} + +static float rocmfp4_block_mse_for_scale_unweighted_finite( + const float * x, int n, int e, float best_err) { + const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); + const float inv_scale_half = 1.0f / scale_half; + float err = 0.0f; + + for (int i = 0; i < n; ++i) { + const uint8_t q = rocmfp4_best_index_scaled_finite(x[i], inv_scale_half); + const float y = (float) rocmfp4_decode(q) * scale_half; + const float d = x[i] - y; + + err += d*d; + if (err > best_err) { + return err; + } + } + + return err; +} + +static float rocmfp4_block_mse_for_scale_weighted( + const float * x, int n, const float * mse_weights, int e, float best_err) { + const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); + const float inv_scale_half = 1.0f / scale_half; + float err = 0.0f; + + for (int i = 0; i < n; ++i) { + const uint8_t q = rocmfp4_best_index_scaled(x[i], inv_scale_half); + const float y = (float) rocmfp4_decode(q) * scale_half; + const float d = x[i] - y; + + err += mse_weights[i]*d*d; + if (err > best_err) { + return err; + } + } + + return err; +} + +static float rocmfp4_block_mse_for_scale_weighted_finite( + const float * x, int n, const float * mse_weights, int e, float best_err) { + const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e); + const float inv_scale_half = 1.0f / scale_half; + float err = 0.0f; + + for (int i = 0; i < n; ++i) { + const uint8_t q = rocmfp4_best_index_scaled_finite(x[i], inv_scale_half); + const float y = (float) rocmfp4_decode(q) * scale_half; + const float d = x[i] - y; + + err += mse_weights[i]*d*d; + if (err > best_err) { + return err; + } + } + + return err; +} + +static void rocmfp4_prepare_mse_weights( + float * dst, const float * x, int n, const float * quant_weights, float sigma2, + float * max_abs, float * max_abs_weight, bool * all_finite) { + *max_abs = 0.0f; + *max_abs_weight = 0.0f; + *all_finite = true; + + for (int i = 0; i < n; ++i) { + const float qw = quant_weights[i]; + const float ax = fabsf(x[i]); + const float weight = isfinite(qw) && qw > 0.0f ? qw * sqrtf(sigma2 + x[i]*x[i]) : 0.0f; + *all_finite = *all_finite && isfinite(x[i]); + + if (ax > *max_abs) { + *max_abs = ax; + *max_abs_weight = weight; + } else if (ax == *max_abs && weight > *max_abs_weight) { + *max_abs_weight = weight; + } + + // Match llama.cpp's imatrix weighting style for Q4_0: calibration + // importance is scaled by row energy so large activations remain protected. + dst[i] = weight; + } +} + +static int rocmfp4_nearest_scale_ue4m3(float target_scale_half) { + if (!(target_scale_half > 0.0f) || !isfinite(target_scale_half)) { + return 1; + } + + int lo = 1; + int hi = 126; + while (lo < hi) { + const int mid = lo + (hi - lo) / 2; + if (rocmfp4_ue4m3_to_fp32_half((uint8_t) mid) < target_scale_half) { + lo = mid + 1; + } else { + hi = mid; + } + } + + if (lo == 1) { + return 1; + } + + const float hi_scale = rocmfp4_ue4m3_to_fp32_half((uint8_t) lo); + const float lo_scale = rocmfp4_ue4m3_to_fp32_half((uint8_t) (lo - 1)); + + // Match the former ascending nearest scan: exact midpoint ties keep the + // lower scale byte. + return (target_scale_half - lo_scale <= hi_scale - target_scale_half) ? lo - 1 : lo; +} + +static uint8_t rocmfp4_choose_scale_ue4m3_exhaustive_unweighted( + const float * x, int n, float max_abs, bool all_finite) { + const int start_e = rocmfp4_nearest_scale_ue4m3(max_abs / 10.0f); + + int best_e = 0; + float best_err = FLT_MAX; + bool lower_done = false; + + for (int delta = 0; delta <= 125; ++delta) { + const int e0 = start_e - delta; + if (!lower_done && e0 >= 1 && e0 <= 126) { + const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e0); + const float clip_delta = max_abs - 10.0f*scale_half; + if (clip_delta > 0.0f && clip_delta*clip_delta > best_err) { + lower_done = true; + } else { + const float err = all_finite ? + rocmfp4_block_mse_for_scale_unweighted_finite(x, n, e0, best_err) : + rocmfp4_block_mse_for_scale_unweighted(x, n, e0, best_err); + if (err < best_err || (err == best_err && e0 < best_e)) { + best_err = err; + best_e = e0; + } + } + } + + const int e1 = start_e + delta; + if (delta != 0 && e1 >= 1 && e1 <= 126) { + const float err = all_finite ? + rocmfp4_block_mse_for_scale_unweighted_finite(x, n, e1, best_err) : + rocmfp4_block_mse_for_scale_unweighted(x, n, e1, best_err); + if (err < best_err || (err == best_err && e1 < best_e)) { + best_err = err; + best_e = e1; + } + } + + if ((lower_done || e0 <= 1) && e1 >= 126) { + break; + } + } + + return (uint8_t) best_e; +} + +static uint8_t rocmfp4_choose_scale_ue4m3_exhaustive_weighted( + const float * x, int n, const float * mse_weights, float max_abs, float max_abs_weight, bool all_finite) { + const int start_e = rocmfp4_nearest_scale_ue4m3(max_abs / 10.0f); + + int best_e = 0; + float best_err = FLT_MAX; + bool lower_done = false; + + for (int delta = 0; delta <= 125; ++delta) { + const int e0 = start_e - delta; + if (!lower_done && e0 >= 1 && e0 <= 126) { + const float scale_half = rocmfp4_ue4m3_to_fp32_half((uint8_t) e0); + const float clip_delta = max_abs - 10.0f*scale_half; + if (max_abs_weight > 0.0f && clip_delta > 0.0f && max_abs_weight*clip_delta*clip_delta > best_err) { + lower_done = true; + } else { + const float err = all_finite ? + rocmfp4_block_mse_for_scale_weighted_finite(x, n, mse_weights, e0, best_err) : + rocmfp4_block_mse_for_scale_weighted(x, n, mse_weights, e0, best_err); + if (err < best_err || (err == best_err && e0 < best_e)) { + best_err = err; + best_e = e0; + } + } + } + + const int e1 = start_e + delta; + if (delta != 0 && e1 >= 1 && e1 <= 126) { + const float err = all_finite ? + rocmfp4_block_mse_for_scale_weighted_finite(x, n, mse_weights, e1, best_err) : + rocmfp4_block_mse_for_scale_weighted(x, n, mse_weights, e1, best_err); + if (err < best_err || (err == best_err && e1 < best_e)) { + best_err = err; + best_e = e1; + } + } + + if ((lower_done || e0 <= 1) && e1 >= 126) { + break; + } + } + + return (uint8_t) best_e; +} + +static uint8_t rocmfp4_choose_scale_ue4m3(const float * x, int n, const float * quant_weights, float sigma2) { + if (quant_weights) { + assert(n <= QK_ROCMFP4); + float mse_weights_buf[QK_ROCMFP4]; + float weighted_max_abs; + float max_abs_weight; + bool all_finite; + rocmfp4_prepare_mse_weights(mse_weights_buf, x, n, quant_weights, sigma2, &weighted_max_abs, &max_abs_weight, &all_finite); + if (!(weighted_max_abs > 0.0f) || !isfinite(weighted_max_abs)) { + return 0; + } + return rocmfp4_choose_scale_ue4m3_exhaustive_weighted(x, n, mse_weights_buf, weighted_max_abs, max_abs_weight, all_finite); + } + + float max_abs = 0.0f; + bool all_finite = true; + for (int i = 0; i < n; ++i) { + all_finite = all_finite && isfinite(x[i]); + const float ax = fabsf(x[i]); + if (ax > max_abs) { + max_abs = ax; + } + } + + if (!(max_abs > 0.0f) || !isfinite(max_abs)) { + return 0; + } + + return rocmfp4_choose_scale_ue4m3_exhaustive_unweighted(x, n, max_abs, all_finite); +} + +static void rocmfp4_quantize_row_q4_0_weighted( + const float * GGML_RESTRICT x, block_rocmfp4 * GGML_RESTRICT y, int64_t k, const float * GGML_RESTRICT quant_weights) { + assert(k % QK_ROCMFP4 == 0); + + float sum_x2 = 0.0f; + for (int64_t i = 0; i < k; ++i) { + sum_x2 += x[i]*x[i]; + } + const float sigma2 = sum_x2 / (float) k; + + const int64_t nb = k / QK_ROCMFP4; + for (int64_t ib = 0; ib < nb; ++ib) { + const float * xb = x + ib*QK_ROCMFP4; + const float * qw = quant_weights ? quant_weights + ib*QK_ROCMFP4 : NULL; + const uint8_t e0 = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4/2, qw, sigma2); + const uint8_t e1 = rocmfp4_choose_scale_ue4m3(xb + QK_ROCMFP4/2, QK_ROCMFP4/2, qw ? qw + QK_ROCMFP4/2 : NULL, sigma2); + const float scale_half0 = rocmfp4_ue4m3_to_fp32_half(e0); + const float scale_half1 = rocmfp4_ue4m3_to_fp32_half(e1); + const float inv_scale_half0 = scale_half0 > 0.0f ? 1.0f / scale_half0 : 0.0f; + const float inv_scale_half1 = scale_half1 > 0.0f ? 1.0f / scale_half1 : 0.0f; + + y[ib].e[0] = e0; + y[ib].e[1] = e1; + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half0); + const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half1); + y[ib].qs[j] = q0 | (q1 << 4); + } + } +} + +static void rocmfp4_quantize_row_q4_0_fast_weighted( + const float * GGML_RESTRICT x, block_rocmfp4_fast * GGML_RESTRICT y, int64_t k, const float * GGML_RESTRICT quant_weights) { + assert(k % QK_ROCMFP4 == 0); + + float sum_x2 = 0.0f; + for (int64_t i = 0; i < k; ++i) { + sum_x2 += x[i]*x[i]; + } + const float sigma2 = sum_x2 / (float) k; + + const int64_t nb = k / QK_ROCMFP4; + for (int64_t ib = 0; ib < nb; ++ib) { + const float * xb = x + ib*QK_ROCMFP4; + const float * qw = quant_weights ? quant_weights + ib*QK_ROCMFP4 : NULL; + const uint8_t e = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4, qw, sigma2); + const float scale_half = rocmfp4_ue4m3_to_fp32_half(e); + const float inv_scale_half = scale_half > 0.0f ? 1.0f / scale_half : 0.0f; + + y[ib].e = e; + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half); + const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half); + y[ib].qs[j] = q0 | (q1 << 4); + } + } +} + +void rocmfp4_quantize_row_q4_0_ref(const float * GGML_RESTRICT x, block_rocmfp4 * GGML_RESTRICT y, int64_t k) { + assert(k % QK_ROCMFP4 == 0); + + const int64_t nb = k / QK_ROCMFP4; + for (int64_t ib = 0; ib < nb; ++ib) { + const float * xb = x + ib*QK_ROCMFP4; + const uint8_t e0 = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4/2, NULL, 0.0f); + const uint8_t e1 = rocmfp4_choose_scale_ue4m3(xb + QK_ROCMFP4/2, QK_ROCMFP4/2, NULL, 0.0f); + const float scale_half0 = rocmfp4_ue4m3_to_fp32_half(e0); + const float scale_half1 = rocmfp4_ue4m3_to_fp32_half(e1); + const float inv_scale_half0 = scale_half0 > 0.0f ? 1.0f / scale_half0 : 0.0f; + const float inv_scale_half1 = scale_half1 > 0.0f ? 1.0f / scale_half1 : 0.0f; + + y[ib].e[0] = e0; + y[ib].e[1] = e1; + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half0); + const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half1); + y[ib].qs[j] = q0 | (q1 << 4); + } + } +} + +void rocmfp4_quantize_row_q4_0_fast_ref(const float * GGML_RESTRICT x, block_rocmfp4_fast * GGML_RESTRICT y, int64_t k) { + assert(k % QK_ROCMFP4 == 0); + + const int64_t nb = k / QK_ROCMFP4; + for (int64_t ib = 0; ib < nb; ++ib) { + const float * xb = x + ib*QK_ROCMFP4; + const uint8_t e = rocmfp4_choose_scale_ue4m3(xb, QK_ROCMFP4, NULL, 0.0f); + const float scale_half = rocmfp4_ue4m3_to_fp32_half(e); + const float inv_scale_half = scale_half > 0.0f ? 1.0f / scale_half : 0.0f; + + y[ib].e = e; + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + const uint8_t q0 = rocmfp4_best_index_scaled(xb[j], inv_scale_half); + const uint8_t q1 = rocmfp4_best_index_scaled(xb[j + QK_ROCMFP4/2], inv_scale_half); + y[ib].qs[j] = q0 | (q1 << 4); + } + } +} + +// Dequant/vec_dot decode via the shared codebook (kvalues_rocmfp4) and shared +// UE4M3 decoder (ggml_ue4m3_to_fp32, raw*0.5) — bit-identical to the former +// local rocmfp4_decode_table / rocmfp4_ue4m3_to_fp32_half. +void rocmfp4_dequantize_row_q4_0(const block_rocmfp4 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { + assert(k % QK_ROCMFP4 == 0); + + const int64_t nb = k / QK_ROCMFP4; + for (int64_t ib = 0; ib < nb; ++ib) { + const float d0 = ggml_ue4m3_to_fp32(x[ib].e[0]); + const float d1 = ggml_ue4m3_to_fp32(x[ib].e[1]); + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + y[ib*QK_ROCMFP4 + j] = (float) kvalues_rocmfp4[x[ib].qs[j] & 0x0f] * d0; + y[ib*QK_ROCMFP4 + j + QK_ROCMFP4/2] = (float) kvalues_rocmfp4[x[ib].qs[j] >> 4] * d1; + } + } +} + +void rocmfp4_dequantize_row_q4_0_fast(const block_rocmfp4_fast * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { + assert(k % QK_ROCMFP4 == 0); + + const int64_t nb = k / QK_ROCMFP4; + for (int64_t ib = 0; ib < nb; ++ib) { + const float d = ggml_ue4m3_to_fp32(x[ib].e); + + for (int j = 0; j < QK_ROCMFP4/2; ++j) { + y[ib*QK_ROCMFP4 + j] = (float) kvalues_rocmfp4[x[ib].qs[j] & 0x0f] * d; + y[ib*QK_ROCMFP4 + j + QK_ROCMFP4/2] = (float) kvalues_rocmfp4[x[ib].qs[j] >> 4] * d; + } + } +} + +size_t rocmfp4_quantize_q4_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix) { + const size_t row_size = ggml_row_size(GGML_TYPE_Q4_0_ROCMFP4, n_per_row); + + if (!imatrix) { + rocmfp4_quantize_row_q4_0_ref(src, (block_rocmfp4 *) dst, nrows*n_per_row); + return nrows * row_size; + } + + char * qrow = (char *) dst; + for (int64_t row = 0; row < nrows; ++row) { + rocmfp4_quantize_row_q4_0_weighted(src, (block_rocmfp4 *) qrow, n_per_row, imatrix); + src += n_per_row; + qrow += row_size; + } + + return nrows * row_size; +} + +size_t rocmfp4_quantize_q4_0_fast(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix) { + const size_t row_size = ggml_row_size(GGML_TYPE_Q4_0_ROCMFP4_FAST, n_per_row); + + if (!imatrix) { + rocmfp4_quantize_row_q4_0_fast_ref(src, (block_rocmfp4_fast *) dst, nrows*n_per_row); + return nrows * row_size; + } + + char * qrow = (char *) dst; + for (int64_t row = 0; row < nrows; ++row) { + rocmfp4_quantize_row_q4_0_fast_weighted(src, (block_rocmfp4_fast *) qrow, n_per_row, imatrix); + src += n_per_row; + qrow += row_size; + } + + return nrows * row_size; +} + +bool rocmfp4_validate_row_data(const void * data, size_t nbytes) { + if (nbytes % sizeof(block_rocmfp4) != 0) { + return false; + } + + const block_rocmfp4 * blocks = (const block_rocmfp4 *) data; + const size_t nblocks = nbytes / sizeof(block_rocmfp4); + for (size_t i = 0; i < nblocks; ++i) { + if (!rocmfp4_scale_is_valid(blocks[i].e[0]) || !rocmfp4_scale_is_valid(blocks[i].e[1])) { + return false; + } + } + + return true; +} + +bool rocmfp4_validate_row_data_fast(const void * data, size_t nbytes) { + if (nbytes % sizeof(block_rocmfp4_fast) != 0) { + return false; + } + + const block_rocmfp4_fast * blocks = (const block_rocmfp4_fast *) data; + const size_t nblocks = nbytes / sizeof(block_rocmfp4_fast); + for (size_t i = 0; i < nblocks; ++i) { + if (!rocmfp4_scale_is_valid(blocks[i].e)) { + return false; + } + } + + return true; +} + // // 2-6 bit quantization in super-blocks // diff --git a/ggml/src/ggml-quants.h b/ggml/src/ggml-quants.h index 75188f1af18..07f021a3fda 100644 --- a/ggml/src/ggml-quants.h +++ b/ggml/src/ggml-quants.h @@ -26,6 +26,9 @@ GGML_API void quantize_row_q8_1_ref(const float * GGML_RESTRICT x, block_q8_1 * GGML_API void quantize_row_mxfp4_ref(const float * GGML_RESTRICT x, block_mxfp4 * GGML_RESTRICT y, int64_t k); GGML_API void quantize_row_nvfp4_ref(const float * GGML_RESTRICT x, block_nvfp4 * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfp4_quantize_row_q4_0_ref (const float * GGML_RESTRICT x, block_rocmfp4 * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfp4_quantize_row_q4_0_fast_ref(const float * GGML_RESTRICT x, block_rocmfp4_fast * GGML_RESTRICT y, int64_t k); + GGML_API void quantize_row_q2_K_ref(const float * GGML_RESTRICT x, block_q2_K * GGML_RESTRICT y, int64_t k); GGML_API void quantize_row_q3_K_ref(const float * GGML_RESTRICT x, block_q3_K * GGML_RESTRICT y, int64_t k); GGML_API void quantize_row_q4_K_ref(const float * GGML_RESTRICT x, block_q4_K * GGML_RESTRICT y, int64_t k); @@ -55,6 +58,9 @@ GGML_API void dequantize_row_q8_0(const block_q8_0 * GGML_RESTRICT x, float * GG GGML_API void dequantize_row_mxfp4(const block_mxfp4 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void dequantize_row_nvfp4(const block_nvfp4 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfp4_dequantize_row_q4_0 (const block_rocmfp4 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfp4_dequantize_row_q4_0_fast(const block_rocmfp4_fast * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); + GGML_API void dequantize_row_q2_K(const block_q2_K * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void dequantize_row_q3_K(const block_q3_K * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void dequantize_row_q4_K(const block_q4_K * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); @@ -105,6 +111,11 @@ GGML_API size_t quantize_q8_0(const float * GGML_RESTRICT src, void * GGML_RESTR GGML_API size_t quantize_mxfp4(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); GGML_API size_t quantize_nvfp4(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); +GGML_API size_t rocmfp4_quantize_q4_0 (const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); +GGML_API size_t rocmfp4_quantize_q4_0_fast(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); +GGML_API bool rocmfp4_validate_row_data (const void * data, size_t nbytes); +GGML_API bool rocmfp4_validate_row_data_fast(const void * data, size_t nbytes); + GGML_API void iq2xs_init_impl(enum ggml_type type); GGML_API void iq2xs_free_impl(enum ggml_type type); GGML_API void iq3xs_init_impl(int grid_size); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 7dfb8eb0e07..d25f5fe1054 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -9,7 +9,6 @@ // FIXME: required here for quantization functions #include "ggml-quants.h" -#include "../rocmfp4/rocmfp4.h" #ifdef GGML_USE_CPU_HBM #include diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index a62d781f677..cb3a3611328 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8178,6 +8178,8 @@ static constexpr ggml_type rdna35_mmq_types[] = { GGML_TYPE_IQ4_NL, GGML_TYPE_MXFP4, GGML_TYPE_NVFP4, + GGML_TYPE_Q4_0_ROCMFP4, + GGML_TYPE_Q4_0_ROCMFP4_FAST, }; static constexpr mmq_test_shape rdna35_mmq_shapes[] = {