Commit 71ad0590f for llama.cpp

commit 71ad0590f4808b6202f9213d166913858c73b1bc
Author: Pranesh Gonegandla <pranesh.iitp@gmail.com>
Date:   Thu Oct 8 22:04:27 2026 +0530

    CUDA: improve top-k algorithm selection (#28713)

    * CUDA: radix top-k for large row counts

    Replaces CUB's per-row DeviceTopKKernel with a grid-over-rows radix select,
    gated on GGML_CUDA_TOPK_RADIX_MIN_ROWS. On qwen4exp at 34,816 tokens this cuts
    top-k from 1,671,253 launches / 5,761.8 ms to 2,329 / 941.8 ms.

    * CUDA: select the TOP_K implementation by shape

    Replace the nrows/ncols special case with the decision boundary from #28547
    (as implemented in #29278): bitonic for short rows, radix select for several
    long rows, and DeviceTopK or CUB argsort for a single long row. The
    thresholds stay overridable at build time.

    Two refinements on top of that boundary:
    - bitonic stays in use for rows up to a padded 1024 while the rows fit in one
      wave of blocks (nrows <= number of SMs); radix select pays a fixed cost of
      about a dozen launches that only amortizes over more rows
    - with DeviceTopK available, it handles up to two rows

    Radix select now processes rows in chunks so its scratch memory stays bounded,
    and the bitonic path keeps its chunking. HIP and MUSA keep their previous
    thresholds.

    Add perf cases around the bitonic/radix crossover to test-backend-ops.

    * CUDA: make top-k comments less verbose

    * CUDA: remove the TOP_K width limit from supports_op

    * CUDA: use DeviceTopK for single-row TOP_K if available

    * CUDA: avoid ncols overflow in the TOP_K bitonic check

    * CUDA: share the row chunking helper between argsort and top-k

    * CUDA: do the TOP_K radix blocks_per_row math in int64_t

    * CUDA: rename GGML_CUDA_TOP_K_NROWS_THRESHOLD_DEVICETOPK to GGML_CUDA_TOP_K_NROWS_THRESHOLD

    * CUDA: share one sort helper between the bitonic and CUB TOP_K paths

    * CUDA: update the TOP_K TODO, threshold and chunking comments

    * tests: add TOP_K cases that span several row chunks

    * CUDA: use int64_t col in the TOP_K radix loops, fix threshold comment

    * CUDA: limit TOP_K and ARGSORT support to ne[0] <= INT_MAX

    ---------

    Co-authored-by: praneshgo <227579474+praneshgo@users.noreply.github.com>
    Co-authored-by: Pranesh Gonegandla <pgonegandla@nvidia.com>

diff --git a/ggml/src/ggml-cuda/argsort.cu b/ggml/src/ggml-cuda/argsort.cu
index 589b99ff0..101afeba2 100644
--- a/ggml/src/ggml-cuda/argsort.cu
+++ b/ggml/src/ggml-cuda/argsort.cu
@@ -28,21 +28,21 @@ static __global__ void init_offsets(int * offsets, const int ncols, const int nr
 }
 #endif  // STRIDED_ITERATOR_AVAILABLE

-#ifdef GGML_CUDA_USE_CUB
-
-// returns the suggested maximum number of rows to process during one argsort_f32_i32_cuda_cub() call
-int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows) {
-    // perform argsort in chunks up to approximately this size (currently 64MB)
+// returns the suggested maximum number of rows to process at once, given the temporary buffer bytes per row
+int ggml_cuda_chunk_nrows(const size_t row_bytes, const int64_t nrows) {
+    // process rows in chunks up to approximately this size (currently 64MB)
     // to avoid excessive temporary buffers memory usage
     const int chunk_bytes = 1 << 26;

     // calculate how many rows will fit in one chunk (must be at least one)
-    const int chunk_nrows = std::max((int) (chunk_bytes / nb01), 1);
+    const int chunk_nrows = std::max((int) (chunk_bytes / row_bytes), 1);

     // limit the resulting amount to total nrows
     return std::min((int64_t) chunk_nrows, nrows);
 }

+#ifdef GGML_CUDA_USE_CUB
+
 void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
                               const float *    x,
                               int *            dst,
@@ -290,7 +290,7 @@ void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
         return;
     }

-    const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows);
+    const int chunk_nrows = ggml_cuda_chunk_nrows(src0->nb[1], nrows);

     ggml_cuda_pool & pool = ctx.pool();

