diff --git a/ggml/src/ggml-cuda/dsa_attn.cu b/ggml/src/ggml-cuda/dsa_attn.cu index 17b6b6ee3..ba20f8018 100644 --- a/ggml/src/ggml-cuda/dsa_attn.cu +++ b/ggml/src/ggml-cuda/dsa_attn.cu @@ -25,7 +25,6 @@ static __global__ void k_prepare_one_batch_kv(int nk, int ncol, const int * idx, int i = idx[row*stride_idx + col]; if (i < 0) { i = 0; - //return; } auto k_row = (const half *)(k_in + stride_k * i); k_out += (row*ncol + col)*nk; @@ -34,6 +33,21 @@ static __global__ void k_prepare_one_batch_kv(int nk, int ncol, const int * idx, } } +static __global__ void k_prepare_one_batch_kv_q8_0(int nk, int ncol, const int * idx, const char * k_in, + half * k_out, size_t stride_k, size_t stride_idx) { + int row = blockIdx.y; + int col = blockIdx.x; + int i = idx[row*stride_idx + col]; + if (i < 0) { + i = 0; + } + auto k_row = (const block_q8_0 *)(k_in + stride_k * i); + k_out += (row*ncol + col)*nk; + for (int j = threadIdx.x; j < nk; j += blockDim.x) { + k_out[j] = k_row[j/32].d * (half)k_row[j/32].qs[j%32]; + } +} + static __global__ void k_prepare_one_batch_q(int ne0, int ne1, size_t nb1, size_t nb2, const float * q_in, half * q_out) { int i0 = blockIdx.x*blockDim.x + threadIdx.x; @@ -228,7 +242,8 @@ bool ggml_cuda_dsa_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) if (K->ne[1] < 4*indexer->ne[0]) return false; // for efficiency } if (K->ne[2] > 1 || K->ne[3] > 1 || mask->ne[2] > 1 || mask->ne[3] > 1 || Q->ne[3] > 1) return false; - if (K->type != GGML_TYPE_F16 || V->type != GGML_TYPE_F16 || mask->type != GGML_TYPE_F16 || Q->type != GGML_TYPE_F32) return false; + if ((K->type != GGML_TYPE_F16 && K->type != GGML_TYPE_Q8_0) || + (V->type != GGML_TYPE_F16 && V->type != GGML_TYPE_Q8_0) || mask->type != GGML_TYPE_F16 || Q->type != GGML_TYPE_F32) return false; if (K->ne[0] != Q->ne[0]) return false; //printf("%s(%s)\n", __func__, dst->name); @@ -274,13 +289,25 @@ bool ggml_cuda_dsa_attn_ext(ggml_backend_cuda_context & ctx, ggml_tensor * dst) int nrows = last - first; { dim3 grid(indexer->ne[0], nrows, 1); - k_prepare_one_batch_kv<<>>(K->ne[0], indexer->ne[0], - (const int *)indexer->data + stride_idx*first, - (const char *)K->data, k16.get(), K->nb[1], stride_idx); - if (!is_k_view) { - k_prepare_one_batch_kv<<>>(V->ne[0], indexer->ne[0], + if (K->type == GGML_TYPE_F16) { + k_prepare_one_batch_kv<<>>(K->ne[0], indexer->ne[0], (const int *)indexer->data + stride_idx*first, - (const char *)V->data, v16.get(), V->nb[1], stride_idx); + (const char *)K->data, k16.get(), K->nb[1], stride_idx); + } else { + k_prepare_one_batch_kv_q8_0<<>>(K->ne[0], indexer->ne[0], + (const int *)indexer->data + stride_idx*first, + (const char *)K->data, k16.get(), K->nb[1], stride_idx); + } + if (!is_k_view) { + if (V->type == GGML_TYPE_F16) { + k_prepare_one_batch_kv<<>>(V->ne[0], indexer->ne[0], + (const int *)indexer->data + stride_idx*first, + (const char *)V->data, v16.get(), V->nb[1], stride_idx); + } else { + k_prepare_one_batch_kv_q8_0<<>>(V->ne[0], indexer->ne[0], + (const int *)indexer->data + stride_idx*first, + (const char *)V->data, v16.get(), V->nb[1], stride_idx); + } } } {