# llama.cpp gfx1151 optimizations for Qwen3.8 27B

> Source: <https://gist.github.com/pedapudi/f414fbcd838610cc6bd8319c11a6cc87>
> Published: 2026-08-16 04:59:27+00:00

| diff --git a/ggml/src/ggml-cuda/concat.cu b/ggml/src/ggml-cuda/concat.cu | |
| index 6df89013c..e5337ec30 100644 | |
| --- a/ggml/src/ggml-cuda/concat.cu | |
| +++ b/ggml/src/ggml-cuda/concat.cu | |
| @@ -79,6 +79,76 @@ static void concat_cont_cuda(const T * x, | |
| concat_cont<T, 2><<<num_blocks, CUDA_CONCAT_BLOCK_SIZE, 0, stream>>>(x, y, dst, ne00, ne01, ne02, ne0, ne1, ne2); | |
| } | |
| +// Fast LDS-transpose path for a dim-0 concat whose src1 is transposed-contiguous | |
| +// (nb[1] == element size < nb[0]). That is the qwen35 ssm_conv input pattern: | |
| +// concat(conv_state[3, d_inner], transpose(x)[n_tokens, d_inner]) along dim 0. | |
| +// | |
| +// The generic non-cont kernel walks i0 with consecutive threads, so for a | |
| +// transposed src1 every lane of a wave reads a different row -> one memory | |
| +// transaction per lane on both the load and the store side. Staging a 32x32 | |
| +// tile through LDS makes both sides fully coalesced. | |
| +#define CONCAT_TRANSPOSE_TILE 64 // 64x64 measured best on gfx1151 (35 -> 70 GB/s on the qwen35 shape) | |
| +#define CONCAT_TRANSPOSE_NTX 64 | |
| +#define CONCAT_TRANSPOSE_NTY 4 | |
| + | |
| +template <typename T> | |
| +static __global__ void __launch_bounds__(CONCAT_TRANSPOSE_NTX * CONCAT_TRANSPOSE_NTY) | |
| + concat_dim0_transposed_src1( | |
| + const char * __restrict__ src0, | |
| + const char * __restrict__ src1, | |
| + char * __restrict__ dst, | |
| + int64_t ne00, | |
| + int64_t ne0, | |
| + int64_t ne1, | |
| + int64_t ne2, | |
| + uint64_t nb00, uint64_t nb01, uint64_t nb02, uint64_t nb03, | |
| + uint64_t nb10, uint64_t nb11, uint64_t nb12, uint64_t nb13, | |
| + uint64_t nb0, uint64_t nb1, uint64_t nb2, uint64_t nb3) { | |
| + constexpr int TILE = CONCAT_TRANSPOSE_TILE; | |
| + constexpr int NTX = CONCAT_TRANSPOSE_NTX; | |
| + constexpr int NTY = CONCAT_TRANSPOSE_NTY; | |
| + __shared__ T tile[TILE][TILE + 1]; // +1 to avoid LDS bank conflicts on the transposed read | |
| + | |
| + const int tx = threadIdx.x; | |
| + const int ty = threadIdx.y; | |
| + const int i23 = blockIdx.z; | |
| + const int i2 = i23 % (int) ne2; | |
| + const int i3 = i23 / (int) ne2; | |
| + const int64_t i0_base = (int64_t) blockIdx.x * TILE; | |
| + const int64_t i1_base = (int64_t) blockIdx.y * TILE; | |
| + | |
| + // load: consecutive tx -> consecutive i1, i.e. consecutive addresses in src1 | |
| +#pragma unroll | |
| + for (int c = 0; c < TILE; c += NTX) { | |
| + const int64_t load_i1 = i1_base + c + tx; | |
| + for (int r = ty; r < TILE; r += NTY) { | |
| + const int64_t load_i0 = i0_base + r; | |
| + T v = T(0); | |
| + if (load_i0 < ne0 && load_i1 < ne1) { | |
| + if (load_i0 < ne00) { | |
| + v = *(const T *)(src0 + i3*nb03 + i2*nb02 + load_i1*nb01 + load_i0*nb00); | |
| + } else { | |
| + v = *(const T *)(src1 + i3*nb13 + i2*nb12 + (load_i0 - ne00)*nb10 + load_i1*nb11); | |
| + } | |
| + } | |
| + tile[r][c + tx] = v; | |
| + } | |
| + } | |
| + __syncthreads(); | |
| + | |
| + // store: consecutive tx -> consecutive i0, i.e. consecutive addresses in dst | |
| +#pragma unroll | |
| + for (int c = 0; c < TILE; c += NTX) { | |
| + const int64_t store_i0 = i0_base + c + tx; | |
| + for (int r = ty; r < TILE; r += NTY) { | |
| + const int64_t store_i1 = i1_base + r; | |
| + if (store_i0 < ne0 && store_i1 < ne1) { | |
| + *(T *)(dst + i3*nb3 + i2*nb2 + store_i1*nb1 + store_i0*nb0) = tile[c + tx][r]; | |
| + } | |
| + } | |
| + } | |
| +} | |
| + | |
| // non-contiguous kernel (slow) | |
| template <typename T, int dim> | |
| static __global__ void __launch_bounds__(CUDA_CONCAT_BLOCK_SIZE) | |
| @@ -164,6 +234,25 @@ static void concat_cuda(const ggml_tensor * src0, const ggml_tensor * src1, ggml | |
| GGML_ASSERT(!ggml_is_quantized(src0->type)); | |
| dim3 grid_dim(dst->ne[1], dst->ne[2], dst->ne[3]); | |
| + // transposed-src1 dim-0 concat (qwen35 ssm_conv input): stage through LDS | |
| + if (dim == 0 | |
| + && src1->nb[1] == sizeof(T) && src1->nb[0] > src1->nb[1] | |
| + && src0->nb[0] == sizeof(T) | |
| + && src0->ne[1] == dst->ne[1] && src1->ne[1] == dst->ne[1]) { | |
| + constexpr int TILE = CONCAT_TRANSPOSE_TILE; | |
| + const dim3 grid((dst->ne[0] + TILE - 1) / TILE, | |
| + (dst->ne[1] + TILE - 1) / TILE, | |
| + (unsigned) (dst->ne[2] * dst->ne[3])); | |
| + const dim3 block(CONCAT_TRANSPOSE_NTX, CONCAT_TRANSPOSE_NTY, 1); | |
| + concat_dim0_transposed_src1<T><<<grid, block, 0, stream>>>( | |
| + (const char *) src0->data, (const char *) src1->data, (char *) dst->data, | |
| + src0->ne[0], dst->ne[0], dst->ne[1], dst->ne[2], | |
| + src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], | |
| + src1->nb[0], src1->nb[1], src1->nb[2], src1->nb[3], | |
| + dst->nb[0], dst->nb[1], dst->nb[2], dst->nb[3]); | |
| + return; | |
| + } | |
| + | |
| auto launch_kernel = [&](auto dim) { | |
| concat_non_cont<T, dim><<<grid_dim, CUDA_CONCAT_BLOCK_SIZE, 0, stream>>>( | |
| (const char *) src0->data, (const char *) src1->data, (char *) dst->data, | |
| diff --git a/ggml/src/ggml-cuda/fattn-tile.cuh b/ggml/src/ggml-cuda/fattn-tile.cuh | |
| index d1164b852..6611f8096 100644 | |
| --- a/ggml/src/ggml-cuda/fattn-tile.cuh | |
| +++ b/ggml/src/ggml-cuda/fattn-tile.cuh | |
| @@ -312,8 +312,25 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_am | |
| return 0; | |
| } | |
| +// gfx1151, DKQ = DV = 256 (the qwen35 attention head): nbatch_fa 64 -> 128. | |
| +// | |
| +// The generic RDNA table's 64 is the right answer only while the kernel is spilling: with | |
| +// the staging loops fully unrolled this shape sits at 256 VGPRs with 5 spilled, and every | |
| +// config that wants more registers is punished for that rather than on its merits. Once the | |
| +// unroll is capped (see GGML_FA_TILE_UNROLL_STAGE) the optimum moves -- 128 was the *worst* | |
| +// point of the original sweep and is the best one now. LDS goes 37888 -> ~33800 B, still | |
| +// three workgroups per WGP. Measured 8.65 -> 9.40 TFLOPS at 512x4096 and 8.93 -> 9.18 at | |
| +// 512x33280. | |
| +static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_amd_rdna3_5(const int DKQ, const int DV, const int ncols) { | |
| + GGML_CUDA_FATTN_TILE_CONFIG_CASE(256, 256, 32, 256, 2, 128, 64) | |
| + return ggml_cuda_fattn_tile_get_config_amd_rdna(DKQ, DV, ncols); | |
| +} | |
| + | |
| static __host__ uint32_t ggml_cuda_fattn_tile_get_config(const int DKQ, const int DV, const int ncols, const int cc) { | |
| if (GGML_CUDA_CC_IS_AMD(cc)) { | |
| + if (GGML_CUDA_CC_IS_RDNA3_5(cc)) { | |
| + return ggml_cuda_fattn_tile_get_config_amd_rdna3_5(DKQ, DV, ncols); | |
| + } | |
| if (GGML_CUDA_CC_IS_RDNA(cc)) { | |
| return ggml_cuda_fattn_tile_get_config_amd_rdna(DKQ, DV, ncols); | |
| } | |
| @@ -327,7 +344,9 @@ static __host__ uint32_t ggml_cuda_fattn_tile_get_config(const int DKQ, const in | |
| static constexpr __device__ uint32_t ggml_cuda_fattn_tile_get_config(const int DKQ, const int DV, const int ncols) { | |
| #ifdef GGML_USE_HIP | |
| -#ifdef RDNA | |
| +#ifdef RDNA3_5 | |
| + return ggml_cuda_fattn_tile_get_config_amd_rdna3_5(DKQ, DV, ncols); | |
| +#elif defined(RDNA) | |
| return ggml_cuda_fattn_tile_get_config_amd_rdna(DKQ, DV, ncols); | |
| #else | |
| return ggml_cuda_fattn_tile_get_config_amd(DKQ, DV, ncols); | |
| @@ -482,6 +501,19 @@ static __device__ __forceinline__ void flash_attn_tile_load_tile( | |
| // Function that performs a single iteration in for the KQ matrix multiplication: | |
| template <int warp_size, int nwarps, int ncols1, int ncols2, int DKQ, int nbatch_fa, int nbatch_K, | |
| bool use_logit_softcap, bool oob_check, typename T_vec_dot> | |
| +// gfx1151: these two loops stage 24 half2 of K/Q (resp. V/KQ) per iteration, and fully | |
| +// unrolling them lets the scheduler keep several iterations of loads in flight. Combined | |
| +// with this target's raised -amdgpu-unroll-threshold-local that pushes the DKQ=256 kernel | |
| +// to 256 VGPRs with 5 spilled, and the spills then stop it keeping *any* further loads in | |
| +// flight -- the kernel runs at 45% instruction-issue utilisation. Capping the unroll takes | |
| +// it to 127 VGPRs with no spills and +23% at 32k context. Measured on RDNA3.5 only, so it | |
| +// is scoped there; other targets keep the full unroll. | |
| +#if defined(GGML_USE_HIP) && defined(RDNA3_5) | |
| +#define GGML_FA_TILE_UNROLL_STAGE _Pragma("unroll 4") | |
| +#else | |
| +#define GGML_FA_TILE_UNROLL_STAGE _Pragma("unroll") | |
| +#endif | |
| + | |
| static __device__ __forceinline__ void flash_attn_tile_iter_KQ( | |
| T_vec_dot * const Q_tmp, | |
| const half2 * const __restrict__ K_h2, | |
| @@ -504,7 +536,7 @@ static __device__ __forceinline__ void flash_attn_tile_iter_KQ( | |
| #ifdef FAST_FP16_AVAILABLE | |
| static_assert((nbatch_K/2) % cpy_ne == 0, "bad nbatch_K"); | |
| -#pragma unroll | |
| +GGML_FA_TILE_UNROLL_STAGE | |
| for (int k_KQ_1 = 0; k_KQ_1 < nbatch_K/2; k_KQ_1 += cpy_ne) { | |
| __align__(16) half2 K_k[nbatch_fa/(np*warp_size)][cpy_ne]; | |
| __align__(16) half2 Q_k[cpw][cpy_ne]; | |
| @@ -723,7 +755,7 @@ static __device__ __forceinline__ void flash_attn_tile_iter( | |
| __syncthreads(); | |
| #ifdef FAST_FP16_AVAILABLE | |
| -#pragma unroll | |
| +GGML_FA_TILE_UNROLL_STAGE | |
| for (int k1 = 0; k1 < nbatch_V; k1 += np) { | |
| __align__(16) half2 V_k[(DVp/2)/warp_size]; | |
| __align__(16) half2 KQ_k[cpw]; | |
| @@ -755,7 +787,7 @@ static __device__ __forceinline__ void flash_attn_tile_iter( | |
| } | |
| } | |
| #else | |
| -#pragma unroll | |
| +GGML_FA_TILE_UNROLL_STAGE | |
| for (int k1 = 0; k1 < nbatch_V; k1 += np) { | |
| __align__(16) float2 V_k[(DVp/2)/warp_size]; | |
| __align__(16) float KQ_k[cpw]; | |
| diff --git a/ggml/src/ggml-cuda/fattn-wmma-d256.cu b/ggml/src/ggml-cuda/fattn-wmma-d256.cu | |
| new file mode 100644 | |
| index 000000000..908eb00bb | |
| --- /dev/null | |
| +++ b/ggml/src/ggml-cuda/fattn-wmma-d256.cu | |
| @@ -0,0 +1,55 @@ | |
| +#include "fattn-wmma-d256.cuh" | |
| +#include <cstdlib> | |
| + | |
| +template<int ncols1, int ncols2, ggml_type type_K, ggml_type type_V> | |
| +static void launch_wmma_d256(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { | |
| + constexpr int DKQ = 256; | |
| + constexpr int DV = 256; | |
| + launch_fattn<DV, ncols1, ncols2>( | |
| + ctx, dst, flash_attn_wmma_d256<DKQ, DV, ncols1, ncols2, type_K, type_V>, | |
| + /*nwarps =*/ FA_WMMA_D256_NWARPS, | |
| + /*nbytes_shared=*/ getenv("GGML_FA_LDS_ABLATE") | |
| + ? (size_t) atoi(getenv("GGML_FA_LDS_ABLATE")) : ggml_cuda_fattn_wmma_d256_nbytes_shared(), | |
| + /*nbatch_fa =*/ FA_WMMA_D256_FA, | |
| + // false: the kernel dequantises K/V itself, so launch_fattn must not pre-convert | |
| + // the cache to f16 -- that conversion is exactly the DRAM traffic we are avoiding. | |
| + /*need_f16_K =*/ false, | |
| + /*need_f16_V =*/ false, | |
| + /*stream_k =*/ false, | |
| + /*warp_size =*/ 32); | |
| +} | |
| + | |
| +void ggml_cuda_flash_attn_ext_wmma_d256(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { | |
| + const ggml_tensor * Q = dst->src[0]; | |
| + const ggml_tensor * K = dst->src[1]; | |
| + | |
| + GGML_ASSERT(Q->ne[0] == 256); | |
| + | |
| + const int gqa_ratio = Q->ne[2] / K->ne[2]; | |
| + | |
| + // Widest power-of-two head group that divides the GQA ratio, capped so that | |
| + // ncols1 = 64/ncols2 still leaves at least one Q column per block. | |
| + int ncols2 = 1; | |
| + while (gqa_ratio % (2*ncols2) == 0 && ncols2 < 8) { | |
| + ncols2 *= 2; | |
| + } | |
| + | |
| + constexpr int NC = FA_WMMA_D256_NC; | |
| + const ggml_type tK = K->type; | |
| + const ggml_type tV = dst->src[2]->type; | |
| + | |
| +#define FA_WMMA_DISPATCH_KV(N1, N2) \ | |
| + if (tK == GGML_TYPE_F16 && tV == GGML_TYPE_F16 ) { launch_wmma_d256<N1, N2, GGML_TYPE_F16, GGML_TYPE_F16 >(ctx, dst); return; } \ | |
| + if (tK == GGML_TYPE_Q8_0 && tV == GGML_TYPE_Q8_0) { launch_wmma_d256<N1, N2, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0>(ctx, dst); return; } \ | |
| + if (tK == GGML_TYPE_Q4_0 && tV == GGML_TYPE_Q4_0) { launch_wmma_d256<N1, N2, GGML_TYPE_Q4_0, GGML_TYPE_Q4_0>(ctx, dst); return; } \ | |
| + GGML_ABORT("unsupported KV type combination"); | |
| + | |
| + switch (ncols2) { | |
| + case 1: FA_WMMA_DISPATCH_KV(NC/1, 1); break; | |
| + case 2: FA_WMMA_DISPATCH_KV(NC/2, 2); break; | |
| + case 4: FA_WMMA_DISPATCH_KV(NC/4, 4); break; | |
| + case 8: FA_WMMA_DISPATCH_KV(NC/8, 8); break; | |
| + default: GGML_ABORT("fatal error"); | |
| + } | |
| +#undef FA_WMMA_DISPATCH_KV | |
| +} | |
| diff --git a/ggml/src/ggml-cuda/fattn-wmma-d256.cuh b/ggml/src/ggml-cuda/fattn-wmma-d256.cuh | |
| new file mode 100644 | |
| index 000000000..88187a552 | |
| --- /dev/null | |
| +++ b/ggml/src/ggml-cuda/fattn-wmma-d256.cuh | |
| @@ -0,0 +1,426 @@ | |
| +#pragma once | |
| + | |
| +#include "common.cuh" | |
| +#include "fattn-common.cuh" | |
| + | |
| +// WMMA flash attention for RDNA3/3.5 at head size 256. | |
| +// | |
| +// The generic tile kernel emits no WMMA at all, and fattn-mma-f16 gives every warp the | |
| +// full DV in registers, which at DV=256 spills >1 kB/lane to scratch on a 256-VGPR wave. | |
| +// This kernel splits DV across warps instead, so the O accumulator is 64 VGPRs/lane. | |
| +// | |
| +// Fixed shape: 8 warps, 64 columns per block, 32 KV rows per iteration. | |
| +// NQB = 4 q blocks, NKB = 2 kv blocks, NQB*NKB == 8 -> one 16x16 S tile per warp. | |
| +// O phase: warp w owns q block w/NKB and DV slice [(w%NKB)*128, +128) = 8 tiles. | |
| +// | |
| +// WMMA fragment layout (wave32, gfx11): | |
| +// A: lane t supplies row (t % 16) of A, all 16 k | |
| +// B: lane t supplies column (t % 16) of B, all 16 k | |
| +// D: lane t holds COLUMN (t % 16); its 8 slots are rows 2*l + t/16 | |
| + | |
| +#define FA_WMMA_D256_D 256 | |
| +#ifndef FA_WMMA_D256_NC | |
| +#define FA_WMMA_D256_NC 64 | |
| +#endif | |
| +#ifndef FA_WMMA_D256_FA | |
| +#define FA_WMMA_D256_FA 32 | |
| +#endif | |
| +#define FA_WMMA_D256_NQB (FA_WMMA_D256_NC / 16) | |
| +#define FA_WMMA_D256_NKB (FA_WMMA_D256_FA / 16) | |
| +// KSPLIT warps cooperate on each S tile by splitting the head dimension. It buys waves per | |
| +// workgroup without any extra LDS, which is what this kernel is short of: the Q tile alone | |
| +// is 33 kB, so only 2 workgroups fit per WGP however the rest is arranged. | |
| +#define FA_WMMA_D256_KSPLIT 1 | |
| +#define FA_WMMA_D256_NWARPS (FA_WMMA_D256_NQB * FA_WMMA_D256_NKB * FA_WMMA_D256_KSPLIT) | |
| +#define FA_WMMA_D256_NSW (FA_WMMA_D256_NKB * FA_WMMA_D256_KSPLIT) // warps sharing a q block | |
| +#define FA_WMMA_D256_DVW (FA_WMMA_D256_D / FA_WMMA_D256_NSW) | |
| +#define FA_WMMA_D256_OT (FA_WMMA_D256_DVW / 16) | |
| +#define FA_WMMA_D256_KPAD 8 | |
| +#define FA_WMMA_D256_FPAD 8 | |
| + | |
| +static constexpr size_t ggml_cuda_fattn_wmma_d256_nbytes_shared() { | |
| + // Q is staged through LDS once and then held in registers, so its footprint overlaps | |
| + // the K/V/P working set instead of adding to it. Keeping Q resident would put the | |
| + // block at 60 kB and 2 workgroups per WGP; this way it is 3, which measured -20% on | |
| + // the long-context shape. | |
| + const size_t nb_Q = (size_t) FA_WMMA_D256_NC * (FA_WMMA_D256_D + FA_WMMA_D256_KPAD) * 2; | |
| + const size_t nb_K = (size_t) FA_WMMA_D256_FA * (FA_WMMA_D256_D + FA_WMMA_D256_KPAD) * 2; | |
| + const size_t nb_V = (size_t) FA_WMMA_D256_D * (FA_WMMA_D256_FA + FA_WMMA_D256_FPAD) * 2; | |
| + const size_t nb_S = (size_t) FA_WMMA_D256_NQB * FA_WMMA_D256_NKB * 32 * 8 * sizeof(float); | |
| + size_t nb_KV = nb_K > nb_V ? nb_K : nb_V; | |
| + if (nb_S > nb_KV) { nb_KV = nb_S; } | |
| + const size_t nb_P = (size_t) FA_WMMA_D256_NC * (FA_WMMA_D256_FA + FA_WMMA_D256_FPAD) * 2; | |
| + const size_t nb_R = (size_t) 2 * FA_WMMA_D256_NQB * FA_WMMA_D256_NKB * 16 * sizeof(float); | |
| + return nb_Q + nb_KV + nb_P + nb_R; | |
| +} | |
| + | |
| +#if defined(AMD_WMMA_AVAILABLE) && defined(RDNA3) | |
| + | |
| +typedef _Float16 fa_half16 __attribute__((ext_vector_type(16))); | |
| +typedef float fa_float8 __attribute__((ext_vector_type(8))); | |
| + | |
| +union fa_frag16 { fa_half16 h; int4 v[2]; }; | |
| + | |
| +static __device__ __forceinline__ fa_half16 fa_load_frag(const _Float16 * p) { | |
| + fa_frag16 f; | |
| + f.v[0] = *(const int4 *) (p + 0); | |
| + f.v[1] = *(const int4 *) (p + 8); | |
| + return f.h; | |
| +} | |
| + | |
| +// Read 8 consecutive head-dimension elements of one K/V row, dequantising on the way in. | |
| +// The point of doing this here rather than letting launch_fattn pre-convert the cache to | |
| +// f16 is that the DRAM traffic is what bounds this kernel: at head size 256 the arithmetic | |
| +// intensity is (columns per block)/2 MAC per byte, so halving the bytes doubles the roofline. | |
| +template <ggml_type type_KV> | |
| +static __device__ __forceinline__ int4 fa_load8(const char * __restrict__ row, int i0) { | |
| + int4 out; | |
| + _Float16 * o = (_Float16 *) &out; | |
| + if constexpr (type_KV == GGML_TYPE_F16) { | |
| + out = *(const int4 *) (row + i0*2); | |
| + } else if constexpr (type_KV == GGML_TYPE_Q8_0) { | |
| + const block_q8_0 * b = (const block_q8_0 *) row + i0/QK8_0; | |
| + const float d = __half2float(b->d); | |
| + const int j0 = i0 % QK8_0; | |
| +#pragma unroll | |
| + for (int j = 0; j < 8; ++j) { o[j] = (_Float16) (d * (float) b->qs[j0 + j]); } | |
| + } else { | |
| + static_assert(type_KV == GGML_TYPE_Q4_0, "unsupported KV type"); | |
| + const block_q4_0 * b = (const block_q4_0 *) row + i0/QK4_0; | |
| + const float d = __half2float(b->d); | |
| + const int j0 = i0 % QK4_0; | |
| + const int sh = (j0 / 16) * 4; // low nibbles hold the first half of the block | |
| + const uint8_t * q = b->qs + (j0 % 16); | |
| +#pragma unroll | |
| + for (int j = 0; j < 8; ++j) { o[j] = (_Float16) (d * (float) (int) (((q[j] >> sh) & 0x0F) - 8)); } | |
| + } | |
| + return out; | |
| +} | |
| + | |
| +// Single element, for the transposed V staging where each lane walks kv, not the head dim. | |
| +template <ggml_type type_KV> | |
| +static __device__ __forceinline__ _Float16 fa_load1(const char * __restrict__ row, int i) { | |
| + if constexpr (type_KV == GGML_TYPE_F16) { | |
| + return *(const _Float16 *) (row + i*2); | |
| + } else if constexpr (type_KV == GGML_TYPE_Q8_0) { | |
| + const block_q8_0 * b = (const block_q8_0 *) row + i/QK8_0; | |
| + return (_Float16) (__half2float(b->d) * (float) b->qs[i % QK8_0]); | |
| + } else { | |
| + static_assert(type_KV == GGML_TYPE_Q4_0, "unsupported KV type"); | |
| + const block_q4_0 * b = (const block_q4_0 *) row + i/QK4_0; | |
| + const int j = i % QK4_0; | |
| + return (_Float16) (__half2float(b->d) * (float) (int) (((b->qs[j % 16] >> ((j/16)*4)) & 0x0F) - 8)); | |
| + } | |
| +} | |
| + | |
| +template<int DKQ, int DV, int ncols1, int ncols2, ggml_type type_K, ggml_type type_V> | |
| +__launch_bounds__(FA_WMMA_D256_NWARPS*32, 1) | |
| +static __global__ void flash_attn_wmma_d256( | |
| + const char * Q_ptr, | |
| + const char * K_ptr, | |
| + const char * V_ptr, | |
| + const char * mask_ptr, | |
| + const char * sinks_ptr, | |
| + const int * KV_max_ptr, | |
| + float * dst_ptr, | |
| + float2 * dst_meta_ptr, | |
| + const float scale, | |
| + const float max_bias, | |
| + const float m0, | |
| + const float m1, | |
| + const uint32_t n_head_log2, | |
| + const float logit_softcap, | |
| + const int32_t ne00, const uint3 ne01, const int32_t ne02, const int32_t ne03, | |
| + const int32_t nb01, const int32_t nb02, const int32_t nb03, | |
| + const int32_t ne10, const int32_t ne11, const int32_t ne12, const int32_t ne13, | |
| + const int32_t nb11, const int32_t nb12, const int64_t nb13, | |
| + const int32_t nb21, const int32_t nb22, const int64_t nb23, | |
| + const int32_t ne31, const int32_t ne32, const int32_t ne33, | |
| + const int32_t nb31, const int32_t nb32, const int64_t nb33) { | |
| + constexpr int D = FA_WMMA_D256_D; | |
| + constexpr int NC = FA_WMMA_D256_NC; | |
| + constexpr int FAB = FA_WMMA_D256_FA; | |
| + constexpr int NWARPS = FA_WMMA_D256_NWARPS; | |
| + constexpr int NQB = FA_WMMA_D256_NQB; | |
| + constexpr int NKB = FA_WMMA_D256_NKB; | |
| + constexpr int KSPLIT = FA_WMMA_D256_KSPLIT; | |
| + constexpr int NSW = FA_WMMA_D256_NSW; | |
| + constexpr int DVW = FA_WMMA_D256_DVW; | |
| + constexpr int OT = FA_WMMA_D256_OT; | |
| + constexpr int KPAD = FA_WMMA_D256_KPAD; | |
| + constexpr int FPAD = FA_WMMA_D256_FPAD; | |
| + constexpr int WS = 32; | |
| + constexpr int NTHR = NWARPS*WS; | |
| + | |
| + static_assert(DKQ == D && DV == D, "this kernel is head size 256 only"); | |
| + static_assert(ncols1*ncols2 == NC, "bad column count"); | |
| + static_assert(NQB*NKB*KSPLIT == NWARPS, "bad warp count"); | |
| + static_assert(D % (16*KSPLIT) == 0, "bad head split"); | |
| + | |
| + GGML_UNUSED_VARS(max_bias, m0, m1, n_head_log2, logit_softcap, ne00, ne10, ne12, ne13, | |
| + ne31, ne32, nb32, nb33, sinks_ptr); | |
| + | |
| + const int tid = threadIdx.x + threadIdx.y*WS; | |
| + const int warp = threadIdx.y; | |
| + const int lane = threadIdx.x; | |
| + const int l16 = lane % 16; | |
| + const int lhi = lane / 16; | |
| + | |
| + const int col_Q_0 = blockIdx.x * ncols1; | |
| + | |
| + const int sequence = blockIdx.z / (ne02/ncols2); | |
| + const int head0 = blockIdx.z*ncols2 - sequence*ne02; | |
| + const int gqa_ratio = ne02 / ne12; | |
| + | |
| + const float * Q_f = (const float *) (Q_ptr + nb03*sequence + nb02* head0); | |
| + const char * K_base = K_ptr + nb13*sequence + nb12*(head0 / gqa_ratio); | |
| + const char * V_base = V_ptr + nb23*sequence + nb22*(head0 / gqa_ratio); | |
| + const half * maskh = mask_ptr ? (const half *) (mask_ptr + nb33*(sequence % ne33)) : nullptr; | |
| + | |
| + const int stride_mask = nb31 / sizeof(half); | |
| + | |
| + extern __shared__ char fa_smem[]; | |
| + _Float16 * Qs = (_Float16 *) fa_smem; // [NC][D+KPAD] | |
| + _Float16 * Ks = Qs + NC*(D + KPAD); // [FAB][D+KPAD] | |
| + _Float16 * Vt = Ks; // [D][FAB+FPAD], aliases Ks | |
| + constexpr size_t kv_halves = (size_t) FAB*(D + KPAD) > (size_t) D*(FAB + FPAD) | |
| + ? (size_t) FAB*(D + KPAD) : (size_t) D*(FAB + FPAD); | |
| + _Float16 * Ps = Ks + kv_halves; // [NC][FAB+FPAD] | |
| + float * Rm = (float *) (Ps + NC*(FAB + FPAD)); // [NQB*NKB][16] | |
| + float * Rl = Rm + NQB*NKB*16; // [NQB*NKB][16] | |
| + | |
| + ggml_cuda_pdl_sync(); | |
| + | |
| + // ---- Q into LDS, pre-scaled and converted to f16 ---- | |
| + // column jc -> q position col_Q_0 + jc/ncols2, head head0 + jc%ncols2 | |
| + for (int idx = tid*8; idx < NC*D; idx += NTHR*8) { | |
| + const int jc = idx / D; | |
| + const int i0 = idx % D; | |
| + const int j = jc / ncols2; | |
| + const int c = jc % ncols2; | |
| + const float * src = Q_f + c*(nb02/sizeof(float)) + fastmodulo(col_Q_0 + j, ne01)*(nb01/sizeof(float)) + i0; | |
| + fa_frag16 f; | |
| +#pragma unroll | |
| + for (int i = 0; i < 8; ++i) { | |
| + ((_Float16 *) &f.v[0])[i] = (_Float16) (src[i] * scale); | |
| + } | |
| + *(int4 *) &Qs[jc*(D + KPAD) + i0] = f.v[0]; | |
| + } | |
| + | |
| + const int s_tile = warp % (NQB*NKB); | |
| + const int k_half = warp / (NQB*NKB); | |
| + const int qb = s_tile / NKB; | |
| + const int s_kb = s_tile % NKB; | |
| + // warps sharing this q block: NSW of them, each taking one DV slice in the O phase | |
| + const int o_slot = (s_tile % NKB) + NKB*k_half; | |
| + const int o_d0 = o_slot * DVW; | |
| + | |
| + fa_float8 Oacc[OT]; | |
| +#pragma unroll | |
| + for (int i = 0; i < OT; ++i) { Oacc[i] = (fa_float8)(0.0f); } | |
| + | |
| + float m_run[8], l_run[8]; | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { m_run[l] = -FLT_MAX/2.0f; l_run[l] = 0.0f; } | |
| + | |
| + __syncthreads(); | |
| + | |
| + const int k_VKQ_max = KV_max_ptr ? KV_max_ptr[sequence*gridDim.x + blockIdx.x] : ne11; | |
| + | |
| + for (int k_VKQ_0 = blockIdx.y*FAB; k_VKQ_0 < k_VKQ_max; k_VKQ_0 += gridDim.y*FAB) { | |
| + const int kv_sup = k_VKQ_max - k_VKQ_0 < FAB ? k_VKQ_max - k_VKQ_0 : FAB; | |
| + | |
| + // ---- K into LDS ---- | |
| + for (int idx = tid*8; idx < FAB*D; idx += NTHR*8) { | |
| + const int kv = idx / D; | |
| + const int i0 = idx % D; | |
| + int4 v = make_int4(0, 0, 0, 0); | |
| + if (kv < kv_sup) { | |
| + v = fa_load8<type_K>(K_base + (size_t)(k_VKQ_0 + kv)*nb11, i0); | |
| + } | |
| + *(int4 *) &Ks[kv*(D + KPAD) + i0] = v; | |
| + } | |
| + __syncthreads(); | |
| + | |
| + // ---- S = Q * K^T, one tile per warp ---- | |
| + fa_float8 Sacc = (fa_float8)(0.0f); | |
| +#pragma unroll | |
| + for (int k0 = k_half*(D/KSPLIT); k0 < (k_half + 1)*(D/KSPLIT); k0 += 16) { | |
| + const fa_half16 a = fa_load_frag(&Qs[(qb *16 + l16)*(D + KPAD) + k0]); | |
| + const fa_half16 b = fa_load_frag(&Ks[(s_kb*16 + l16)*(D + KPAD) + k0]); | |
| + Sacc = __builtin_amdgcn_wmma_f32_16x16x16_f16_w32(a, b, Sacc); | |
| + } | |
| + // combine the KSPLIT partial head ranges; K is dead now so its buffer is scratch | |
| + if constexpr (KSPLIT > 1) { | |
| + float * Sred = (float *) Ks; | |
| + __syncthreads(); | |
| + if (k_half != 0) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Sred[(s_tile*32 + lane)*8 + l] = Sacc[l]; } | |
| + } | |
| + __syncthreads(); | |
| + if (k_half == 0) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Sacc[l] += Sred[(s_tile*32 + lane)*8 + l]; } | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Sred[(s_tile*32 + lane)*8 + l] = Sacc[l]; } | |
| + } | |
| + __syncthreads(); | |
| + if (k_half != 0) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Sacc[l] = Sred[(s_tile*32 + lane)*8 + l]; } | |
| + } | |
| + } | |
| + // slot l -> block column jc = qb*16 + 2*l + lhi, kv = s_kb*16 + l16 | |
| + | |
| + // ---- mask / out-of-bounds ---- | |
| + const int kv_loc = s_kb*16 + l16; | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { | |
| + const int jc = qb*16 + 2*l + lhi; | |
| + float s = Sacc[l]; | |
| + if (kv_loc >= kv_sup) { | |
| + s = -FLT_MAX/2.0f; | |
| + } else if (maskh) { | |
| + s += __half2float(maskh[(size_t)(col_Q_0 + jc/ncols2)*stride_mask + k_VKQ_0 + kv_loc]); | |
| + } | |
| + Sacc[l] = s; | |
| + } | |
| + | |
| + // ---- row max: reduce over the 16 lanes of the subgroup, then over NKB warps ---- | |
| + float m_t[8]; | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { m_t[l] = Sacc[l]; } | |
| +#pragma unroll | |
| + for (int off = 1; off < 16; off <<= 1) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { m_t[l] = fmaxf(m_t[l], __shfl_xor_sync(0xFFFFFFFF, m_t[l], off, WS)); } | |
| + } | |
| + if (l16 == 0 && k_half == 0) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Rm[s_tile*16 + 2*l + lhi] = m_t[l]; } | |
| + } | |
| + __syncthreads(); | |
| + | |
| + float m_new[8], rescale[8]; | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { | |
| + float mx = Rm[(qb*NKB + 0)*16 + 2*l + lhi]; | |
| +#pragma unroll | |
| + for (int w = 1; w < NKB; ++w) { mx = fmaxf(mx, Rm[(qb*NKB + w)*16 + 2*l + lhi]); } | |
| + m_new[l] = fmaxf(m_run[l], mx); | |
| + rescale[l] = expf(m_run[l] - m_new[l]); | |
| + m_run[l] = m_new[l]; | |
| + } | |
| + | |
| + // ---- P = exp(S - m), row sums, P into LDS ---- | |
| + float l_t[8]; | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { | |
| + const float p = expf(Sacc[l] - m_new[l]); | |
| + l_t[l] = p; | |
| + if (k_half == 0) { | |
| + Ps[(qb*16 + 2*l + lhi)*(FAB + FPAD) + s_kb*16 + l16] = (_Float16) p; | |
| + } | |
| + } | |
| +#pragma unroll | |
| + for (int off = 1; off < 16; off <<= 1) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { l_t[l] += __shfl_xor_sync(0xFFFFFFFF, l_t[l], off, WS); } | |
| + } | |
| + if (l16 == 0 && k_half == 0) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Rl[s_tile*16 + 2*l + lhi] = l_t[l]; } | |
| + } | |
| + | |
| +#pragma unroll | |
| + for (int t = 0; t < OT; ++t) { | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { Oacc[t][l] *= rescale[l]; } | |
| + } | |
| + __syncthreads(); // Rl ready and K no longer needed | |
| + | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { | |
| + float sum = 0.0f; | |
| +#pragma unroll | |
| + for (int w = 0; w < NKB; ++w) { sum += Rl[(qb*NKB + w)*16 + 2*l + lhi]; } | |
| + l_run[l] = l_run[l]*rescale[l] + sum; | |
| + } | |
| + | |
| + // ---- V transposed into the K buffer: Vt[dv][kv] ---- | |
| + // one lane owns 8 consecutive kv of a single dv, so the LDS write is one int4 | |
| + for (int w = tid; w < (FAB/8)*D; w += NTHR) { | |
| + const int dv = w % D; | |
| + const int kc = w / D; | |
| + fa_frag16 g; | |
| +#pragma unroll | |
| + for (int j = 0; j < 8; ++j) { | |
| + const int kv = kc*8 + j; | |
| + ((_Float16 *) &g.v[0])[j] = kv < kv_sup | |
| + ? fa_load1<type_V>(V_base + (size_t)(k_VKQ_0 + kv)*nb21, dv) : (_Float16) 0.0f; | |
| + } | |
| + *(int4 *) &Vt[(size_t) dv*(FAB + FPAD) + kc*8] = g.v[0]; | |
| + } | |
| + __syncthreads(); | |
| + | |
| + // ---- O += P * V ---- | |
| +#pragma unroll | |
| + for (int t = 0; t < OT; ++t) { | |
| + const int d0 = o_d0 + t*16; | |
| +#pragma unroll | |
| + for (int k0 = 0; k0 < FAB; k0 += 16) { | |
| + const fa_half16 a = fa_load_frag(&Ps[(qb*16 + l16)*(FAB + FPAD) + k0]); | |
| + const fa_half16 b = fa_load_frag(&Vt[(size_t)(d0 + l16)*(FAB + FPAD) + k0]); | |
| + Oacc[t] = __builtin_amdgcn_wmma_f32_16x16x16_f16_w32(a, b, Oacc[t]); | |
| + } | |
| + } | |
| + __syncthreads(); | |
| + } | |
| + | |
| + // ---- write back ---- | |
| +#pragma unroll | |
| + for (int l = 0; l < 8; ++l) { | |
| + const int jc = qb*16 + 2*l + lhi; | |
| + const int j = jc / ncols2; | |
| + const int c = jc % ncols2; | |
| + | |
| + if (ncols1 > 1 && col_Q_0 + j >= int(ne01.z)) { | |
| + continue; | |
| + } | |
| + | |
| + const float out_scale = gridDim.y == 1 ? 1.0f/l_run[l] : 1.0f; | |
| + const int j_dst_unrolled = ((sequence*int(ne01.z) + col_Q_0 + j)*ne02 + head0 + c)*gridDim.y + blockIdx.y; | |
| + | |
| +#pragma unroll | |
| + for (int t = 0; t < OT; ++t) { | |
| + dst_ptr[(size_t) j_dst_unrolled*DV + o_d0 + t*16 + l16] = Oacc[t][l] * out_scale; | |
| + } | |
| + | |
| + if (gridDim.y != 1 && l16 == 0 && o_d0 == 0) { | |
| + dst_meta_ptr[j_dst_unrolled] = make_float2(m_run[l], l_run[l]); | |
| + } | |
| + } | |
| +} | |
| + | |
| +#else // AMD_WMMA_AVAILABLE && RDNA3 | |
| + | |
| +template<int DKQ, int DV, int ncols1, int ncols2, ggml_type type_K, ggml_type type_V> | |
| +__launch_bounds__(FA_WMMA_D256_NWARPS*32, 1) | |
| +static __global__ void flash_attn_wmma_d256( | |
| + const char *, const char *, const char *, const char *, const char *, | |
| + const int *, float *, float2 *, | |
| + const float, const float, const float, const float, const uint32_t, const float, | |
| + const int32_t, const uint3, const int32_t, const int32_t, | |
| + const int32_t, const int32_t, const int32_t, | |
| + const int32_t, const int32_t, const int32_t, const int32_t, | |
| + const int32_t, const int32_t, const int64_t, | |
| + const int32_t, const int32_t, const int64_t, | |
| + const int32_t, const int32_t, const int32_t, | |
| + const int32_t, const int32_t, const int64_t) { | |
| + NO_DEVICE_CODE; | |
| +} | |
| + | |
| +#endif // AMD_WMMA_AVAILABLE && RDNA3 | |
| + | |
| +void ggml_cuda_flash_attn_ext_wmma_d256(ggml_backend_cuda_context & ctx, ggml_tensor * dst); | |
| diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu | |
| index ab7a3b297..f95c57b01 100644 | |
| --- a/ggml/src/ggml-cuda/fattn.cu | |
| +++ b/ggml/src/ggml-cuda/fattn.cu | |
| @@ -3,7 +3,9 @@ | |
| #include "fattn-mma-f16.cuh" | |
| #include "fattn-tile.cuh" | |
| #include "fattn-vec.cuh" | |
| +#include "fattn-wmma-d256.cuh" | |
| #include "fattn.cuh" | |
| +#include <cstdlib> | |
| template <int DKQ, int DV, int ncols2> | |
| static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { | |
| @@ -333,6 +335,7 @@ enum best_fattn_kernel { | |
| BEST_FATTN_KERNEL_TILE = 200, | |
| BEST_FATTN_KERNEL_VEC = 100, | |
| BEST_FATTN_KERNEL_MMA_F16 = 400, | |
| + BEST_FATTN_KERNEL_WMMA_D256 = 500, | |
| }; | |
| static bool ggml_cuda_fattn_kv_type_supported(ggml_type type) { | |
| @@ -511,6 +514,28 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const | |
| } | |
| } | |
| + // gfx1151, head size 256: the tile kernel emits no WMMA at all there, so use the | |
| + // dedicated DV-split WMMA kernel. It needs 64 columns per block to pay for its LDS | |
| + // footprint, so only for batches that can fill them, and only for the plain f16 case. | |
| + float fa_logit_softcap = 0.0f; | |
| + memcpy(&fa_logit_softcap, (const float *) KQV->op_params + 2, sizeof(float)); | |
| + // Opt-in: 1.3-1.4x the tile kernel on short/medium KV in isolation, but its 60 kB LDS | |
| + // footprint fits only 2 workgroups per WGP against the tile kernel's 3, and at long | |
| + // context that costs more than the matrix units gain. See docs/exp205. | |
| + // | |
| + // Quantised K/V are read natively (no pre-conversion to f16) so that the DRAM traffic | |
| + // actually drops. Measured: it does not pay. q8_0 and q4_0 are 1.7-1.8x *slower* than | |
| + // f16 here, because this kernel is bound by instruction issue and memory latency rather | |
| + // than by bytes, and per-element dequantisation costs more than the bytes it saves. | |
| + if (GGML_CUDA_CC_IS_RDNA3_5(cc) && getenv("GGML_CUDA_FA_WMMA256") | |
| + && Q->ne[0] == 256 && V->ne[0] == 256 | |
| + && K->type == V->type | |
| + && (K->type == GGML_TYPE_F16 || K->type == GGML_TYPE_Q8_0 || K->type == GGML_TYPE_Q4_0) | |
| + && mask && !dst->src[4] && max_bias == 0.0f && fa_logit_softcap == 0.0f | |
| + && gqa_opt_applies && Q->ne[1] * gqa_ratio_eff >= FA_WMMA_D256_NC) { | |
| + return BEST_FATTN_KERNEL_WMMA_D256; | |
| + } | |
| + | |
| // AMD WMMA is always faster than the tile kernel if the full tile width of 16 can be utilized. | |
| if ((amd_wmma_available(cc) && gqa_opt_applies && Q->ne[0] <= 128) && Q->ne[0] != 40 && Q->ne[0] != 72 && Q->ne[1] * gqa_ratio_eff > 8) { | |
| return BEST_FATTN_KERNEL_MMA_F16; | |
| @@ -553,6 +578,8 @@ size_t ggml_cuda_flash_attn_ext_get_alloc_size(int device, const ggml_tensor * d | |
| need_f16_K = true; | |
| need_f16_V = true; | |
| break; | |
| + case BEST_FATTN_KERNEL_WMMA_D256: | |
| + break; // reads K/V natively, quantised or not | |
| case BEST_FATTN_KERNEL_VEC: | |
| need_f16_K = K->type == GGML_TYPE_F32; | |
| need_f16_V = V->type == GGML_TYPE_F32; | |
| @@ -567,8 +594,35 @@ size_t ggml_cuda_flash_attn_ext_get_alloc_size(int device, const ggml_tensor * d | |
| return f16_extra.end - (uintptr_t) dst->data; | |
| } | |
| -void ggml_cuda_flash_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { | |
| - ggml_cuda_set_device(ctx.device); | |
| +// Max Q columns per flash-attention launch, 0 to disable splitting. | |
| +// One launch with many Q columns puts many column tiles on the GPU at the same time. The tiles | |
| +// then read different parts of K/V and stop sharing them in cache, which only hurts when the KV | |
| +// cache is much larger than the last level cache. Splitting Q keeps the number of concurrent | |
| +// tiles low. On gfx1151 with ubatch 2048 this is 1.4x faster at 32k context and 1.8x at 64k. | |
| +static int ggml_cuda_fattn_q_chunk(const ggml_tensor * dst) { | |
| + static const int chunk = getenv("GGML_CUDA_FA_QCHUNK") ? atoi(getenv("GGML_CUDA_FA_QCHUNK")) : 256; | |
| + | |
| + if (chunk <= 0 || !GGML_CUDA_CC_IS_RDNA3_5(ggml_cuda_info().devices[ggml_cuda_get_device()].cc)) { | |
| + return 0; | |
| + } | |
| + | |
| + // Only worth it if K/V does not fit in cache anyway, else the extra launches just add overhead. | |
| + static const int min_kv = getenv("GGML_CUDA_FA_QCHUNK_MINKV") ? atoi(getenv("GGML_CUDA_FA_QCHUNK_MINKV")) : 8192; | |
| + | |
| + if (dst->src[1]->ne[1] < min_kv) { | |
| + return 0; | |
| + } | |
| + | |
| + // Q rows are dst dimension 2. If dst has more than one dimension 3 then a row range of dst is | |
| + // not one continuous block, but the kernels write dst as continuous memory. Do not split. | |
| + if (dst->ne[3] != 1) { | |
| + return 0; | |
| + } | |
| + | |
| + return chunk; | |
| +} | |
| + | |
| +static void ggml_cuda_flash_attn_ext_dispatch(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { | |
| switch (ggml_cuda_get_best_fattn_kernel(ggml_cuda_get_device(), dst)) { | |
| case BEST_FATTN_KERNEL_NONE: | |
| GGML_ABORT("fatal error"); | |
| @@ -581,6 +635,50 @@ void ggml_cuda_flash_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst | |
| case BEST_FATTN_KERNEL_MMA_F16: | |
| ggml_cuda_flash_attn_ext_mma_f16(ctx, dst); | |
| break; | |
| + case BEST_FATTN_KERNEL_WMMA_D256: | |
| + ggml_cuda_flash_attn_ext_wmma_d256(ctx, dst); | |
| + break; | |
| + } | |
| +} | |
| + | |
| +void ggml_cuda_flash_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { | |
| + ggml_cuda_set_device(ctx.device); | |
| + | |
| + const ggml_tensor * Q = dst->src[0]; | |
| + const ggml_tensor * mask = dst->src[3]; | |
| + | |
| + const int chunk = ggml_cuda_fattn_q_chunk(dst); | |
| + if (chunk == 0 || Q->ne[1] <= chunk) { | |
| + ggml_cuda_flash_attn_ext_dispatch(ctx, dst); | |
| + return; | |
| + } | |
| + | |
| + // Q rows are dst rows in dimension 2, mask rows are in dimension 1. | |
| + ggml_tensor Q_chunk = *Q; | |
| + ggml_tensor mask_chunk = mask ? *mask : ggml_tensor(); | |
| + ggml_tensor dst_chunk = *dst; | |
| + | |
| + dst_chunk.src[0] = &Q_chunk; | |
| + if (mask) { | |
| + dst_chunk.src[3] = &mask_chunk; | |
| + } | |
| + | |
| + for (int64_t i0 = 0; i0 < Q->ne[1]; i0 += chunk) { | |
| + const int64_t ncols = std::min<int64_t>(chunk, Q->ne[1] - i0); | |
| + | |
| + Q_chunk.ne[1] = ncols; | |
| + Q_chunk.data = (char *) Q->data + i0*Q->nb[1]; | |
| + | |
| + // Keep all remaining mask rows so that reads past ncols stay in bounds. | |
| + if (mask) { | |
| + mask_chunk.ne[1] = mask->ne[1] - i0; | |
| + mask_chunk.data = (char *) mask->data + i0*mask->nb[1]; | |
| + } | |
| + | |
| + dst_chunk.ne[2] = ncols; | |
| + dst_chunk.data = (char *) dst->data + i0*dst->nb[2]; | |
| + | |
| + ggml_cuda_flash_attn_ext_dispatch(ctx, &dst_chunk); | |
| } | |
| } | |
| diff --git a/ggml/src/ggml-cuda/gated_delta_net.cu b/ggml/src/ggml-cuda/gated_delta_net.cu | |
| index 1b431a724..97ab34724 100644 | |
| --- a/ggml/src/ggml-cuda/gated_delta_net.cu | |
| +++ b/ggml/src/ggml-cuda/gated_delta_net.cu | |
| @@ -1,6 +1,17 @@ | |
| #include "gated_delta_net.cuh" | |
| +#include <cstdlib> | |
| #include "ggml-cuda/common.cuh" | |
| +// Lanes cooperating on one state column in the column-parallel kernel. | |
| +// 0 disables it and falls back to the warp-per-column kernel. Overridable for tuning. | |
| +static int ggml_cuda_gdn_split() { | |
| + static const int split = [] { | |
| + const char * s = getenv("GGML_CUDA_GDN_SPLIT"); | |
| + return s ? atoi(s) : 2; | |
| + }(); | |
| + return split; | |
| +} | |
| + | |
| template <int S_v, bool KDA, bool keep_rs_t> | |
| __global__ void __launch_bounds__((ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v) * 4, 2) | |
| gated_delta_net_cuda(const float * q, | |
| @@ -166,6 +177,152 @@ gated_delta_net_cuda(const float * q, | |
| } | |
| } | |
| +// Column-parallel variant of the kernel above. | |
| +// | |
| +// The warp-per-column kernel gives one *warp* one state column and splits the column's | |
| +// rows across the warp's lanes. Every warp therefore reads the head's whole q and k | |
| +// vector, and with the columns spread over S_v/nwarps blocks that means q and k get | |
| +// re-read S_v times per head per token. That redundancy, not the arithmetic, is what | |
| +// bounds the kernel on prefill. | |
| +// | |
| +// Here a block owns *all* S_v columns of one head, so q/k/g are read exactly once per | |
| +// head per token into LDS. Each column is handled by SPLIT adjacent lanes, so a thread | |
| +// holds S_v/SPLIT state values in registers and the row reduction costs log2(SPLIT) | |
| +// shuffles instead of log2(warp_size). SPLIT trades registers (and hence occupancy) | |
| +// against those shuffles. | |
| +template <int S_v, int SPLIT, bool KDA, bool keep_rs_t> | |
| +__global__ void __launch_bounds__(S_v * SPLIT, 1) | |
| +gated_delta_net_col_cuda(const float * __restrict__ q, | |
| + const float * __restrict__ k, | |
| + const float * __restrict__ v, | |
| + const float * __restrict__ g, | |
| + const float * __restrict__ beta, | |
| + const float * __restrict__ curr_state, | |
| + float * __restrict__ dst, | |
| + float * __restrict__ state, | |
| + int64_t H, | |
| + int64_t n_tokens, | |
| + int64_t n_seqs, | |
| + int64_t sq1, int64_t sq2, int64_t sq3, | |
| + int64_t sv1, int64_t sv2, int64_t sv3, | |
| + int64_t sb1, int64_t sb2, int64_t sb3, | |
| + const uint3 neqk1_magic, | |
| + const uint3 rq3_magic, | |
| + float scale, | |
| + int64_t state_slot_stride, | |
| + int K) { | |
| + GGML_UNUSED(n_seqs); | |
| + constexpr int RPT = S_v / SPLIT; // state rows held per thread | |
| + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); | |
| + static_assert(S_v % SPLIT == 0, "SPLIT must divide S_v"); | |
| + static_assert(SPLIT <= warp_size, "the row reduction must stay inside one wave"); | |
| + | |
| + const uint32_t h_idx = blockIdx.x; | |
| + const uint32_t sequence = blockIdx.y; | |
| + const int tid = threadIdx.x; | |
| + const int col = tid / SPLIT; // SPLIT adjacent lanes share a column | |
| + const int part = tid % SPLIT; | |
| + const int row0 = part * RPT; | |
| + | |
| + const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic); | |
| + const uint32_t iq3 = fastdiv(sequence, rq3_magic); | |
| + | |
| + float * attn_data = dst + (sequence * n_tokens * H + h_idx) * S_v; | |
| + | |
| + const int64_t state_in_offset = sequence * H * S_v * S_v + h_idx * S_v * S_v; | |
| + const int64_t state_out_offset = (sequence * H + h_idx) * S_v * S_v; | |
| + state += state_out_offset; | |
| + curr_state += state_in_offset + col * S_v + row0; | |
| + | |
| + // state is stored transposed: M[col][i] = S[i][col], so row `col` is contiguous | |
| + float s_col[RPT]; | |
| + | |
| + // q/k (and the per-row gate for KDA) are identical for every column, so stage them in | |
| + // LDS once per token. Double-buffered so one barrier per token is enough. | |
| + __shared__ float sk[2][S_v]; | |
| + __shared__ float sq[2][S_v]; | |
| + __shared__ float sg[2][KDA ? S_v : 1]; | |
| + | |
| + ggml_cuda_pdl_sync(); | |
| +#pragma unroll | |
| + for (int j = 0; j < RPT; j++) { | |
| + s_col[j] = curr_state[j]; | |
| + } | |
| + | |
| + for (int t = 0; t < n_tokens; t++) { | |
| + const float * __restrict__ q_t = q + iq3 * sq3 + t * sq2 + iq1 * sq1; | |
| + const float * __restrict__ k_t = k + iq3 * sq3 + t * sq2 + iq1 * sq1; | |
| + const float * __restrict__ v_t = v + sequence * sv3 + t * sv2 + h_idx * sv1; | |
| + | |
| + const int64_t gb_offset = sequence * sb3 + t * sb2 + h_idx * sb1; | |
| + const float beta_val = beta[gb_offset]; | |
| + const float * __restrict__ g_t = g + gb_offset * (KDA ? S_v : 1); | |
| + | |
| + const int buf = t & 1; | |
| + if (tid < S_v) { | |
| + sk[buf][tid] = k_t[tid]; | |
| + sq[buf][tid] = q_t[tid]; | |
| + if (KDA) { | |
| + sg[buf][tid] = expf(g_t[tid]); | |
| + } | |
| + } | |
| + __syncthreads(); | |
| + | |
| + const float * __restrict__ lk = sk[buf] + row0; | |
| + const float * __restrict__ lq = sq[buf] + row0; | |
| + const float * __restrict__ lg = sg[buf] + (KDA ? row0 : 0); | |
| + | |
| + const float g_val = KDA ? 0.0f : expf(*g_t); | |
| + | |
| + float kv = 0.0f; | |
| +#pragma unroll | |
| + for (int j = 0; j < RPT; j++) { | |
| + kv = KDA ? fmaf(lg[j] * s_col[j], lk[j], kv) : fmaf(s_col[j], lk[j], kv); | |
| + } | |
| +#pragma unroll | |
| + for (int o = 1; o < SPLIT; o <<= 1) { | |
| + kv += __shfl_xor_sync(0xFFFFFFFF, kv, o, warp_size); | |
| + } | |
| + | |
| + const float delta = (v_t[col] - (KDA ? kv : g_val * kv)) * beta_val; | |
| + | |
| + float attn = 0.0f; | |
| +#pragma unroll | |
| + for (int j = 0; j < RPT; j++) { | |
| + s_col[j] = fmaf(KDA ? lg[j] : g_val, s_col[j], lk[j] * delta); | |
| + attn = fmaf(s_col[j], lq[j], attn); | |
| + } | |
| +#pragma unroll | |
| + for (int o = 1; o < SPLIT; o <<= 1) { | |
| + attn += __shfl_xor_sync(0xFFFFFFFF, attn, o, warp_size); | |
| + } | |
| + | |
| + if (part == 0) { | |
| + attn_data[col] = attn * scale; | |
| + } | |
| + attn_data += S_v * H; | |
| + | |
| + if constexpr (keep_rs_t) { | |
| + const int target_slot = (int) n_tokens - 1 - t; | |
| + if (target_slot >= 0 && target_slot < K) { | |
| + float * slot = state + target_slot * state_slot_stride + col * S_v + row0; | |
| +#pragma unroll | |
| + for (int j = 0; j < RPT; j++) { | |
| + slot[j] = s_col[j]; | |
| + } | |
| + } | |
| + } | |
| + } | |
| + | |
| + if constexpr (!keep_rs_t) { | |
| + float * out = state + col * S_v + row0; | |
| +#pragma unroll | |
| + for (int j = 0; j < RPT; j++) { | |
| + out[j] = s_col[j]; | |
| + } | |
| + } | |
| +} | |
| + | |
| template <bool KDA, bool keep_rs_t> | |
| static void launch_gated_delta_net( | |
| const float * q_d, const float * k_d, const float * v_d, | |
| @@ -186,6 +343,42 @@ static void launch_gated_delta_net( | |
| const uint3 neqk1_magic = init_fastdiv_values(neqk1); | |
| const uint3 rq3_magic = init_fastdiv_values(rq3); | |
| + // Column-parallel kernel: one block per (head, sequence), all S_v columns inside it. | |
| + // Needs at least a full wave per block, so only used from S_v = 64 up. | |
| + { | |
| + const int gdn_split = ggml_cuda_gdn_split(); | |
| + // For a single token the warp-per-column kernel wins: there is no q/k reuse to | |
| + // recover and its narrower blocks start faster. | |
| + if (gdn_split > 0 && n_tokens >= 16 && (S_v == 64 || S_v == 128)) { | |
| + const dim3 grid_col(H, n_seqs, 1); | |
| + const dim3 block_col(S_v * gdn_split, 1, 1); | |
| + const ggml_cuda_kernel_launch_params lp = ggml_cuda_kernel_launch_params(grid_col, block_col, 0, stream); | |
| +#define GDN_COL_LAUNCH(SV, SP) \ | |
| + ggml_cuda_kernel_launch(gated_delta_net_col_cuda<SV, SP, KDA, keep_rs_t>, lp, \ | |
| + q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H, \ | |
| + n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, \ | |
| + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K) | |
| + if (S_v == 64) { | |
| + switch (gdn_split) { | |
| + case 1: GDN_COL_LAUNCH(64, 1); return; | |
| + case 2: GDN_COL_LAUNCH(64, 2); return; | |
| + case 4: GDN_COL_LAUNCH(64, 4); return; | |
| + case 8: GDN_COL_LAUNCH(64, 8); return; | |
| + default: break; | |
| + } | |
| + } else { | |
| + switch (gdn_split) { | |
| + case 1: GDN_COL_LAUNCH(128, 1); return; | |
| + case 2: GDN_COL_LAUNCH(128, 2); return; | |
| + case 4: GDN_COL_LAUNCH(128, 4); return; | |
| + case 8: GDN_COL_LAUNCH(128, 8); return; | |
| + default: break; | |
| + } | |
| + } | |
| +#undef GDN_COL_LAUNCH | |
| + } | |
| + } | |
| + | |
| const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, stream); | |
| switch (S_v) { | |
| case 16: | |
| diff --git a/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh b/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh | |
| index 180b2d937..9838fd175 100644 | |
| --- a/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh | |
| +++ b/ggml/src/ggml-cuda/mmq-config-rdna3-5.cuh | |
| @@ -1,290 +1,290 @@ | |
| 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 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_Q1_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q1_0, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 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_Q2_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_0, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 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_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_0, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 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_Q4_1, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_1, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 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_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_0, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 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_Q5_1, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_1, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 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_Q8_0, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q8_0, 128, 2, 64, 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, 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_Q2_K, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q2_K, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 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_Q3_K, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q3_K, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 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_Q4_K, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q4_K, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 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_Q5_K, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q5_K, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 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_Q6_K, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_Q6_K, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 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_IQ1_S, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ1_S, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 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_XXS, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XXS, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 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_XS, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_XS, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 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_IQ2_S, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ2_S, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 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_XXS, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_XXS, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 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_IQ3_S, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ3_S, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 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_XS, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_XS, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 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_IQ4_NL, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_IQ4_NL, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 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_MXFP4, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_MXFP4, 128, 2, 64, 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, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 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); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 80, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 96, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 112, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); | |
| + CASE(GGML_TYPE_NVFP4, 128, 2, 64, 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); | |
| + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 128, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); | |
| } | |
| diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu | |
| index 0589e65bd..3a9d732ca 100644 | |
| --- a/ggml/src/ggml-cuda/mmvq.cu | |
| +++ b/ggml/src/ggml-cuda/mmvq.cu | |
| @@ -63,12 +63,17 @@ static constexpr __host__ __device__ int get_vdr_mmvq(ggml_type type) { | |
| } | |
| } | |
| +#ifndef GGML_MMVQ_RDNA35_NWARPS | |
| +#define GGML_MMVQ_RDNA35_NWARPS 2 | |
| +#endif | |
| + | |
| enum mmvq_parameter_table_id { | |
| MMVQ_PARAMETERS_GENERIC = 0, | |
| MMVQ_PARAMETERS_TURING, | |
| MMVQ_PARAMETERS_GCN, | |
| MMVQ_PARAMETERS_RDNA2, | |
| MMVQ_PARAMETERS_RDNA3_0, | |
| + MMVQ_PARAMETERS_RDNA3_5, | |
| MMVQ_PARAMETERS_RDNA4 | |
| }; | |
| @@ -77,7 +82,9 @@ static constexpr __device__ mmvq_parameter_table_id get_device_table_id() { | |
| return MMVQ_PARAMETERS_RDNA4; | |
| #elif defined(RDNA3_0) | |
| return MMVQ_PARAMETERS_RDNA3_0; | |
| -#elif defined(RDNA2) || defined(RDNA3_5) | |
| +#elif defined(RDNA3_5) | |
| + return MMVQ_PARAMETERS_RDNA3_5; | |
| +#elif defined(RDNA2) | |
| return MMVQ_PARAMETERS_RDNA2; | |
| #elif defined(GCN) || defined(CDNA) | |
| return MMVQ_PARAMETERS_GCN; | |
| @@ -95,7 +102,10 @@ static __host__ mmvq_parameter_table_id get_device_table_id(int cc) { | |
| if (GGML_CUDA_CC_IS_RDNA3_0(cc)) { | |
| return MMVQ_PARAMETERS_RDNA3_0; | |
| } | |
| - if (GGML_CUDA_CC_IS_RDNA2(cc) || GGML_CUDA_CC_IS_RDNA3_5(cc)) { | |
| + if (GGML_CUDA_CC_IS_RDNA3_5(cc)) { | |
| + return MMVQ_PARAMETERS_RDNA3_5; | |
| + } | |
| + if (GGML_CUDA_CC_IS_RDNA2(cc)) { | |
| return MMVQ_PARAMETERS_RDNA2; | |
| } | |
| if (GGML_CUDA_CC_IS_GCN(cc) || GGML_CUDA_CC_IS_CDNA(cc)) { | |
| @@ -406,6 +416,31 @@ static constexpr __host__ __device__ int calc_nwarps(ggml_type type, int ncols_d | |
| } | |
| return 1; | |
| } | |
| + if (table_id == MMVQ_PARAMETERS_RDNA3_5) { | |
| + // Strix Halo: mat-vec decode is bandwidth bound, and one warp per block does not | |
| + // issue enough loads to keep the memory system busy. GGML_MMVQ_RDNA35_NWARPS | |
| + // warps cooperate on a row instead. | |
| + if (ncols_dst == 1) { | |
| + switch (type) { | |
| + case GGML_TYPE_Q4_0: | |
| + case GGML_TYPE_Q4_1: | |
| + case GGML_TYPE_Q5_0: | |
| + case GGML_TYPE_Q5_1: | |
| + case GGML_TYPE_Q8_0: | |
| + case GGML_TYPE_Q2_K: | |
| + case GGML_TYPE_Q3_K: | |
| + case GGML_TYPE_Q4_K: | |
| + case GGML_TYPE_Q5_K: | |
| + case GGML_TYPE_Q6_K: | |
| + case GGML_TYPE_IQ4_NL: | |
| + case GGML_TYPE_IQ4_XS: | |
| + return GGML_MMVQ_RDNA35_NWARPS; | |
| + default: | |
| + return 1; | |
| + } | |
| + } | |
| + return 1; | |
| + } | |
| if (table_id == MMVQ_PARAMETERS_RDNA3_0) { | |
| // RDNA3 (W7900): stricter whitelist than RDNA4. | |
| // Q2_K / Q5_K / IQ4_XS regress in full quant sweeps. | |
| diff --git a/ggml/src/ggml-cuda/quantize.cu b/ggml/src/ggml-cuda/quantize.cu | |
| index bcc772395..dd42db0a0 100644 | |
| --- a/ggml/src/ggml-cuda/quantize.cu | |
| +++ b/ggml/src/ggml-cuda/quantize.cu | |
| @@ -1,4 +1,5 @@ | |
| #include "quantize.cuh" | |
| +#include <cstdlib> | |
| #include <cstdint> | |
| #if defined(BLACKWELL_MMA_AVAILABLE) | |
| @@ -458,19 +459,25 @@ template <mmq_q8_1_ds_layout ds_layout, bool scatter> | |
| static __global__ void quantize_mmq_q8_1( | |
| const float * __restrict__ x, const int32_t * __restrict__ ids, void * __restrict__ vy, | |
| const int64_t ne00, const int64_t s01, const int64_t s02, const int64_t s03, | |
| - const int64_t ne0, const int ne1, const int ne2, const int n_expert_used) { | |
| + const int64_t ne0, const int ne1, const int ne2, const int n_expert_used, const int chunks) { | |
| constexpr int vals_per_scale = ds_layout == MMQ_Q8_1_DS_LAYOUT_D2S6 ? 64 : 32; | |
| constexpr int vals_per_sum = ds_layout == MMQ_Q8_1_DS_LAYOUT_D2S6 ? 16 : 32; | |
| - const int64_t i0 = ((int64_t)blockDim.x*blockIdx.y + threadIdx.x)*4; | |
| + ggml_cuda_pdl_sync(); | |
| + | |
| + // Each block walks `chunks` consecutive column chunks of one row. One chunk is only | |
| + // 4*blockDim.x floats, so a chunk per block leaves the reads scattered over ne1 | |
| + // separate rows; walking several keeps a longer contiguous run per row. | |
| + for (int chunk = 0; chunk < chunks; ++chunk) { | |
| + | |
| + const int64_t i0 = ((int64_t)blockDim.x*(blockIdx.y*(int64_t)chunks + chunk) + threadIdx.x)*4; | |
| if (i0 >= ne0) { | |
| - return; | |
| + break; | |
| } | |
| const int64_t i00 = i0; | |
| - ggml_cuda_pdl_sync(); | |
| int64_t base_idx; | |
| if constexpr (scatter) { | |
| @@ -529,7 +536,7 @@ static __global__ void quantize_mmq_q8_1( | |
| const int64_t i = ids[(int64_t) blockIdx.x * n_expert_used + slot]; | |
| ib = k_block*ne1 + i; | |
| } else { | |
| - const int64_t ib0 = blockIdx.z*((int64_t)gridDim.x*gridDim.y*blockDim.x/QK8_1); // first block of channel | |
| + const int64_t ib0 = blockIdx.z*(ne0/QK8_1_MMQ)*(int64_t)ne1; // first block of channel | |
| ib = ib0 + k_block*ne1 + blockIdx.x; | |
| } | |
| @@ -552,6 +559,8 @@ static __global__ void quantize_mmq_q8_1( | |
| } | |
| } | |
| } | |
| + | |
| + } // chunk | |
| GGML_UNUSED(n_expert_used); | |
| } | |
| @@ -572,6 +581,14 @@ void quantize_row_q8_1_cuda( | |
| GGML_UNUSED(type_src0); | |
| } | |
| +static int ggml_cuda_quantize_mmq_chunks() { | |
| + static const int chunks = [] { | |
| + const char * s = getenv("GGML_CUDA_QUANTIZE_CHUNKS"); | |
| + return s ? atoi(s) : 2; | |
| + }(); | |
| + return chunks; | |
| +} | |
| + | |
| void quantize_mmq_q8_1_cuda( | |
| const float * x, const int32_t * ids, void * vy, const ggml_type type_src0, | |
| const int64_t ne00, const int64_t s01, const int64_t s02, const int64_t s03, | |
| @@ -580,21 +597,23 @@ void quantize_mmq_q8_1_cuda( | |
| GGML_ASSERT(ne0 % QK8_1_MMQ == 0); | |
| // ne1 tends to assume the highest values, therefore use it as the "x" dimension of the CUDA grid: | |
| - const int64_t block_num_y = (ne0 + 4*CUDA_QUANTIZE_BLOCK_SIZE_MMQ - 1) / (4*CUDA_QUANTIZE_BLOCK_SIZE_MMQ); | |
| + const int64_t block_num_y_full = (ne0 + 4*CUDA_QUANTIZE_BLOCK_SIZE_MMQ - 1) / (4*CUDA_QUANTIZE_BLOCK_SIZE_MMQ); | |
| + const int chunks = ggml_cuda_quantize_mmq_chunks(); | |
| + const int64_t block_num_y = (block_num_y_full + chunks - 1) / chunks; | |
| const dim3 num_blocks(ne1, block_num_y, ne2*ne3); | |
| const dim3 block_size(CUDA_QUANTIZE_BLOCK_SIZE_MMQ, 1, 1); | |
| switch (mmq_get_q8_1_ds_layout(type_src0)) { | |
| case MMQ_Q8_1_DS_LAYOUT_D4: | |
| quantize_mmq_q8_1<MMQ_Q8_1_DS_LAYOUT_D4, false> | |
| - <<<num_blocks, block_size, 0, stream>>>(x, ids, vy, ne00, s01, s02, s03, ne0, ne1, ne2, /*n_expert_used=*/0); | |
| + <<<num_blocks, block_size, 0, stream>>>(x, ids, vy, ne00, s01, s02, s03, ne0, ne1, ne2, /*n_expert_used=*/0, chunks); | |
| break; | |
| case MMQ_Q8_1_DS_LAYOUT_DS4: | |
| quantize_mmq_q8_1<MMQ_Q8_1_DS_LAYOUT_DS4, false> | |
| - <<<num_blocks, block_size, 0, stream>>>(x, ids, vy, ne00, s01, s02, s03, ne0, ne1, ne2, /*n_expert_used=*/0); | |
| + <<<num_blocks, block_size, 0, stream>>>(x, ids, vy, ne00, s01, s02, s03, ne0, ne1, ne2, /*n_expert_used=*/0, chunks); | |
| break; | |
| case MMQ_Q8_1_DS_LAYOUT_D2S6: | |
| quantize_mmq_q8_1<MMQ_Q8_1_DS_LAYOUT_D2S6, false> | |
| - <<<num_blocks, block_size, 0, stream>>>(x, ids, vy, ne00, s01, s02, s03, ne0, ne1, ne2, /*n_expert_used=*/0); | |
| + <<<num_blocks, block_size, 0, stream>>>(x, ids, vy, ne00, s01, s02, s03, ne0, ne1, ne2, /*n_expert_used=*/0, chunks); | |
| break; | |
| default: | |
| GGML_ABORT("fatal error"); | |
| @@ -611,20 +630,21 @@ void quantize_scatter_mmq_q8_1_cuda( | |
| GGML_ASSERT(ne0 % QK8_1_MMQ == 0); | |
| const int64_t block_num_y = (ne0 + 4*CUDA_QUANTIZE_BLOCK_SIZE_MMQ - 1) / (4*CUDA_QUANTIZE_BLOCK_SIZE_MMQ); | |
| + const int chunks = 1; | |
| const dim3 num_blocks(n_tokens, block_num_y, 1); | |
| const dim3 block_size(CUDA_QUANTIZE_BLOCK_SIZE_MMQ, 1, 1); | |
| switch (mmq_get_q8_1_ds_layout(type_src0)) { | |
| case MMQ_Q8_1_DS_LAYOUT_D4: | |
| quantize_mmq_q8_1<MMQ_Q8_1_DS_LAYOUT_D4, true><<<num_blocks, block_size, 0, stream>>>( | |
| - x, ids_src1_inv, vy, ne00, /*s01=*/0, /*s02=*/stride_token, /*s03=*/0, ne0, /*ne1=*/(int) nrows_dst, /*ne2=*/1, n_expert_used); | |
| + x, ids_src1_inv, vy, ne00, /*s01=*/0, /*s02=*/stride_token, /*s03=*/0, ne0, /*ne1=*/(int) nrows_dst, /*ne2=*/1, n_expert_used, chunks); | |
| break; | |
| case MMQ_Q8_1_DS_LAYOUT_DS4: | |
| quantize_mmq_q8_1<MMQ_Q8_1_DS_LAYOUT_DS4, true><<<num_blocks, block_size, 0, stream>>>( | |
| - x, ids_src1_inv, vy, ne00, /*s01=*/0, /*s02=*/stride_token, /*s03=*/0, ne0, /*ne1=*/(int) nrows_dst, /*ne2=*/1, n_expert_used); | |
| + x, ids_src1_inv, vy, ne00, /*s01=*/0, /*s02=*/stride_token, /*s03=*/0, ne0, /*ne1=*/(int) nrows_dst, /*ne2=*/1, n_expert_used, chunks); | |
| break; | |
| case MMQ_Q8_1_DS_LAYOUT_D2S6: | |
| quantize_mmq_q8_1<MMQ_Q8_1_DS_LAYOUT_D2S6, true><<<num_blocks, block_size, 0, stream>>>( | |
| - x, ids_src1_inv, vy, ne00, /*s01=*/0, /*s02=*/stride_token, /*s03=*/0, ne0, /*ne1=*/(int) nrows_dst, /*ne2=*/1, n_expert_used); | |
| + x, ids_src1_inv, vy, ne00, /*s01=*/0, /*s02=*/stride_token, /*s03=*/0, ne0, /*ne1=*/(int) nrows_dst, /*ne2=*/1, n_expert_used, chunks); | |
| break; | |
| default: | |
| GGML_ABORT("fatal error"); |
