diff --git a/ggml/src/ggml-cuda/indexer_topk.cu b/ggml/src/ggml-cuda/indexer_topk.cu index 46b9980ea..ce89db37f 100644 --- a/ggml/src/ggml-cuda/indexer_topk.cu +++ b/ggml/src/ggml-cuda/indexer_topk.cu @@ -88,8 +88,10 @@ void ggml_cuda_op_indexer_topk(ggml_backend_cuda_context & ctx, ggml_tensor * ds constexpr int k_block_size = 256; if (k->type == GGML_TYPE_F16 && q->type == GGML_TYPE_F32) { - constexpr int k_max_rows = 16; - int max_rows = std::min(k_max_rows, q->ne[2]); + constexpr size_t k_max_buf_size = 1 << 28; + size_t per_row = size_t(n_kv)*(q->ne[1]*sizeof(half) + sizeof(int) + sizeof(float)) + q->ne[0]*q->ne[1]*sizeof(half); + int max_rows = (k_max_buf_size + per_row - 1)/per_row; + max_rows = std::min(max_rows, q->ne[2]); int nstep = (q->ne[2] + max_rows - 1)/max_rows; ggml_cuda_pool_alloc kq(ctx.pool(), int64_t(n_kv)*q->ne[1]*max_rows); @@ -103,15 +105,17 @@ void ggml_cuda_op_indexer_topk(ggml_backend_cuda_context & ctx, ggml_tensor * ds const half alpha = 1.0f; const half beta = 0.0f; + CUBLAS_CHECK(cublasSetStream(ctx.cublas_handle(ctx.device), ctx.stream())); + for (int istep = 0; istep < nstep; ++istep) { int first_row = max_rows*istep; - int last_row = std::min(first_row + k_max_rows, int(q->ne[2])); + if (first_row >= int(q->ne[2])) break; + int last_row = std::min(first_row + max_rows, int(q->ne[2])); int nrows = last_row - first_row; to_fp16_cuda((const float *)q->data + q->ne[0]*q->ne[1]*first_row, q_f16.get(), q->ne[0]*q->ne[1]*nrows, 1, ctx.stream()); CUDA_CHECK(cudaGetLastError()); - CUBLAS_CHECK(cublasSetStream(ctx.cublas_handle(ctx.device), ctx.stream())); CUBLAS_CHECK(cublasGemmEx(ctx.cublas_handle(ctx.device), CUBLAS_OP_T, CUBLAS_OP_N, k->ne[1], q->ne[1]*nrows, q->ne[0], &alpha, (const half *)k->data, CUDA_R_16F, k->ne[0],