diff --git a/ggml/src/ggml-cuda/argsort.cuh b/ggml/src/ggml-cuda/argsort.cuh
index c9adfcb98..86df5fda2 100644
--- a/ggml/src/ggml-cuda/argsort.cuh
+++ b/ggml/src/ggml-cuda/argsort.cuh
@@ -4,8 +4,9 @@

 void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst);

+int ggml_cuda_chunk_nrows(const size_t row_bytes, const int64_t nrows);
+
 #ifdef GGML_CUDA_USE_CUB
-int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows);
 void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
                               const float *    x,
                               int *            dst,
diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu
index 995037f1b..8e0c7658e 100644
--- a/ggml/src/ggml-cuda/ggml-cuda.cu
+++ b/ggml/src/ggml-cuda/ggml-cuda.cu
@@ -5689,11 +5689,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
         case GGML_OP_SUM:
             return ggml_is_contiguous_rows(op->src[0]);
         case GGML_OP_TOP_K:
-#if defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB)
-            return true;
-#else
-            return op->src[0]->ne[0] <= 1024;
-#endif // defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB)
+            return op->src[0]->ne[0] <= INT_MAX;
         case GGML_OP_ARGSORT:
 #ifndef GGML_CUDA_USE_CUB
             {
@@ -5705,7 +5701,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
                 return ncols_pad * sizeof(int) <= ggml_cuda_info().devices[dev_ctx->device].smpb;
             }
 #else
-            return true;
+            return op->src[0]->ne[0] <= INT_MAX;
 #endif
         case GGML_OP_SUM_ROWS:
             return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && ggml_is_contiguous_rows(op->src[0]);
diff --git a/ggml/src/ggml-cuda/top-k.cu b/ggml/src/ggml-cuda/top-k.cu
index 3ffbba839..b0f609990 100644
--- a/ggml/src/ggml-cuda/top-k.cu
+++ b/ggml/src/ggml-cuda/top-k.cu
@@ -1,6 +1,29 @@
 #include "argsort.cuh"
 #include "top-k.cuh"

+// Adjusted implementation thresholds from #28547, can be overridden at build time
+#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC
+#    if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
+// not measured on HIP/MUSA, keep the old split
+#        define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC 1024
+#    else
+#        define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC 512
+#    endif
+#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC
+
+#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT
+#    define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT 4096
+#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT
+
+// bitonic up to this width while nrows fits in one wave of SMs, 0 disables
+#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS
+#    if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
+#        define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS 0
+#    else
+#        define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS 1024
+#    endif
+#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS
+
 #ifdef GGML_CUDA_USE_CUB
 #    include <cub/cub.cuh>
 // DeviceTopK has a race condition before CCCL 3.4.3.
@@ -14,6 +37,15 @@ using namespace cub;
 #    endif  // CCCL >= 3.4.3
 #endif      // GGML_CUDA_USE_CUB

+// max rows for the per-row DeviceTopK / CUB argsort path before switching to radix / bitonic
+#ifndef GGML_CUDA_TOP_K_NROWS_THRESHOLD
+#    ifdef CUB_TOP_K_AVAILABLE
+#        define GGML_CUDA_TOP_K_NROWS_THRESHOLD 2
+#    else
+#        define GGML_CUDA_TOP_K_NROWS_THRESHOLD 1
+#    endif
+#endif // GGML_CUDA_TOP_K_NROWS_THRESHOLD
+
 #ifdef CUB_TOP_K_AVAILABLE

 static void top_k_cub(ggml_cuda_pool & pool,
@@ -40,7 +72,7 @@ static void top_k_cub(ggml_cuda_pool & pool,
                          ncols, k, env));
 }

