Implement exact Qwen sparse attention

This commit is contained in:
Georg Bauer
2026-09-03 21:36:33 +02:00
parent 87ccf67d0c
commit 714b39f7f3
3 changed files with 1086 additions and 84 deletions

View File

@@ -483,43 +483,205 @@ kernel void kernel_qwen_store_kv_bf16(
cache[base + width + index] = qwen_to_bf16(value[index]);
}
kernel void kernel_qwen_qsa_store_raw(
constant qwen_kernel_args &args [[buffer(0)]],
device ushort *raw [[buffer(1)]],
device const float *projected [[buffer(2)]],
device float *query [[buffer(3)]],
uint index [[thread_position_in_grid]]) {
const uint dim = args.u[0];
const uint query_width = args.u[1] * dim;
if (index < query_width) {
query[index] = projected[index];
} else if (index < query_width + dim) {
raw[(ulong)args.u[2] * dim + index - query_width] = qwen_to_bf16(projected[index]);
}
}
kernel void kernel_qwen_qsa_pool_key(
constant qwen_kernel_args &args [[buffer(0)]],
device ushort *pooled [[buffer(1)]],
device const ushort *raw [[buffer(2)]],
device const uchar *weight [[buffer(5)]],
uint column [[thread_position_in_grid]]) {
const uint dim = args.u[0];
if (column >= dim) return;
const ulong raw_start = (ulong)args.u[3] * dim;
float mean = 0.0f;
for (uint token = 0; token < args.u[2]; token++) {
mean += qwen_bf16(raw[raw_start + (ulong)token * dim + column]);
}
const ushort mean_bf16 = qwen_to_bf16(mean / (float)args.u[2]);
float variance = 0.0f;
for (uint i = 0; i < dim; i++) {
float item = 0.0f;
for (uint token = 0; token < args.u[2]; token++) {
item += qwen_bf16(raw[raw_start + (ulong)token * dim + i]);
}
item = qwen_bf16(qwen_to_bf16(item / (float)args.u[2]));
variance = fma(item, item, variance);
}
const float scale = rsqrt(variance / (float)dim + args.f[0]);
float value = qwen_bf16(mean_bf16) * scale *
(1.0f + qwen_bf16(qwen_weight_u16(weight, args.u[13], column)));
const uint rotary = 64u;
if (column < rotary) {
const uint rotary_half = rotary / 2u;
const uint pair = column < rotary_half ? column + rotary_half : column - rotary_half;
float paired_mean = 0.0f;
for (uint token = 0; token < args.u[2]; token++) {
paired_mean += qwen_bf16(raw[raw_start + (ulong)token * dim + pair]);
}
paired_mean = qwen_bf16(qwen_to_bf16(paired_mean / (float)args.u[2]));
const float paired = paired_mean * scale *
(1.0f + qwen_bf16(qwen_weight_u16(weight, args.u[13], pair)));
const float theta = (float)args.u[3] *
pow(args.f[1], -2.0f * (float)(column % rotary_half) / (float)rotary);
value = value * cos(theta) + (column < rotary_half ? -paired : paired) * sin(theta);
}
pooled[(ulong)args.u[1] * dim + column] = qwen_to_bf16(value);
}
kernel void kernel_qwen_qsa_scores(
constant qwen_kernel_args &args [[buffer(0)]],
device float *scores [[buffer(1)]],
device const float *query [[buffer(2)]],
device const ushort *pooled [[buffer(3)]],
uint block [[thread_position_in_grid]]) {
const uint dim = args.u[0];
if (block >= args.u[2]) return;
float score = 0.0f;
for (uint head = 0; head < args.u[1]; head++) {
float head_score = 0.0f;
for (uint i = 0; i < dim; i++) {
head_score = fma(query[(ulong)head * dim + i],
qwen_bf16(pooled[(ulong)block * dim + i]),
head_score);
}
score += max(head_score, 0.0f);
}
scores[block] = score * args.f[0];
}
kernel void kernel_qwen_qsa_sort_blocks(
constant qwen_kernel_args &args [[buffer(0)]],
device int *selected [[buffer(1)]],
uint gid [[thread_position_in_grid]]) {
if (gid != 0u) return;
for (uint i = 1; i < args.u[0]; i++) {
const int value = selected[i];
uint j = i;
while (j > 0u && selected[j - 1u] > value) {
selected[j] = selected[j - 1u];
j--;
}
selected[j] = value;
}
}
static inline uint qwen_qsa_token(
device const int *selected,
uint ordinal,
uint selected_count,
uint ratio,
uint tail_start) {
const uint selected_tokens = selected_count * ratio;
return ordinal < selected_tokens
? (uint)selected[ordinal / ratio] * ratio + ordinal % ratio
: tail_start + ordinal - selected_tokens;
}
kernel void kernel_qwen_sparse_attention(
constant qwen_kernel_args &args [[buffer(0)]],
device float *out [[buffer(1)]],
device const float *query [[buffer(2)]],
device const ushort *cache [[buffer(3)]],
device const int *selected [[buffer(4)]],
uint2 gid [[thread_position_in_grid]],
uint lane [[thread_index_in_threadgroup]]) {
const uint heads = args.u[0];
const uint kv_heads = args.u[1];
const uint dim = args.u[2];
const uint column = gid.x;
const uint head = gid.y;
if (column >= dim || head >= heads) return;
const uint kv_head = head / (heads / kv_heads);
const uint tokens = args.u[4] * args.u[5] + args.u[7];
const float attention_scale = rsqrt((float)dim);
threadgroup float probabilities[2051];
if (lane == 0u) {
float max_score = -INFINITY;
for (uint ordinal = 0; ordinal < tokens; ordinal++) {
const uint token = qwen_qsa_token(selected, ordinal, args.u[4], args.u[5], args.u[6]);
const ulong base = (ulong)token * kv_heads * dim * 2u + (ulong)kv_head * dim;
float score = 0.0f;
for (uint i = 0; i < dim; i++) {
score = fma(query[(ulong)head * dim + i], qwen_bf16(cache[base + i]), score);
}
probabilities[ordinal] = score * attention_scale;
max_score = max(max_score, probabilities[ordinal]);
}
float denominator = 0.0f;
for (uint ordinal = 0; ordinal < tokens; ordinal++) {
probabilities[ordinal] = exp(probabilities[ordinal] - max_score);
denominator += probabilities[ordinal];
}
for (uint ordinal = 0; ordinal < tokens; ordinal++) {
probabilities[ordinal] /= denominator;
}
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float value = 0.0f;
for (uint ordinal = 0; ordinal < tokens; ordinal++) {
const uint token = qwen_qsa_token(selected, ordinal, args.u[4], args.u[5], args.u[6]);
const ulong base = (ulong)token * kv_heads * dim * 2u + (ulong)kv_head * dim;
value = fma(probabilities[ordinal],
qwen_bf16(cache[base + kv_heads * dim + column]),
value);
}
out[(ulong)head * dim + column] = value;
}
kernel void kernel_qwen_dense_attention(
constant qwen_kernel_args &args [[buffer(0)]],
device float *out [[buffer(1)]],
device const float *query [[buffer(2)]],
device const ushort *cache [[buffer(3)]],
uint head [[thread_position_in_grid]]) {
uint2 gid [[thread_position_in_grid]],
uint lane [[thread_index_in_threadgroup]]) {
const uint heads = args.u[0];
const uint kv_heads = args.u[1];
const uint dim = args.u[2];
const uint tokens = args.u[3];
if (head >= heads) return;
const uint column = gid.x;
const uint head = gid.y;
if (column >= dim || head >= heads) return;
const uint kv_head = head / (heads / kv_heads);
float max_score = -INFINITY;
for (uint token = 0; token < tokens; token++) {
const ulong base = (ulong)token * kv_heads * dim * 2u + (ulong)kv_head * dim;
float score = 0.0f;
for (uint i = 0; i < dim; i++) score = fma(query[(ulong)head * dim + i], qwen_bf16(cache[base + i]), score);
max_score = max(max_score, score * rsqrt((float)dim));
}
float denominator = 0.0f;
for (uint token = 0; token < tokens; token++) {
const ulong base = (ulong)token * kv_heads * dim * 2u + (ulong)kv_head * dim;
float score = 0.0f;
for (uint i = 0; i < dim; i++) score = fma(query[(ulong)head * dim + i], qwen_bf16(cache[base + i]), score);
denominator += exp(score * rsqrt((float)dim) - max_score);
}
for (uint column = 0; column < dim; column++) {
float value = 0.0f;
const float attention_scale = rsqrt((float)dim);
threadgroup float probabilities[2048];
if (lane == 0u) {
float max_score = -INFINITY;
for (uint token = 0; token < tokens; token++) {
const ulong base = (ulong)token * kv_heads * dim * 2u + (ulong)kv_head * dim;
float score = 0.0f;
for (uint i = 0; i < dim; i++) score = fma(query[(ulong)head * dim + i], qwen_bf16(cache[base + i]), score);
const float probability = exp(score * rsqrt((float)dim) - max_score) / denominator;
value = fma(probability, qwen_bf16(cache[base + kv_heads * dim + column]), value);
probabilities[token] = score * attention_scale;
max_score = max(max_score, probabilities[token]);
}
out[(ulong)head * dim + column] = value;
float denominator = 0.0f;
for (uint token = 0; token < tokens; token++) {
probabilities[token] = exp(probabilities[token] - max_score);
denominator += probabilities[token];
}
for (uint token = 0; token < tokens; token++) probabilities[token] /= denominator;
}
threadgroup_barrier(mem_flags::mem_threadgroup);
float value = 0.0f;
for (uint token = 0; token < tokens; token++) {
const ulong base = (ulong)token * kv_heads * dim * 2u + (ulong)kv_head * dim;
value = fma(probabilities[token], qwen_bf16(cache[base + kv_heads * dim + column]), value);
}
out[(ulong)head * dim + column] = value;
}
kernel void kernel_qwen_gate_attention(