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:
@@ -0,0 +1,144 @@
|
||||
#version 450
|
||||
|
||||
#extension GL_EXT_control_flow_attributes : enable
|
||||
#extension GL_EXT_shader_16bit_storage : require
|
||||
|
||||
#include "types.glsl"
|
||||
|
||||
layout(constant_id = 0) const int BLOCK_SIZE = 1024;
|
||||
layout(constant_id = 1) const int QSA = 0; // 1: fuse the qwen4 QSA indexer gather + f16 mask
|
||||
|
||||
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
|
||||
|
||||
layout (binding = 0) readonly buffer A {float data_a[];}; // input values, or QSA block scores [n_tps, n_blocks, n_stream]
|
||||
layout (binding = 1) writeonly buffer D {int data_d[];}; // [k, ...]
|
||||
layout (binding = 2) readonly buffer CB {int cell_blk[];}; // QSA: cell->block map [n_kv, n_stream]
|
||||
layout (binding = 3) readonly buffer M {float16_t mask[];}; // QSA: raw f16 kq_mask [n_kv, n_tps, n_stream]
|
||||
layout (binding = 4) buffer S {float scratch[];}; // QSA: [nrows, n_kv] gathered inputs
|
||||
|
||||
layout (push_constant) uniform parameter {
|
||||
uint ncols;
|
||||
uint k;
|
||||
uint nrows;
|
||||
uint n_tps; // QSA only
|
||||
uint n_blocks; // QSA only
|
||||
uint n_stream; // QSA only
|
||||
} p;
|
||||
|
||||
#define RADIX_BITS 8
|
||||
#define RADIX_SIZE (1 << RADIX_BITS)
|
||||
|
||||
shared uint histo[RADIX_SIZE];
|
||||
shared uint sh_bucket;
|
||||
shared uint sh_above;
|
||||
shared uint out_count;
|
||||
|
||||
// order-preserving float -> uint mapping
|
||||
uint f2ui(float x) {
|
||||
uint y = floatBitsToUint(x);
|
||||
if ((y & 0x80000000u) != 0u) {
|
||||
y ^= 0xFFFFFFFFu;
|
||||
} else {
|
||||
y |= 0x80000000u;
|
||||
}
|
||||
return y;
|
||||
}
|
||||
|
||||
// QSA element i of row (t,s): score[cell_blk[i,s], t, s] + mask[i,t,s]
|
||||
float gather(uint row, uint i) {
|
||||
const uint t = row % p.n_tps;
|
||||
const uint s = row / p.n_tps;
|
||||
const uint block = uint(cell_blk[s * p.ncols + i]);
|
||||
const float a = data_a[(s * p.n_blocks + block) * p.n_tps + t];
|
||||
const float m = float(mask[(s * p.n_tps + t) * p.ncols + i]);
|
||||
return a + m;
|
||||
}
|
||||
|
||||
float load(uint row, uint i, bool first) {
|
||||
if (QSA == 0) {
|
||||
return data_a[row * p.ncols + i];
|
||||
}
|
||||
// materialize the scattered gather on the first pass and reuse it after; each
|
||||
// invocation only touches its own scratch entries, so no barrier is needed
|
||||
const uint off = row * p.ncols + i;
|
||||
if (first) {
|
||||
const float v = gather(row, i);
|
||||
scratch[off] = v;
|
||||
return v;
|
||||
}
|
||||
return scratch[off];
|
||||
}
|
||||
|
||||
// one workgroup per row: radix-select the K-th largest, then compact it plus enough ties
|
||||
void topk(const uint row) {
|
||||
const uint tid = gl_LocalInvocationID.x;
|
||||
const uint ncols = p.ncols;
|
||||
const uint row_out = row * p.k;
|
||||
|
||||
uint prefix = 0; // fixed high bits of the threshold key
|
||||
uint desired = p.k; // count still needed from the candidate range
|
||||
|
||||
[[unroll]] for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
|
||||
for (uint i = tid; i < RADIX_SIZE; i += BLOCK_SIZE) {
|
||||
histo[i] = 0;
|
||||
}
|
||||
barrier();
|
||||
|
||||
const bool first = (shift == 32 - RADIX_BITS);
|
||||
const uint hi_mask = (shift + RADIX_BITS >= 32) ? 0u : (0xFFFFFFFFu << uint(shift + RADIX_BITS));
|
||||
const uint prefix_hi = prefix & hi_mask;
|
||||
for (uint i = tid; i < ncols; i += BLOCK_SIZE) {
|
||||
const uint key = f2ui(load(row, i, first));
|
||||
if ((key & hi_mask) == prefix_hi) {
|
||||
atomicAdd(histo[(key >> uint(shift)) & (RADIX_SIZE - 1)], 1u);
|
||||
}
|
||||
}
|
||||
barrier();
|
||||
|
||||
// top-down scan for the bucket holding the K-th value
|
||||
if (tid == 0) {
|
||||
uint acc = 0;
|
||||
uint b = 0;
|
||||
for (int bb = RADIX_SIZE - 1; bb >= 0; --bb) {
|
||||
const uint c = histo[bb];
|
||||
if (acc + c >= desired) { b = uint(bb); break; }
|
||||
acc += c;
|
||||
}
|
||||
sh_bucket = b;
|
||||
sh_above = acc;
|
||||
}
|
||||
barrier();
|
||||
|
||||
prefix |= sh_bucket << uint(shift);
|
||||
desired -= sh_above;
|
||||
barrier();
|
||||
}
|
||||
|
||||
if (tid == 0) {
|
||||
out_count = 0;
|
||||
}
|
||||
barrier();
|
||||
|
||||
// emit everything above the threshold, then fill the rest from ties
|
||||
const uint threshold = prefix;
|
||||
for (uint i = tid; i < ncols; i += BLOCK_SIZE) {
|
||||
if (f2ui(load(row, i, false)) > threshold) {
|
||||
data_d[row_out + atomicAdd(out_count, 1u)] = int(i);
|
||||
}
|
||||
}
|
||||
barrier();
|
||||
for (uint i = tid; i < ncols; i += BLOCK_SIZE) {
|
||||
if (f2ui(load(row, i, false)) == threshold) {
|
||||
const uint pos = atomicAdd(out_count, 1u);
|
||||
if (pos < p.k) {
|
||||
data_d[row_out + pos] = int(i);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void main() {
|
||||
for (uint row = gl_WorkGroupID.y; row < p.nrows; row += gl_NumWorkGroups.y) {
|
||||
topk(row);
|
||||
}
|
||||
}
|
||||
@@ -1028,6 +1028,7 @@ void process_shaders() {
|
||||
|
||||
string_to_spv("topk_argsort_f32", "topk_argsort.comp", {{"A_TYPE", "float"}});
|
||||
string_to_spv("topk_nary_search_f32", "topk_nary_search.comp", {{"A_TYPE", "float"}});
|
||||
string_to_spv("topk_radix_select_f32", "topk_radix_select.comp", {{"A_TYPE", "float"}});
|
||||
|
||||
string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}}));
|
||||
string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}}));
|
||||
|
||||
Reference in New Issue
Block a user