vulkan: top_k radix select for k >= 1024 for Qwen 3.8 Flash Next (#28032)

* vulkan: add top-k radix sort shader for k >= 1024

* add Qwen 3.8 Flash Next top-k tests

* add top-k qsa fusion

* clean up code
This commit is contained in:
Ruben Ortlam
2026-08-31 07:04:34 +02:00
committed by GitHub
parent 9723942adc
commit daef7b6874
4 changed files with 471 additions and 10 deletions
+97
View File
@@ -6287,6 +6287,87 @@ struct test_top_k : public test_case {
}
};
// qwen4exp QSA indexer top-k fusion: expand per-block scores to cells, add the f16 mask, top-k.
struct test_topk_qsa : public test_case {
const int64_t n_blocks;
const int64_t n_kv;
const int64_t n_tps;
const int64_t n_stream;
const int width;
ggml_tensor * out {};
std::string op_desc(ggml_tensor * t) override {
GGML_UNUSED(t);
return "TOPK_QSA";
}
std::string vars() override {
return VARS_TO_STR5(n_blocks, n_kv, n_tps, n_stream, width);
}
test_topk_qsa(int64_t n_blocks = 512, int64_t n_kv = 2048, int64_t n_tps = 2, int64_t n_stream = 1, int width = 1500)
: n_blocks(n_blocks), n_kv(n_kv), n_tps(n_tps), n_stream(n_stream), width(width) {}
double max_err() override { return 0.0; }
bool run_whole_graph() override { return true; }
ggml_tensor * build_graph(ggml_context * ctx) override {
ggml_tensor * score = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_blocks, n_tps, n_stream);
ggml_set_name(score, "score");
ggml_tensor * cell_blk = ggml_new_tensor_2d(ctx, GGML_TYPE_I32, n_kv, n_stream);
ggml_set_name(cell_blk, "cell_blk");
ggml_tensor * kq_mask = ggml_new_tensor_3d(ctx, GGML_TYPE_F16, n_kv, n_tps, n_stream);
ggml_set_name(kq_mask, "kq_mask");
ggml_tensor * a = ggml_cont(ctx, ggml_permute(ctx, score, 1, 0, 2, 3));
ggml_tensor * e = ggml_get_rows(ctx, a, cell_blk);
e = ggml_cont(ctx, ggml_permute(ctx, e, 1, 0, 2, 3));
ggml_tensor * m = ggml_cast(ctx, kq_mask, GGML_TYPE_F32);
e = ggml_add(ctx, e, ggml_reshape_3d(ctx, m, n_kv, n_tps, n_stream));
out = ggml_top_k(ctx, e, width);
ggml_set_name(out, "out");
return out;
}
std::vector<ggml_tensor *> fusion_test_nodes() override { return { out }; }
// distinct mask ramp + small scores keep every cell value unique, so no top-k ties
void initialize_tensors(ggml_context * ctx) override {
for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {
if (t->op != GGML_OP_NONE) {
continue;
}
if (t->type == GGML_TYPE_I32) {
std::vector<int32_t> data(ggml_nelements(t));
for (auto & v : data) { v = rand() % n_blocks; }
ggml_backend_tensor_set(t, data.data(), 0, data.size() * sizeof(int32_t));
} else if (t->type == GGML_TYPE_F16) {
std::vector<ggml_fp16_t> data(ggml_nelements(t));
for (int64_t r = 0; r < ggml_nrows(t); r++) {
for (int64_t i = 0; i < n_kv; i++) {
data[r * n_kv + i] = ggml_fp32_to_fp16((float) i);
}
}
ggml_backend_tensor_set(t, data.data(), 0, data.size() * sizeof(ggml_fp16_t));
} else {
init_tensor_uniform(t, 0.0f, 0.5f);
}
}
}
// top-k output order is unspecified; compare as a set of indices
double err(const float * a, const float * b, size_t n) override {
std::vector<int32_t> ia(n), ib(n);
double diff = 0.0;
for (size_t i = 0; i < n; i++) {
ia[i] = (int32_t) a[i];
ib[i] = (int32_t) b[i];
diff += std::fabs(a[i] - ia[i]) + std::fabs(b[i] - ib[i]);
}
return diff + jdst(ia.data(), ib.data(), n);
}
};
enum MoeGatingFunc {
GATING_FUNC_SOFTMAX,
GATING_FUNC_SIGMOID,
@@ -9813,6 +9894,22 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {2049, 2, 1, 3}, k));
}
// Large-k, including multi-row and ties (qwen4exp)
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 1024, 1, 1, 1 }, 1024));
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 2048, 2, 1, 1 }, 1024));
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 4096, 1, 1, 1 }, 2048));
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 8192, 2, 1, 1 }, 2051));
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 33024, 1, 1, 1 }, 2051));
test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, { 33024, 4, 1, 1 }, 2051));
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));
// 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));
test_cases.emplace_back(new test_topk_qsa(256, 2048, 4, 2, 2000));
test_cases.emplace_back(new test_topk_qsa(64, 256, 2, 1, 200)); // small k: unfused fallback
// exhaustive top_k tests
//for (int i = 1; i < 9999; ++i) {
// test_cases.emplace_back(new test_top_k(GGML_TYPE_F32, {i, 2, 1, 3}, rand() % i + 1));