Commit 8a1a9b512 for llama.cpp
commit 8a1a9b5126126e5228b95fa909d4b08fac65e8b3
Author: ynankani <ynankani@nvidia.com>
Date: Fri Oct 9 06:32:52 2026 +0000
CUDA: pass src1 precision to host MMQ config helpers (#30168)
* CUDA: pass src1 precision to host MMQ config helpers
Signed-off-by: ynankani <ynankani@nvidia.com>
* make src1 prec explicit
Signed-off-by: ynankani <ynankani@nvidia.com>
---------
Signed-off-by: ynankani <ynankani@nvidia.com>
diff --git a/ggml/src/ggml-cuda/mmq-load-tiles.cuh b/ggml/src/ggml-cuda/mmq-load-tiles.cuh
index d5654868b..0e59c0347 100644
--- a/ggml/src/ggml-cuda/mmq-load-tiles.cuh
+++ b/ggml/src/ggml-cuda/mmq-load-tiles.cuh
@@ -7,9 +7,9 @@
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q1_0(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -98,9 +98,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_0(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -187,9 +187,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_0(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -250,9 +250,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_1(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -313,9 +313,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_0(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -393,9 +393,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_1(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -471,9 +471,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q8_0(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -537,9 +537,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_K(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -598,9 +598,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q3_K(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -711,9 +711,9 @@ static __device__ __forceinline__ int unpack_scales_q45_K(const int * scales, co
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_K(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -822,9 +822,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_K(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -946,9 +946,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q6_K(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1036,9 +1036,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq1_s(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1098,9 +1098,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_xxs(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1162,9 +1162,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_xs(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1227,9 +1227,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_s(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1295,9 +1295,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq3_xxs(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1359,9 +1359,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq3_s(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1428,9 +1428,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq4_xs(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1495,9 +1495,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq4_nl(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1564,9 +1564,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_mxfp4(
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);
+ constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
int * x_qs = (int *) x_tile;
@@ -1670,7 +1670,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
const char * __restrict__ x, int * __restrict__ x_tile, const int kb0, 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, prec_src1) / warp_size;
diff --git a/ggml/src/ggml-cuda/mmq-vec-dot.cuh b/ggml/src/ggml-cuda/mmq-vec-dot.cuh
index b39e4d579..a75dd046e 100644
--- a/ggml/src/ggml-cuda/mmq-vec-dot.cuh
+++ b/ggml/src/ggml-cuda/mmq-vec-dot.cuh
@@ -10,8 +10,8 @@ using namespace ggml_cuda_mma;
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_0_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_0, I);
const int * x_qs = (const int *) x;
@@ -60,8 +60,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_1_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_1, I);
const int * x_qs = (const int *) x;
@@ -110,8 +110,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q8_0, I);
const int * x_qs = (const int *) x;
@@ -148,8 +148,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
typedef tile<16, 8, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -203,8 +203,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
typedef tile< 8, 8, int> tile_B;
typedef tile<16, 8, int> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -281,8 +281,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q5_1, I);
const int * x_qs = (const int *) x;
@@ -318,8 +318,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 8, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -368,8 +368,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile< 8, 8, int> tile_B;
typedef tile<16, 8, int> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -442,8 +442,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(type, I);
const int * x_qs = (const int *) x;
@@ -474,7 +474,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
// Used for Q3_K, IQ2_S, and IQ2_XS:
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
#if defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
constexpr data_layout input_layout = get_input_data_layout();
@@ -483,7 +483,7 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -533,7 +533,7 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
typedef tile<16, 8, int> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -610,8 +610,8 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q2_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q2_K, I);
const int * x_qs = (const int *) x;
@@ -680,8 +680,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 4, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -749,8 +749,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile< 8, 4, int> tile_B;
typedef tile<16, 8, int> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -870,8 +870,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q3_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q3_K, I);
const int * x_qs = (const int *) x;
@@ -905,8 +905,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_K, I);
const int * x_qs = (const int *) x;
@@ -940,8 +940,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q5_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q5_K, I);
const int * x_qs = (const int *) x;
@@ -975,8 +975,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q6_K_q8_1_dp4a(
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q6_K, I);
const int * x_qs = (const int *) x;
@@ -1015,8 +1015,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 4, int, input_layout> tile_B;
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -1066,8 +1066,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile< 8, 4, int> tile_B;
typedef tile<16, 8, int> tile_C;
- constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
@@ -1181,7 +1181,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
typedef tile<16, 8, float> tile_C;
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q4);
constexpr int ntx = rows_per_warp / tile_C::I;
constexpr int nfrags = MMQ_TILE_NE_K / tile_A::J;
diff --git a/ggml/src/ggml-cuda/mmq.cu b/ggml/src/ggml-cuda/mmq.cu
index 027a7e60d..91b89999e 100644
--- a/ggml/src/ggml-cuda/mmq.cu
+++ b/ggml/src/ggml-cuda/mmq.cu
@@ -8,66 +8,66 @@
static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream, const ggml_prec prec_src1) {
switch (args.type_x) {
case GGML_TYPE_Q1_0:
- mul_mat_q_case<GGML_TYPE_Q1_0>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q1_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q2_0:
- mul_mat_q_case<GGML_TYPE_Q2_0>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q2_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q4_0:
- mul_mat_q_case<GGML_TYPE_Q4_0>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q4_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q4_1:
- mul_mat_q_case<GGML_TYPE_Q4_1>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q4_1, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q5_0:
- mul_mat_q_case<GGML_TYPE_Q5_0>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q5_0, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q5_1:
- mul_mat_q_case<GGML_TYPE_Q5_1>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q5_1, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q8_0:
- mul_mat_q_case<GGML_TYPE_Q8_0>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q8_0, GGML_PREC_Q8>(ctx, args, stream);
break;
// -----------------------------------------------------------------------
case GGML_TYPE_Q2_K:
- mul_mat_q_case<GGML_TYPE_Q2_K>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q2_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q3_K:
- mul_mat_q_case<GGML_TYPE_Q3_K>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q3_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q4_K:
- mul_mat_q_case<GGML_TYPE_Q4_K>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q4_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q5_K:
- mul_mat_q_case<GGML_TYPE_Q5_K>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q5_K, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_Q6_K:
- mul_mat_q_case<GGML_TYPE_Q6_K>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_Q6_K, GGML_PREC_Q8>(ctx, args, stream);
break;
// -----------------------------------------------------------------------
case GGML_TYPE_IQ1_S:
- mul_mat_q_case<GGML_TYPE_IQ1_S>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ1_S, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ2_XXS:
- mul_mat_q_case<GGML_TYPE_IQ2_XXS>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ2_XXS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ2_XS:
- mul_mat_q_case<GGML_TYPE_IQ2_XS>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ2_XS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ2_S:
- mul_mat_q_case<GGML_TYPE_IQ2_S>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ2_S, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ3_XXS:
- mul_mat_q_case<GGML_TYPE_IQ3_XXS>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ3_XXS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ3_S:
- mul_mat_q_case<GGML_TYPE_IQ3_S>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ3_S, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ4_XS:
- mul_mat_q_case<GGML_TYPE_IQ4_XS>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ4_XS, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_IQ4_NL:
- mul_mat_q_case<GGML_TYPE_IQ4_NL>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_IQ4_NL, GGML_PREC_Q8>(ctx, args, stream);
break;
// -----------------------------------------------------------------------
case GGML_TYPE_MXFP4:
@@ -76,14 +76,14 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q4>(ctx, args, stream);
break;
}
- mul_mat_q_case<GGML_TYPE_MXFP4>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q8>(ctx, args, stream);
break;
case GGML_TYPE_NVFP4:
if (prec_src1 == GGML_PREC_Q4) {
mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q4>(ctx, args, stream);
break;
}
- mul_mat_q_case<GGML_TYPE_NVFP4>(ctx, args, stream);
+ mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q8>(ctx, args, stream);
break;
default:
GGML_ABORT("fatal error");
diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh
index 198aab4b5..d3e7b75ea 100644
--- a/ggml/src/ggml-cuda/mmq.cuh
+++ b/ggml/src/ggml-cuda/mmq.cuh
@@ -227,7 +227,7 @@ struct ggml_cuda_mmq_config {
#undef CASE
-static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1 = GGML_PREC_Q8) {
+static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
if (GGML_CUDA_CC_IS_AMD(cc)) {
if (GGML_CUDA_CC_IS_GCN(cc)) {
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
@@ -262,7 +262,7 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback);
}
-static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
#ifdef GGML_USE_HIP
#ifdef GCN
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
@@ -295,79 +295,77 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
GGML_UNUSED_VARS(type, J, fallback, prec_src1);
}
-// FIXME all of the host functions are missing prec_src1, this can lead to inconsitent behavior.
-
-static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).type;
+static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).type;
}
-static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type;
}
-static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).nthreads;
}
-static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).occupancy;
}
-static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).I;
+static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).I;
}
-static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).I;
}
-static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).J;
+static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).J;
}
-static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).J;
}
-static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).sram_layout;
+static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).sram_layout;
}
-static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).sram_layout;
}
-static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).K_vram;
+static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).K_vram;
}
-static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).K_vram;
}
-static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).stream_k;
+static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).stream_k;
}
-static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).stream_k;
}
-static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_config(type, J, fallback, cc).fallback;
+static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).fallback;
}
-static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).fallback;
}
// ---------------------------------------------------------------------------------------------
-static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const int J, const bool fallback, const int cc) {
- return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc));
+static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc, prec_src1));
}
-static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
+static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, prec_src1));
}
@@ -375,8 +373,8 @@ static __host__ bool ggml_cuda_mmq_needs_fallback(const int64_t nrows_x) {
return nrows_x % 128 != 0;
}
-static constexpr __device__ int ggml_cuda_mmq_get_rows_per_warp(ggml_type type, int J, bool fallback) {
- return ggml_cuda_mmq_get_config(type, J, fallback).rows_per_warp();
+static constexpr __device__ int ggml_cuda_mmq_get_rows_per_warp(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
+ return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).rows_per_warp();
}
#define MMQ_DP4A_TXS_Q4_0 tile_x_sizes{I*MMQ_TILE_NE_K + I, I*MMQ_TILE_NE_K/QI4_0 + I/QI4_0, 0}
@@ -432,12 +430,12 @@ static __host__ int ggml_cuda_mmq_get_nbytes_shared_x(const ggml_cuda_mmq_config
#include "mmq-load-tiles.cuh"
#include "mmq-vec-dot.cuh"
-template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_write_back_dp4a(
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_write_back_dp4a(
const float * __restrict__ sum, const int32_t * __restrict__ ids_dst, float * __restrict__ dst,
const float * __restrict__ y_scale, const int stride, const int i_max, const int j_max) {
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 nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
+ constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
const bool y_scale_used = y_scale != nullptr;
@@ -471,7 +469,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
}
}
-template<ggml_type type, int J, bool fallback>
+template<ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma(
const float * __restrict__ sum, const int * __restrict__ ids_dst, float * __restrict__ dst,
const float * __restrict__ y_scale, const int stride, const int i_max, const int j_max) {
@@ -482,7 +480,7 @@ static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma(
typedef tile<16, 8, int> tile_C;
#endif // defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
- constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
+ constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
const int i0 = (threadIdx.y / ntx) * (ntx*tile_C::I);
@@ -536,7 +534,7 @@ struct ggml_cuda_mmq_util_funcs {
vdr(vdr), load_tiles(load_tiles), vec_dot(vec_dot), write_back(write_back) {}
};
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_funcs() {
if (!ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).use_mma_data_layout()) {
switch (type) {
@@ -545,136 +543,136 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
VDR_Q1_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q1_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q2_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q2_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q2_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q4_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q4_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q4_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_1:
return ggml_cuda_mmq_util_funcs(
VDR_Q4_1_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q4_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q4_1_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q5_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q5_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_1:
return ggml_cuda_mmq_util_funcs(
VDR_Q5_1_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q5_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q8_0:
return ggml_cuda_mmq_util_funcs(
VDR_Q8_0_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q8_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_Q2_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q2_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q2_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q2_K_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q3_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q3_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q3_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q3_K_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q4_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q4_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q4_K_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q5_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q5_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q5_K_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_Q6_K:
return ggml_cuda_mmq_util_funcs(
VDR_Q6_K_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_q6_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q6_K_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_IQ1_S:
return ggml_cuda_mmq_util_funcs(
VDR_IQ1_S_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq1_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XXS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ2_XXS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq2_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ2_XS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq2_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_S:
return ggml_cuda_mmq_util_funcs(
VDR_IQ2_S_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq2_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_XXS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ3_XXS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq3_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_S:
return ggml_cuda_mmq_util_funcs(
VDR_IQ3_S_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq3_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_XS:
return ggml_cuda_mmq_util_funcs(
VDR_IQ4_XS_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq4_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_NL:
return ggml_cuda_mmq_util_funcs(
VDR_IQ4_NL_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_iq4_nl<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_MXFP4:
return ggml_cuda_mmq_util_funcs(
VDR_MXFP4_Q8_1_MMQ,
ggml_cuda_mmq_load_tiles_mxfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
case GGML_TYPE_NVFP4:
return ggml_cuda_mmq_util_funcs(
VDR_NVFP4_Q8_1_MMQ,
- ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback>,
+ ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
- ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
+ ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
default:
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
}
@@ -690,7 +688,7 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
-1,
ggml_cuda_mmq_load_tiles_mxfp4_fp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
}
break;
case GGML_TYPE_NVFP4:
@@ -699,7 +697,7 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
-1,
ggml_cuda_mmq_load_tiles_nvfp4_nvfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
}
break;
default:
@@ -715,164 +713,164 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
-1,
ggml_cuda_mmq_load_tiles_q1_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q2_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q2_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q4_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_DS4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_1:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q4_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q5_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_1:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q5_1<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q8_0:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q8_0<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_Q2_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q2_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q2_K_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q3_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q3_K<type, J, fallback>,
- ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q4_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q4_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q5_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q5_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_Q6_K:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_q6_K<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q6_K_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_IQ1_S:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq1_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XXS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq2_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_XS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq2_xs<type, J, fallback>,
- ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ2_S:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq2_s<type, J, fallback>,
- ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_XXS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq3_xxs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ3_S:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq3_s<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_XS:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq4_xs<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_IQ4_NL:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_iq4_nl<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
// ---------------------------------------------------------------------------------------------
case GGML_TYPE_MXFP4:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_mxfp4<type, J, fallback>,
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
case GGML_TYPE_NVFP4:
return ggml_cuda_mmq_util_funcs(
-1,
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
- ggml_cuda_mmq_write_back_mma<type, J, fallback>);
+ ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
default:
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
}
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ int ggml_cuda_mmq_get_vdr() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vdr;
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_load_tiles_t ggml_cuda_mmq_get_load_tiles() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().load_tiles;
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_vec_dot_t ggml_cuda_mmq_get_vec_dot() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vec_dot;
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static constexpr __device__ ggml_cuda_mmq_write_back_t ggml_cuda_mmq_get_write_back() {
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().write_back;
}
// ---------------------------------------------------------------------------------------------
-template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1>
static __device__ __forceinline__ void mul_mat_q_process_tile(
const char * __restrict__ x, const int offset_x, const int * __restrict__ y,
const int * __restrict__ ids_dst, float * __restrict__ dst, float * __restrict__ tmp_fixup,
@@ -953,7 +951,7 @@ static __device__ __forceinline__ void mul_mat_q_process_tile(
// The mul_mat_q kernel implements "stream-k" work partitioning as described in https://arxiv.org/abs/2301.03598
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1), ggml_cuda_mmq_get_occupancy(type, J, fallback, prec_src1))
static __global__ void mul_mat_q(
const char * __restrict__ x, const int * __restrict__ y, const int32_t * __restrict__ ids_dst,
@@ -1240,7 +1238,7 @@ static __global__ void mul_mat_q(
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1)/2, 1)
static __global__ void mul_mat_q_stream_k_fixup(
const int32_t * __restrict__ ids_dst, const int32_t * __restrict__ expert_bounds, float * __restrict__ dst,
@@ -1395,7 +1393,7 @@ static size_t mmq_get_nbytes_shared(const ggml_cuda_mmq_config & config, const i
return nbs_ids + nbs_x + GGML_PAD(nbs_y, config.nthreads*sizeof(int));
}
-template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
const int id = ggml_cuda_get_device();
const int cc = ggml_cuda_info().devices[id].cc;
@@ -1477,7 +1475,7 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
ntx_fd);
}
-template <ggml_type type, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, bool fallback, ggml_prec prec_src1>
void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
switch (args.J_best) {
case 8:
@@ -1535,7 +1533,7 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
}
}
-template <ggml_type type, ggml_prec prec_src1 = GGML_PREC_Q8>
+template <ggml_type type, ggml_prec prec_src1>
void mul_mat_q_case(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
if (ggml_cuda_mmq_needs_fallback(args.nrows_x)) {
constexpr bool fallback = true;
@@ -1547,7 +1545,7 @@ void mul_mat_q_case(ggml_backend_cuda_context & ctx, const mmq_args & args, cuda
}
#define DECL_MMQ_CASE(type) \
- template void mul_mat_q_case<type>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
+ template void mul_mat_q_case<type, GGML_PREC_Q8>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
// FP4 variant: uses native FP4 MMA instead of keeping src1 at Q8_1.
#define DECL_MMQ_CASE_W4A4(type) \