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:
@@ -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));
|
||||
|
||||
Reference in New Issue
Block a user