Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
272 changes: 220 additions & 52 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp

Large diffs are not rendered by default.

120 changes: 0 additions & 120 deletions ggml/src/ggml-vulkan/vulkan-shaders/topk_radix.comp

This file was deleted.

165 changes: 165 additions & 0 deletions ggml/src/ggml-vulkan/vulkan-shaders/topk_radix_select.comp
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
#version 450

#extension GL_EXT_control_flow_attributes : enable
#extension GL_EXT_shader_16bit_storage : require
#extension GL_KHR_shader_subgroup_basic : enable
#extension GL_KHR_shader_subgroup_ballot : enable

#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 sg_cnt[64]; // per-subgroup hit counts for the slot scan

// 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();
}


// Emit everything above the threshold, then fill the rest from ties. Slots come from
// an exclusive scan over the candidate flags, one BLOCK_SIZE chunk at a time in ascending
// index order, so the output is identical on every run. The previous atomicAdd slot
// counter made the ORDER scheduling-dependent, and at the tie boundary the SET as well:
// the QSA width is top_k + ratio - 1, so the boundary block's ratio tied cells race for
// ratio - 1 slots and a different cell lost each run (non-repeatable output at depth).
// With the scan the lowest-indexed tied cells win.
const uint threshold = prefix;
uint base = 0;
[[dont_unroll]] for (uint pass = 0; pass < 2; ++pass) {
for (uint c = 0; c < ncols; c += BLOCK_SIZE) {
const uint i = c + tid;
bool hit = false;
if (i < ncols) {
const uint key = f2ui(load(row, i, false));
hit = (pass == 0) ? (key > threshold) : (key == threshold);
}
const uvec4 ballot = subgroupBallot(hit);
const uint rank_sg = subgroupBallotExclusiveBitCount(ballot);
const uint cnt_sg = subgroupBallotBitCount(ballot);
if (subgroupElect()) {
sg_cnt[gl_SubgroupID] = cnt_sg;
}
barrier();
uint sg_base = 0;
uint total = 0;
for (uint sg = 0; sg < gl_NumSubgroups; ++sg) {
const uint v = sg_cnt[sg];
sg_base += (sg < gl_SubgroupID) ? v : 0;
total += v;
}
const uint slot = base + sg_base + rank_sg;
if (hit && slot < p.k) {
data_d[row_out + slot] = int(i);
}
base += total;
barrier();
}
}
}

void main() {
for (uint row = gl_WorkGroupID.y; row < p.nrows; row += gl_NumWorkGroups.y) {
topk(row);
}
}
3 changes: 1 addition & 2 deletions ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1028,8 +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_f32", "topk_radix.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"}}));
Expand Down
Loading
Loading