diff --git a/src/ggml-cuda/convert.cu b/src/ggml-cuda/convert.cu index 61630a3..03ca5fe 100644 --- a/src/ggml-cuda/convert.cu +++ b/src/ggml-cuda/convert.cu @@ -709,6 +709,262 @@ to_bf16_cuda_t ggml_get_to_bf16_cuda(ggml_type type) { } } +// ---- sasori ternary trit-plane types (TQ{K}P / TQ{K}P_G{g} / TQ{K}P_Q4K6) ---- +// K trit-planes summed. The per-plane 2-bit packing mirrors block_tq2_0's strided layout, so the +// byte index equals the thread id within a plane; only the scale decode differs per family. This +// dequant mirrors dequantize_row_tq{K}p* / q4k6_dequantize in ggml-quants.c. MVP path: dequant -> +// cuBLAS fp16 GEMM (no mmvq/mmq). Activation goes through fp16 here, not Q8_K as on CPU, so expect +// tiny numeric drift vs the CPU vec_dot (same class of caveat as PR ggml-org/llama.cpp#11183). + +// f16-scale K-plane dequant: fixed TQ{K}P == G=256, plus TQ{K}P_G{32,64,128}. nsub = 256/G scales/plane. +template +static __global__ void dequantize_block_tqkp_gv(const void * __restrict__ vx, dst_t * __restrict__ yy) { + const int64_t i = blockIdx.x; // super-block (256 weights) + const int64_t tid = threadIdx.x; // 0..63: byte within a plane + const int nsub = 256 / G; + const int qb = QK_TQKP_SUPER / 4; // 64 bytes per plane + const int bpb = K*qb + K*nsub*2; // bytes per super-block + const uint8_t * __restrict__ blk = (const uint8_t *) vx + i*bpb; + const half * __restrict__ d = (const half *) (blk + K*qb); + const int n = tid / 32; + const int l = tid - 32*n; + dst_t * y = yy + i*QK_TQKP_SUPER + 128*n; +#pragma unroll + for (int bp = 0; bp < 4; ++bp) { + const int j = 128*n + bp*32 + l; // logical position in the super-block + const int sub = j / G; + float acc = 0.0f; +#pragma unroll + for (int kk = 0; kk < K; ++kk) { + const int t = ((blk[kk*qb + tid] >> (bp*2)) & 3) - 1; // trit in {-1,0,1} + acc += __half2float(d[kk*nsub + sub]) * (float) t; + } + y[l + bp*32] = (dst_t) acc; + } +} + +// q4k6-scale K-plane dequant (g=32 fixed): per-plane fp16 super-scale, sign-flipped by the sub bit7, +// times a 6-bit unsigned sub-scale. Mirrors q4k6_dequantize() in ggml-quants.c bit-for-bit. +template +static __global__ void dequantize_block_tqkp_q4k6(const void * __restrict__ vx, dst_t * __restrict__ yy) { + const int64_t i = blockIdx.x; + const int64_t tid = threadIdx.x; // 0..63 + const int nsub = QK_TQKP_SUPER / Q4K6_G; // 8 + const int qb = QK_TQKP_SUPER / 4; // 64 + const int bpb = K*qb + K*nsub + K*2; + const uint8_t * __restrict__ blk = (const uint8_t *) vx + i*bpb; + const uint8_t * __restrict__ sb = blk + K*qb; // K*nsub sub-scale bytes + const half * __restrict__ sup = (const half *)(blk + K*qb + K*nsub); // K fp16 super-scales + const int n = tid / 32; + const int l = tid - 32*n; + dst_t * y = yy + i*QK_TQKP_SUPER + 128*n; +#pragma unroll + for (int bp = 0; bp < 4; ++bp) { + const int j = 128*n + bp*32 + l; + const int c = j / Q4K6_G; // sub-group index 0..7 + float acc = 0.0f; +#pragma unroll + for (int kk = 0; kk < K; ++kk) { + const uint8_t s = sb[kk*nsub + c]; + uint32_t bits = __float_as_uint(__half2float(sup[kk])); + bits ^= ((uint32_t)(s & 0x80u)) << 24; // flip sign iff sub bit7 set + const float sc = __uint_as_float(bits) * (float)(s & 0x3F); + const int t = ((blk[kk*qb + tid] >> (bp*2)) & 3) - 1; + acc += sc * (float) t; + } + y[l + bp*32] = (dst_t) acc; + } +} + +#define DEQUANT_TQKP_GV_CUDA(NAME, K, G) \ + template \ + static void NAME(const void * vx, dst_t * y, const int64_t k, cudaStream_t stream) { \ + const int nb = k / QK_TQKP_SUPER; \ + dequantize_block_tqkp_gv<<>>(vx, y); \ + } +#define DEQUANT_TQKP_Q4K6_CUDA(NAME, K) \ + template \ + static void NAME(const void * vx, dst_t * y, const int64_t k, cudaStream_t stream) { \ + const int nb = k / QK_TQKP_SUPER; \ + dequantize_block_tqkp_q4k6<<>>(vx, y); \ + } + +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq1p_cuda, 1, 256) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq2p_cuda, 2, 256) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq3p_cuda, 3, 256) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq4p_cuda, 4, 256) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq1p_g32_cuda, 1, 32) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq1p_g64_cuda, 1, 64) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq1p_g128_cuda, 1, 128) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq2p_g32_cuda, 2, 32) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq2p_g64_cuda, 2, 64) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq2p_g128_cuda, 2, 128) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq3p_g32_cuda, 3, 32) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq3p_g64_cuda, 3, 64) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq3p_g128_cuda, 3, 128) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq4p_g32_cuda, 4, 32) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq4p_g64_cuda, 4, 64) +DEQUANT_TQKP_GV_CUDA(dequantize_row_tq4p_g128_cuda, 4, 128) +DEQUANT_TQKP_Q4K6_CUDA(dequantize_row_tq1p_q4k6_cuda, 1) +DEQUANT_TQKP_Q4K6_CUDA(dequantize_row_tq2p_q4k6_cuda, 2) +DEQUANT_TQKP_Q4K6_CUDA(dequantize_row_tq3p_q4k6_cuda, 3) + + +// ---- sasori fused ternary mat-vec (decode path) ---- +// dst[row] = sum_col W[row,col] * x[col], reading PACKED TQ{K}P weights directly (no fp16 spill) -> +// ~K*2/8 bytes/weight instead of 2 (fp16), i.e. much less weight bandwidth on the bandwidth-bound +// decode. Activation x is fp32 (src1). Strided-256 K-plane layout mirrors dequantize_block_tqkp_gv. +// One warp per output row; correctness-first (not yet memory-coalesced). MVP: f16-scale families only +// (TQ{K}P == g256 + _G{32,64,128}); q4k6 still uses the dequant->cuBLAS path. +template +static __global__ void mul_mat_vec_tqkp_kernel( + const void * __restrict__ vx, const float * __restrict__ x, float * __restrict__ dst, + const int ncols, const int nrows) { + const int row = blockIdx.x; + if (row >= nrows) return; + const int tid = threadIdx.x; + const int warp = tid >> 5; + const int lane = tid & 31; + const int nwarps = blockDim.x >> 5; + const int nsub = 256 / G; + const int qb = 64; + const int bpb = K*qb + K*nsub*2; + const int nsuper = ncols / 256; + const uint8_t * __restrict__ row_base = (const uint8_t *) vx + (size_t) row * nsuper * bpb; + float acc = 0.0f; + for (int sb = warp; sb < nsuper; sb += nwarps) { // warps split super-blocks + const uint8_t * __restrict__ blk = row_base + (size_t) sb * bpb; + const half * __restrict__ d = (const half *) (blk + K*qb); + const float * __restrict__ xs = x + sb*256; + #pragma unroll + for (int tb = lane; tb < 64; tb += 32) { // coalesced byte reads within a super-block + const int n = tb >> 5, l = tb & 31; + #pragma unroll + for (int bp = 0; bp < 4; ++bp) { + const int j = 128*n + bp*32 + l; + const int sub = j / G; + float w = 0.0f; + #pragma unroll + for (int kk = 0; kk < K; ++kk) { + const int code = (blk[kk*qb + tb] >> (bp*2)) & 3; + w += __half2float(d[kk*nsub + sub]) * (float) (code - 1); + } + acc += w * xs[j]; + } + } + } + acc = warp_reduce_sum(acc); + __shared__ float sh[32]; + if (lane == 0) sh[warp] = acc; + __syncthreads(); + if (warp == 0) { + float v = (lane < nwarps) ? sh[lane] : 0.0f; + v = warp_reduce_sum(v); + if (lane == 0) dst[row] = v; + } +} + + +// q4k6-scale fused mat-vec (g=32): per-plane fp16 super-scale sign-flipped by sub bit7, times a 6-bit +// unsigned sub-scale. Mirrors dequantize_block_tqkp_q4k6's scale decode; same coalesced structure. +template +static __global__ void mul_mat_vec_tqkp_q4k6_kernel( + const void * __restrict__ vx, const float * __restrict__ x, float * __restrict__ dst, + const int ncols, const int nrows) { + const int row = blockIdx.x; + if (row >= nrows) return; + const int tid = threadIdx.x; + const int warp = tid >> 5; + const int lane = tid & 31; + const int nwarps = blockDim.x >> 5; + const int nsub = 256 / Q4K6_G; + const int qb = 64; + const int bpb = K*qb + K*nsub + K*2; + const int nsuper = ncols / 256; + const uint8_t * __restrict__ row_base = (const uint8_t *) vx + (size_t) row * nsuper * bpb; + float acc = 0.0f; + for (int sb = warp; sb < nsuper; sb += nwarps) { + const uint8_t * __restrict__ blk = row_base + (size_t) sb * bpb; + const uint8_t * __restrict__ sbsc = blk + K*qb; + const half * __restrict__ sup = (const half *)(blk + K*qb + K*nsub); + const float * __restrict__ xs = x + sb*256; + #pragma unroll + for (int tb = lane; tb < 64; tb += 32) { + const int n = tb >> 5, l = tb & 31; + #pragma unroll + for (int bp = 0; bp < 4; ++bp) { + const int j = 128*n + bp*32 + l; + const int c = j / Q4K6_G; + float w = 0.0f; + #pragma unroll + for (int kk = 0; kk < K; ++kk) { + const uint8_t s = sbsc[kk*nsub + c]; + uint32_t bits = __float_as_uint(__half2float(sup[kk])); + bits ^= ((uint32_t)(s & 0x80u)) << 24; + const float sc = __uint_as_float(bits) * (float)(s & 0x3F); + const int t = ((blk[kk*qb + tb] >> (bp*2)) & 3) - 1; + w += sc * (float) t; + } + acc += w * xs[j]; + } + } + } + acc = warp_reduce_sum(acc); + __shared__ float sh[32]; + if (lane == 0) sh[warp] = acc; + __syncthreads(); + if (warp == 0) { + float v = (lane < nwarps) ? sh[lane] : 0.0f; + v = warp_reduce_sum(v); + if (lane == 0) dst[row] = v; + } +} + +void ggml_cuda_mul_mat_vec_tqkp(ggml_backend_cuda_context & ctx, + const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) { + const int ncols = src0->ne[0]; + const int nrows = src0->ne[1]; + const float * x = (const float *) src1->data; + float * d = (float *) dst->data; + cudaStream_t stream = ctx.stream(); + const dim3 grid(nrows, 1, 1); + switch (src0->type) { + case GGML_TYPE_TQ1P: mul_mat_vec_tqkp_kernel<1,256><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ2P: mul_mat_vec_tqkp_kernel<2,256><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ3P: mul_mat_vec_tqkp_kernel<3,256><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ4P: mul_mat_vec_tqkp_kernel<4,256><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ1P_G32: mul_mat_vec_tqkp_kernel<1,32 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ1P_G64: mul_mat_vec_tqkp_kernel<1,64 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ1P_G128: mul_mat_vec_tqkp_kernel<1,128><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ2P_G32: mul_mat_vec_tqkp_kernel<2,32 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ2P_G64: mul_mat_vec_tqkp_kernel<2,64 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ2P_G128: mul_mat_vec_tqkp_kernel<2,128><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ3P_G32: mul_mat_vec_tqkp_kernel<3,32 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ3P_G64: mul_mat_vec_tqkp_kernel<3,64 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ3P_G128: mul_mat_vec_tqkp_kernel<3,128><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ4P_G32: mul_mat_vec_tqkp_kernel<4,32 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ4P_G64: mul_mat_vec_tqkp_kernel<4,64 ><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ4P_G128: mul_mat_vec_tqkp_kernel<4,128><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ1P_Q4K6: mul_mat_vec_tqkp_q4k6_kernel<1><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ2P_Q4K6: mul_mat_vec_tqkp_q4k6_kernel<2><<>>(src0->data,x,d,ncols,nrows); break; + case GGML_TYPE_TQ3P_Q4K6: mul_mat_vec_tqkp_q4k6_kernel<3><<>>(src0->data,x,d,ncols,nrows); break; + default: GGML_ABORT("mul_mat_vec_tqkp: unsupported type"); + } +} + +bool ggml_cuda_is_tqkp_fused(enum ggml_type t) { + switch (t) { + case GGML_TYPE_TQ1P: case GGML_TYPE_TQ2P: case GGML_TYPE_TQ3P: case GGML_TYPE_TQ4P: + case GGML_TYPE_TQ1P_G32: case GGML_TYPE_TQ1P_G64: case GGML_TYPE_TQ1P_G128: + case GGML_TYPE_TQ2P_G32: case GGML_TYPE_TQ2P_G64: case GGML_TYPE_TQ2P_G128: + case GGML_TYPE_TQ3P_G32: case GGML_TYPE_TQ3P_G64: case GGML_TYPE_TQ3P_G128: + case GGML_TYPE_TQ4P_G32: case GGML_TYPE_TQ4P_G64: case GGML_TYPE_TQ4P_G128: + case GGML_TYPE_TQ1P_Q4K6: case GGML_TYPE_TQ2P_Q4K6: case GGML_TYPE_TQ3P_Q4K6: + return true; + default: return false; + } +} + to_fp16_cuda_t ggml_get_to_fp16_cuda(ggml_type type) { switch (type) { case GGML_TYPE_Q1_0: @@ -736,6 +992,44 @@ to_fp16_cuda_t ggml_get_to_fp16_cuda(ggml_type type) { return dequantize_row_q5_K_cuda; case GGML_TYPE_Q6_K: return dequantize_row_q6_K_cuda; + case GGML_TYPE_TQ1P: + return dequantize_row_tq1p_cuda; + case GGML_TYPE_TQ2P: + return dequantize_row_tq2p_cuda; + case GGML_TYPE_TQ3P: + return dequantize_row_tq3p_cuda; + case GGML_TYPE_TQ4P: + return dequantize_row_tq4p_cuda; + case GGML_TYPE_TQ1P_G32: + return dequantize_row_tq1p_g32_cuda; + case GGML_TYPE_TQ1P_G64: + return dequantize_row_tq1p_g64_cuda; + case GGML_TYPE_TQ1P_G128: + return dequantize_row_tq1p_g128_cuda; + case GGML_TYPE_TQ2P_G32: + return dequantize_row_tq2p_g32_cuda; + case GGML_TYPE_TQ2P_G64: + return dequantize_row_tq2p_g64_cuda; + case GGML_TYPE_TQ2P_G128: + return dequantize_row_tq2p_g128_cuda; + case GGML_TYPE_TQ3P_G32: + return dequantize_row_tq3p_g32_cuda; + case GGML_TYPE_TQ3P_G64: + return dequantize_row_tq3p_g64_cuda; + case GGML_TYPE_TQ3P_G128: + return dequantize_row_tq3p_g128_cuda; + case GGML_TYPE_TQ4P_G32: + return dequantize_row_tq4p_g32_cuda; + case GGML_TYPE_TQ4P_G64: + return dequantize_row_tq4p_g64_cuda; + case GGML_TYPE_TQ4P_G128: + return dequantize_row_tq4p_g128_cuda; + case GGML_TYPE_TQ1P_Q4K6: + return dequantize_row_tq1p_q4k6_cuda; + case GGML_TYPE_TQ2P_Q4K6: + return dequantize_row_tq2p_q4k6_cuda; + case GGML_TYPE_TQ3P_Q4K6: + return dequantize_row_tq3p_q4k6_cuda; case GGML_TYPE_IQ2_XXS: return dequantize_row_iq2_xxs_cuda; case GGML_TYPE_IQ2_XS: @@ -791,6 +1085,44 @@ to_fp32_cuda_t ggml_get_to_fp32_cuda(ggml_type type) { return dequantize_row_q5_K_cuda; case GGML_TYPE_Q6_K: return dequantize_row_q6_K_cuda; + case GGML_TYPE_TQ1P: + return dequantize_row_tq1p_cuda; + case GGML_TYPE_TQ2P: + return dequantize_row_tq2p_cuda; + case GGML_TYPE_TQ3P: + return dequantize_row_tq3p_cuda; + case GGML_TYPE_TQ4P: + return dequantize_row_tq4p_cuda; + case GGML_TYPE_TQ1P_G32: + return dequantize_row_tq1p_g32_cuda; + case GGML_TYPE_TQ1P_G64: + return dequantize_row_tq1p_g64_cuda; + case GGML_TYPE_TQ1P_G128: + return dequantize_row_tq1p_g128_cuda; + case GGML_TYPE_TQ2P_G32: + return dequantize_row_tq2p_g32_cuda; + case GGML_TYPE_TQ2P_G64: + return dequantize_row_tq2p_g64_cuda; + case GGML_TYPE_TQ2P_G128: + return dequantize_row_tq2p_g128_cuda; + case GGML_TYPE_TQ3P_G32: + return dequantize_row_tq3p_g32_cuda; + case GGML_TYPE_TQ3P_G64: + return dequantize_row_tq3p_g64_cuda; + case GGML_TYPE_TQ3P_G128: + return dequantize_row_tq3p_g128_cuda; + case GGML_TYPE_TQ4P_G32: + return dequantize_row_tq4p_g32_cuda; + case GGML_TYPE_TQ4P_G64: + return dequantize_row_tq4p_g64_cuda; + case GGML_TYPE_TQ4P_G128: + return dequantize_row_tq4p_g128_cuda; + case GGML_TYPE_TQ1P_Q4K6: + return dequantize_row_tq1p_q4k6_cuda; + case GGML_TYPE_TQ2P_Q4K6: + return dequantize_row_tq2p_q4k6_cuda; + case GGML_TYPE_TQ3P_Q4K6: + return dequantize_row_tq3p_q4k6_cuda; case GGML_TYPE_IQ2_XXS: return dequantize_row_iq2_xxs_cuda; case GGML_TYPE_IQ2_XS: diff --git a/src/ggml-cuda/convert.cuh b/src/ggml-cuda/convert.cuh index f5d37c7..707fcc8 100644 --- a/src/ggml-cuda/convert.cuh +++ b/src/ggml-cuda/convert.cuh @@ -64,3 +64,8 @@ template return float(x); } } + +// sasori fused ternary mat-vec (decode): reads packed TQ{K}P weights directly (f16-scale families). +void ggml_cuda_mul_mat_vec_tqkp(ggml_backend_cuda_context & ctx, + const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst); +bool ggml_cuda_is_tqkp_fused(enum ggml_type t); diff --git a/src/ggml-cuda/ggml-cuda.cu b/src/ggml-cuda/ggml-cuda.cu index f5293ad..0509752 100644 --- a/src/ggml-cuda/ggml-cuda.cu +++ b/src/ggml-cuda/ggml-cuda.cu @@ -2500,6 +2500,20 @@ static bool ggml_cuda_should_fuse_mul_mat_vec_f(const ggml_tensor * tensor) { } static bool ggml_cuda_should_fuse_mul_mat_vec_q(const ggml_tensor * tensor) { + // sasori: ternary trit-plane types have no fused mmvq kernel -> never fuse (would GGML_ABORT); + // they take the dequant -> cuBLAS path via ggml_get_to_fp16_cuda. + switch (tensor->src[0]->type) { + case GGML_TYPE_TQ1P: case GGML_TYPE_TQ2P: case GGML_TYPE_TQ3P: + case GGML_TYPE_TQ4P: case GGML_TYPE_TQ1P_G32: case GGML_TYPE_TQ1P_G64: + case GGML_TYPE_TQ1P_G128: case GGML_TYPE_TQ2P_G32: case GGML_TYPE_TQ2P_G64: + case GGML_TYPE_TQ2P_G128: case GGML_TYPE_TQ3P_G32: case GGML_TYPE_TQ3P_G64: + case GGML_TYPE_TQ3P_G128: case GGML_TYPE_TQ4P_G32: case GGML_TYPE_TQ4P_G64: + case GGML_TYPE_TQ4P_G128: case GGML_TYPE_TQ1P_Q4K6: case GGML_TYPE_TQ2P_Q4K6: + case GGML_TYPE_TQ3P_Q4K6: + return false; + default: + break; + } ggml_tensor * src0 = tensor->src[0]; ggml_tensor * src1 = tensor->src[1]; const ggml_tensor * dst = tensor; @@ -2604,6 +2618,17 @@ static void ggml_cuda_mul_mat(ggml_backend_cuda_context & ctx, const ggml_tensor return; } + // sasori: fused ternary mat-vec for the DECODE case (ne11==1) — reads packed TQ weights directly + // instead of the dequant->cuBLAS path (much less weight bandwidth). f16-scale families only. + // Toggle off with SASORI_NO_FUSED_MMVEC=1 (A/B correctness vs the dequant->cuBLAS reference). + static const bool sasori_fused_off = getenv("SASORI_NO_FUSED_MMVEC") != nullptr; + if (!sasori_fused_off && !split && ggml_cuda_is_tqkp_fused(src0->type) && + src1->type == GGML_TYPE_F32 && src1->ne[1] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1 && + ggml_is_contiguous(src1) && ggml_is_contiguous(src0)) { + ggml_cuda_mul_mat_vec_tqkp(ctx, src0, src1, dst); + return; + } + if (!split && use_mul_mat_vec_f) { // the custom F16 vector kernel can be used over batched cuBLAS GEMM // but this is only faster for GPUs without tensor cores or with a thin src0 matrix (particularly KQV in attention) @@ -5166,6 +5191,25 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g case GGML_TYPE_Q5_K: case GGML_TYPE_Q6_K: case GGML_TYPE_Q8_K: + case GGML_TYPE_TQ1P: + case GGML_TYPE_TQ2P: + case GGML_TYPE_TQ3P: + case GGML_TYPE_TQ4P: + case GGML_TYPE_TQ1P_G32: + case GGML_TYPE_TQ1P_G64: + case GGML_TYPE_TQ1P_G128: + case GGML_TYPE_TQ2P_G32: + case GGML_TYPE_TQ2P_G64: + case GGML_TYPE_TQ2P_G128: + case GGML_TYPE_TQ3P_G32: + case GGML_TYPE_TQ3P_G64: + case GGML_TYPE_TQ3P_G128: + case GGML_TYPE_TQ4P_G32: + case GGML_TYPE_TQ4P_G64: + case GGML_TYPE_TQ4P_G128: + case GGML_TYPE_TQ1P_Q4K6: + case GGML_TYPE_TQ2P_Q4K6: + case GGML_TYPE_TQ3P_Q4K6: case GGML_TYPE_IQ1_M: case GGML_TYPE_IQ1_S: case GGML_TYPE_IQ2_S: diff --git a/src/ggml-cuda/mmvq.cu b/src/ggml-cuda/mmvq.cu index 4b04265..555d74b 100644 --- a/src/ggml-cuda/mmvq.cu +++ b/src/ggml-cuda/mmvq.cu @@ -278,6 +278,21 @@ int get_mmvq_mmid_max_batch(ggml_type type, int cc) { } bool ggml_cuda_should_use_mmvq(enum ggml_type type, int cc, int64_t ne11) { + // sasori ternary trit-plane types have no mmvq (or mmq) kernel yet: force the + // dequant -> cuBLAS GEMM path (they are wired into ggml_get_to_fp16_cuda). Without + // this, use_mul_mat_vec_q would be true and mul_mat_vec_q_switch_type would GGML_ABORT. + switch (type) { + case GGML_TYPE_TQ1P: case GGML_TYPE_TQ2P: case GGML_TYPE_TQ3P: + case GGML_TYPE_TQ4P: case GGML_TYPE_TQ1P_G32: case GGML_TYPE_TQ1P_G64: + case GGML_TYPE_TQ1P_G128: case GGML_TYPE_TQ2P_G32: case GGML_TYPE_TQ2P_G64: + case GGML_TYPE_TQ2P_G128: case GGML_TYPE_TQ3P_G32: case GGML_TYPE_TQ3P_G64: + case GGML_TYPE_TQ3P_G128: case GGML_TYPE_TQ4P_G32: case GGML_TYPE_TQ4P_G64: + case GGML_TYPE_TQ4P_G128: case GGML_TYPE_TQ1P_Q4K6: case GGML_TYPE_TQ2P_Q4K6: + case GGML_TYPE_TQ3P_Q4K6: + return false; + default: + break; + } if (GGML_CUDA_CC_IS_CDNA(cc)) { if (GGML_CUDA_CC_IS_CDNA1(cc)) { switch (type) {