-#elif defined(GGML_CUDA_USE_CUB)  // CUB_TOP_K_AVAILABLE
+#endif                            // CUB_TOP_K_AVAILABLE

 static int next_power_of_2(int x) {
     int n = 1;
@@ -50,10 +82,6 @@ static int next_power_of_2(int x) {
     return n;
 }

-#endif                            // CUB_TOP_K_AVAILABLE
-
-#if !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP)
-
 static __device__ __forceinline__ uint32_t top_k_float_to_ordered(float value) {
     const uint32_t bits = __float_as_uint(value);
     const uint32_t mask = (uint32_t) (-(int32_t) (bits >> 31)) | 0x80000000U;
@@ -95,7 +123,7 @@ static __global__ void top_k_radix_histogram(
     __syncthreads();

     const top_k_radix_state state = states[row];
-    for (int col = row_block * BLOCK_SIZE + tid;
+    for (int64_t col = row_block * BLOCK_SIZE + tid;
          col < ncols;
          col += blocks_per_row * BLOCK_SIZE) {
         const uint32_t key = top_k_float_to_ordered(row_src[col]);
@@ -165,7 +193,7 @@ static __global__ void top_k_radix_gather(
     int * row_dst = dst + (size_t) row * k;
     top_k_radix_state * state = &states[row];

-    for (int col = row_block * BLOCK_SIZE + tid;
+    for (int64_t col = row_block * BLOCK_SIZE + tid;
          col < ncols;
          col += blocks_per_row * BLOCK_SIZE) {
         const uint32_t key = top_k_float_to_ordered(row_src[col]);
@@ -183,36 +211,72 @@ static __global__ void top_k_radix_gather(

 static void top_k_radix_cuda(
         ggml_cuda_pool & pool,
-        const float * src, int * dst, int ncols, int nrows, int k, cudaStream_t stream) {
+        const float * src, int * dst, int ncols, int64_t nrows, int k, cudaStream_t stream) {
     constexpr int BLOCK_SIZE = 256;
     constexpr int RADIX_BITS = 8;
     constexpr int NBINS = 1 << RADIX_BITS;
-    const int blocks_per_row = std::min((ncols + 1023) / 1024, 64);
+    const int blocks_per_row = (int) std::min<int64_t>(((int64_t) ncols + 1023) / 1024, 64);
+
+    // chunk the rows to bound the histogram memory to 64 MB
+    const int64_t chunk_nrows = ggml_cuda_chunk_nrows((size_t) blocks_per_row * NBINS * sizeof(int), nrows);

-    ggml_cuda_pool_alloc<top_k_radix_state> states_alloc(pool, nrows);
-    ggml_cuda_pool_alloc<int> histograms_alloc(pool, (size_t) nrows * blocks_per_row * NBINS);
+    ggml_cuda_pool_alloc<top_k_radix_state> states_alloc(pool, chunk_nrows);
+    ggml_cuda_pool_alloc<int> histograms_alloc(pool, (size_t) chunk_nrows * blocks_per_row * NBINS);
     top_k_radix_state * states = states_alloc.get();
     int * histograms = histograms_alloc.get();

-    top_k_radix_init<<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows, k);
+    for (int64_t i = 0; i < nrows; i += chunk_nrows) {
+        const int iter_nrows = std::min(chunk_nrows, nrows - i);
+
+        top_k_radix_init<<<(iter_nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, iter_nrows, k);
+
+        const dim3 row_grid(blocks_per_row * iter_nrows);
+        for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
+            top_k_radix_histogram<BLOCK_SIZE, RADIX_BITS>
+                <<<row_grid, BLOCK_SIZE, 0, stream>>>(
+                    src, states, histograms, ncols, blocks_per_row, shift);
+            top_k_radix_select<BLOCK_SIZE, RADIX_BITS>
+                <<<iter_nrows, BLOCK_SIZE, 0, stream>>>(histograms, states, blocks_per_row, shift);
+        }

-    const dim3 row_grid(blocks_per_row * nrows);
-    for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
-        top_k_radix_histogram<BLOCK_SIZE, RADIX_BITS>
+        top_k_radix_reset_counters
+            <<<(iter_nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, iter_nrows);
+        top_k_radix_gather<BLOCK_SIZE>
             <<<row_grid, BLOCK_SIZE, 0, stream>>>(
-                src, states, histograms, ncols, blocks_per_row, shift);
-        top_k_radix_select<BLOCK_SIZE, RADIX_BITS>
-            <<<nrows, BLOCK_SIZE, 0, stream>>>(histograms, states, blocks_per_row, shift);
-    }
+                src, dst, states, ncols, k, blocks_per_row);

-    top_k_radix_reset_counters
-        <<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows);
-    top_k_radix_gather<BLOCK_SIZE>
-        <<<row_grid, BLOCK_SIZE, 0, stream>>>(
-            src, dst, states, ncols, k, blocks_per_row);
+        src += (size_t) ncols * iter_nrows;
+        dst += (size_t) k     * iter_nrows;
+    }
 }

-#endif // !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP)
+static void top_k_argsort_cuda(
+        ggml_cuda_pool & pool,
+        const float * src, int * dst, int ncols, int64_t nrows, int k, bool use_cub, cudaStream_t stream) {
+    const int64_t chunk_nrows = ggml_cuda_chunk_nrows((size_t) ncols * sizeof(int), nrows);
+
+    ggml_cuda_pool_alloc<int> tmp_alloc(pool, (size_t) ncols * chunk_nrows);
+    int * tmp = tmp_alloc.get();
+
+    for (int64_t i = 0; i < nrows; i += chunk_nrows) {
+        const int iter_nrows = std::min(chunk_nrows, nrows - i);
+
+        if (use_cub) {
+#ifdef GGML_CUDA_USE_CUB
+            argsort_f32_i32_cuda_cub(pool, src, tmp, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
+#else
+            GGML_ABORT("CUB is not available");
+#endif // GGML_CUDA_USE_CUB
+        } else {
+            argsort_f32_i32_cuda_bitonic(src, tmp, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
+        }
+        CUDA_CHECK(cudaMemcpy2DAsync(dst, k * sizeof(int), tmp, ncols * sizeof(int), k * sizeof(int), iter_nrows,
+                                     cudaMemcpyDeviceToDevice, stream));
+
+        src += (size_t) ncols * iter_nrows;
+        dst += (size_t) k     * iter_nrows;
+    }
+}

 void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     const ggml_tensor * src0   = dst->src[0];
@@ -229,51 +293,45 @@ void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     const int64_t    nrows = ggml_nrows(src0);
     const int64_t    k     = dst->ne[0];
     ggml_cuda_pool & pool  = ctx.pool();
-#ifdef CUB_TOP_K_AVAILABLE
-    // TODO: Switch to `DeviceSegmentedTopK` for multi-row TopK once implemented
-    // https://github.com/NVIDIA/cccl/issues/6391
-    // TODO: investigate if there exists a point where parallelized argsort is faster than sequential top-k
-    for (int i = 0; i < nrows; i++) {
-        top_k_cub(pool, src0_d + i * ncols, dst_d + i * k, ncols, k, stream);
-    }
-#elif defined(GGML_CUDA_USE_CUB)  // CUB_TOP_K_AVAILABLE
-    // Fall back to argsort + copy
-    const int    ncols_pad      = next_power_of_2(ncols);
-    const size_t shared_mem     = ncols_pad * sizeof(int);
-    const size_t max_shared_mem = ggml_cuda_info().devices[ggml_cuda_get_device()].smpb;
-    const bool   use_bitonic    = shared_mem <= max_shared_mem && ncols <= 1024;
-    const int    chunk_nrows    = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows);

-    ggml_cuda_pool_alloc<int> temp_dst_alloc(pool, ncols * chunk_nrows);
-    int *                     tmp_dst = temp_dst_alloc.get();
+    const int device = ggml_cuda_get_device();

-    for (int64_t i = 0; i < nrows; i += chunk_nrows) {
-        int iter_nrows = std::min((int64_t) chunk_nrows, nrows - i);
-
-        if (use_bitonic) {
-            argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
-        } else {
-            argsort_f32_i32_cuda_cub(pool, src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
+#ifdef CUB_TOP_K_AVAILABLE
+    // a single row always uses DeviceTopK if available
+    const bool bitonic_short    = nrows > 1 && ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC;
+#else
+    const bool bitonic_short    = ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC;
+#endif // CUB_TOP_K_AVAILABLE
+    const bool bitonic_few_rows = nrows > GGML_CUDA_TOP_K_NROWS_THRESHOLD &&
+                                  ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS &&
+                                  nrows <= ggml_cuda_info().devices[device].nsm;
+
+    if (bitonic_short || bitonic_few_rows) {
+        // the padded row must fit in shared memory
+        const int ncols_pad = next_power_of_2(ncols);
+        if (ncols_pad * sizeof(int) <= ggml_cuda_info().devices[device].smpb) {
+            top_k_argsort_cuda(pool, src0_d, dst_d, ncols, nrows, k, false, stream);
+            return;
         }
-        CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), iter_nrows,
-                                     cudaMemcpyDeviceToDevice, stream));
-
-        src0_d += ncols * iter_nrows;
-        dst_d  += k     * iter_nrows;
     }
-#else                             // GGML_CUDA_USE_CUB
-#if defined(GGML_USE_HIP)
-    if (ncols > 1024) {
+
+    if (nrows > GGML_CUDA_TOP_K_NROWS_THRESHOLD) {
         top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
+        return;
+    }
+
+#ifdef CUB_TOP_K_AVAILABLE
+    // TODO: Assess perf of `DeviceBatchedTopK` for multi-row TopK & CCCL >= 3.5.0, re-running perf sweep of https://github.com/ggml-org/llama.cpp/pull/28713
+    for (int64_t i = 0; i < nrows; i++) {
+        top_k_cub(pool, src0_d + i * ncols, dst_d + i * k, ncols, k, stream);
+    }
+#elif defined(GGML_CUDA_USE_CUB)  // CUB_TOP_K_AVAILABLE
+    if (ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT) {
+        top_k_argsort_cuda(pool, src0_d, dst_d, ncols, nrows, k, true, stream);
     } else {
-#endif // defined(GGML_USE_HIP)
-        ggml_cuda_pool_alloc<int> temp_dst_alloc(pool, ncols * nrows);
-        int *                     tmp_dst = temp_dst_alloc.get();
-        argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream);
-        CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), nrows,
-                                     cudaMemcpyDeviceToDevice, stream));
-#if defined(GGML_USE_HIP)
+        top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
     }
-#endif // defined(GGML_USE_HIP)
-#endif
+#else                             // GGML_CUDA_USE_CUB
+    top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
+#endif                            // CUB_TOP_K_AVAILABLE
 }
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 40ad2bdf8..7b7e6cae1 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11266,6 +11266,10 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 8192,  2, 1, 1 }, 2051, true));
     test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 33024, 4, 1, 1 }, 2051, true));

+    // rows that CUDA processes in several chunks (bitonic, radix)
+    test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 500,  40000, 1, 1 }, 16));
+    test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 1100, 33000, 1, 1 }, 16));
+
     // qwen4exp QSA indexer top-k fusion (get_rows + f16 mask + top_k)
     test_cases.emplace_back(new test_topk_qsa(512,  2048,  1, 1, 1500));
     test_cases.emplace_back(new test_topk_qsa(512,  2048,  2, 1, 1500));
@@ -12228,6 +12232,12 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
             test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {cols, nrows, 1, 1}, 2048));
         }
     }
+    // bitonic vs radix crossover
+    for (auto cols : {520, 1000, 2048, 3000, 4096}) {
+        for (auto nrows : {2, 16, 32, 48, 64, 128, 256, 512}) {
+            test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {cols, nrows, 1, 1}, 16));
+        }
+    }
     // backend sampler: one row of the vocab (llama-sampler.cpp top_k)
     for (auto k : {20, 40}) {
         test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {151936, 1, 1, 1}, k));