Implement exact Qwen sparse attention

This commit is contained in:
Georg Bauer
2026-09-03 21:36:33 +02:00
parent 87ccf67d0c
commit dda70c3793
3 changed files with 909 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(

View File

@@ -19,6 +19,13 @@ const ATTN_KV_HEADS: u32 = 2;
const ATTN_DIM: u32 = 256;
const ATTN_WIDTH: u32 = ATTN_HEADS * ATTN_DIM;
const ATTN_KV_WIDTH: u32 = ATTN_KV_HEADS * ATTN_DIM;
const QSA_HEADS: u32 = 4;
const QSA_KV_HEADS: u32 = 1;
const QSA_DIM: u32 = 128;
const QSA_WIDTH: u32 = (QSA_HEADS + QSA_KV_HEADS) * QSA_DIM;
const QSA_RATIO: u32 = 4;
const QSA_TOP_K: u32 = 512;
const MTP_TOKEN_RESERVE: u32 = 3;
const EXPERTS: u32 = 512;
const EXPERTS_USED: usize = 10;
const EXPERT_WIDTH: u32 = 640;
@@ -31,7 +38,7 @@ const PLE_CONV_STATE: u32 = 9;
const EOS_TOKEN: i32 = 248_044;
const PLE_ROW_BYTES: usize = 100;
const CHECKPOINT_MAGIC: &[u8; 8] = b"DS4QWN01";
const CHECKPOINT_VERSION: u32 = 2;
const CHECKPOINT_VERSION: u32 = 3;
const CHECKPOINT_CHUNK: usize = 8 * 1024 * 1024;
#[derive(Clone, Copy)]
@@ -49,8 +56,15 @@ struct Affine<'a> {
}
enum LayerState {
Gdn { conv: Buffer, recurrent: Buffer },
Attention { kv: Buffer },
Gdn {
conv: Buffer,
recurrent: Buffer,
},
Attention {
kv: Buffer,
qsa_raw: Buffer,
qsa_pooled: Buffer,
},
}
struct Scratch {
@@ -81,6 +95,10 @@ struct Scratch {
k_rope: Buffer,
v: Buffer,
attention: Buffer,
qsa_qk: Buffer,
qsa_q: Buffer,
qsa_scores: Buffer,
qsa_selected: Buffer,
ple_packed: Buffer,
ple_scales: Buffer,
ple_biases: Buffer,
@@ -94,7 +112,8 @@ struct Scratch {
}
impl Scratch {
fn new() -> Result<Self, String> {
fn new(context: u32) -> Result<Self, String> {
let qsa_blocks = qsa_block_capacity(context)?;
Ok(Self {
hidden: Buffer::floats(HIDDEN.into())?,
hc: Buffer::floats(HC_WIDTH.into())?,
@@ -123,6 +142,10 @@ impl Scratch {
k_rope: Buffer::floats(ATTN_KV_WIDTH.into())?,
v: Buffer::floats(ATTN_KV_WIDTH.into())?,
attention: Buffer::floats(ATTN_WIDTH.into())?,
qsa_qk: Buffer::floats(QSA_WIDTH.into())?,
qsa_q: Buffer::floats((QSA_HEADS * QSA_DIM).into())?,
qsa_scores: Buffer::floats(qsa_blocks.into())?,
qsa_selected: Buffer::bytes(u64::from(QSA_TOP_K) * 4)?,
ple_packed: Buffer::bytes((PLE_HEADS as u64) * 80)?,
ple_scales: Buffer::bytes((PLE_HEADS as u64) * 10)?,
ple_biases: Buffer::bytes((PLE_HEADS as u64) * 10)?,
@@ -181,7 +204,7 @@ impl QwenExecutor {
states,
ple_contract,
ple_state: allocate_ple_state()?,
scratch: Scratch::new()?,
scratch: Scratch::new(context)?,
logits: vec![0.0; VOCAB as usize],
tokens: Vec::new(),
position: 0,
@@ -201,13 +224,6 @@ impl QwenExecutor {
self.context
));
}
if self.position + 1 > DENSE_BUDGET {
return Err(
"Qwen sparse QSA selection is required beyond 2048 tokens; native QSA belongs to issue #97"
.into(),
);
}
self.begin_token(token)?;
for layer in 0..LAYERS {
@@ -688,15 +704,69 @@ impl QwenExecutor {
}
fn attention(&self, prefix: &str, layer: usize) -> Result<(), String> {
if self.position + 1 > DENSE_BUDGET {
return Err(
"Qwen sparse QSA selection is required beyond 2048 tokens; native QSA belongs to issue #97"
.into(),
);
}
let LayerState::Attention { kv } = &self.states[layer] else {
let LayerState::Attention {
kv,
qsa_raw,
qsa_pooled,
} = &self.states[layer]
else {
return Err("Qwen attention graph received GDN state".into());
};
self.affine_mv_into(
&self.affine(
&format!("{prefix}.self_attn.indexer.index_qk_proj"),
HIDDEN,
QSA_WIDTH,
None,
)?,
&self.scratch.block,
&self.scratch.qsa_qk,
HIDDEN,
QSA_WIDTH,
)?;
let mut index = args();
index.u[0] = QSA_DIM;
index.u[1] = QSA_HEADS;
index.u[2] = self.position;
self.dispatch(
c"kernel_qwen_qsa_store_raw",
qsa_raw,
Some(&self.scratch.qsa_qk),
Some(&self.scratch.qsa_q),
None,
&[],
&index,
QSA_WIDTH,
1,
)?;
self.indexer_norm_rope(
&self.scratch.qsa_q,
&self.scratch.qsa_qk,
self.weight(&format!("{prefix}.self_attn.indexer.q_layernorm.weight"))?,
QSA_HEADS,
self.position,
)?;
if (self.position + 1).is_multiple_of(QSA_RATIO) {
let mut pool = args();
pool.u[0] = QSA_DIM;
pool.u[1] = self.position / QSA_RATIO;
pool.u[2] = QSA_RATIO;
pool.u[3] = self.position + 1 - QSA_RATIO;
pool.f[0] = 1.0e-6;
pool.f[1] = 10_000_000.0;
self.dispatch(
c"kernel_qwen_qsa_pool_key",
qsa_pooled,
Some(qsa_raw),
None,
None,
&[self
.view(self.weight(&format!("{prefix}.self_attn.indexer.k_layernorm.weight"))?)],
&pool,
QSA_DIM,
1,
)?;
}
self.affine_mv_into(
&self.affine(
&format!("{prefix}.self_attn.q_proj"),
@@ -773,22 +843,7 @@ impl QwenExecutor {
ATTN_KV_WIDTH,
1,
)?;
let mut dense = args();
dense.u[0] = ATTN_HEADS;
dense.u[1] = ATTN_KV_HEADS;
dense.u[2] = ATTN_DIM;
dense.u[3] = self.position + 1;
self.dispatch(
c"kernel_qwen_dense_attention",
&self.scratch.q,
Some(&self.scratch.attention),
Some(kv),
None,
&[],
&dense,
ATTN_HEADS,
1,
)?;
self.attend(kv, qsa_pooled)?;
let mut gate = args();
gate.u[0] = ATTN_WIDTH;
self.dispatch(
@@ -816,6 +871,114 @@ impl QwenExecutor {
)
}
fn indexer_norm_rope(
&self,
input: &Buffer,
output: &Buffer,
weight: Weight<'_>,
heads: u32,
position: u32,
) -> Result<(), String> {
let mut values = args();
values.u[0] = QSA_DIM;
values.u[1] = 64;
values.u[2] = heads;
values.u[3] = position;
values.f[0] = 1.0e-6;
values.f[1] = 10_000_000.0;
self.dispatch(
c"kernel_qwen_head_norm_rope",
output,
Some(input),
None,
None,
&[self.view(weight)],
&values,
QSA_DIM,
heads,
)
}
fn attend(&self, kv: &Buffer, qsa_pooled: &Buffer) -> Result<(), String> {
let tokens = self.position + 1;
let mut values = args();
values.u[0] = ATTN_HEADS;
values.u[1] = ATTN_KV_HEADS;
values.u[2] = ATTN_DIM;
values.u[3] = tokens;
if tokens <= DENSE_BUDGET {
return self.dispatch(
c"kernel_qwen_dense_attention",
&self.scratch.q,
Some(&self.scratch.attention),
Some(kv),
None,
&[],
&values,
ATTN_DIM,
ATTN_HEADS,
);
}
let complete_blocks = tokens / QSA_RATIO;
let mut score = args();
score.u[0] = QSA_DIM;
score.u[1] = QSA_HEADS;
score.u[2] = complete_blocks;
score.f[0] = (QSA_DIM as f32).sqrt().recip();
self.dispatch(
c"kernel_qwen_qsa_scores",
&self.scratch.qsa_scores,
Some(&self.scratch.qsa_qk),
Some(qsa_pooled),
None,
&[],
&score,
complete_blocks,
1,
)?;
call(
unsafe {
ds4_gpu_indexer_topk_tensor(
self.scratch.qsa_selected.raw(),
self.scratch.qsa_scores.raw(),
complete_blocks,
1,
QSA_TOP_K,
)
},
"selecting Qwen QSA blocks",
)?;
let mut order = args();
order.u[0] = QSA_TOP_K;
self.dispatch(
c"kernel_qwen_qsa_sort_blocks",
&self.scratch.qsa_selected,
None,
None,
None,
&[],
&order,
1,
1,
)?;
values.u[4] = QSA_TOP_K;
values.u[5] = QSA_RATIO;
values.u[6] = complete_blocks * QSA_RATIO;
values.u[7] = tokens - complete_blocks * QSA_RATIO;
self.dispatch(
c"kernel_qwen_sparse_attention",
&self.scratch.q,
Some(&self.scratch.attention),
Some(kv),
Some(&self.scratch.qsa_selected),
&[],
&values,
ATTN_DIM,
ATTN_HEADS,
)
}
fn head_norm_rope(
&self,
input: &Buffer,
@@ -1364,15 +1527,16 @@ impl QwenExecutor {
progress,
)?;
}
LayerState::Attention { kv } => {
write_buffer(
&mut file,
kv,
0,
u64::from(self.position) * ATTN_KV_WIDTH as u64 * 4,
&mut chunk,
progress,
)?;
LayerState::Attention {
kv,
qsa_raw,
qsa_pooled,
} => {
let [kv_bytes, raw_bytes, pooled_bytes] =
attention_checkpoint_bytes(self.position);
write_buffer(&mut file, kv, 0, kv_bytes, &mut chunk, progress)?;
write_buffer(&mut file, qsa_raw, 0, raw_bytes, &mut chunk, progress)?;
write_buffer(&mut file, qsa_pooled, 0, pooled_bytes, &mut chunk, progress)?;
}
}
}
@@ -1469,14 +1633,16 @@ impl QwenExecutor {
progress,
)?;
}
LayerState::Attention { kv } => read_buffer(
&mut file,
LayerState::Attention {
kv,
0,
u64::from(position) * ATTN_KV_WIDTH as u64 * 4,
&mut chunk,
progress,
)?,
qsa_raw,
qsa_pooled,
} => {
let [kv_bytes, raw_bytes, pooled_bytes] = attention_checkpoint_bytes(position);
read_buffer(&mut file, kv, 0, kv_bytes, &mut chunk, progress)?;
read_buffer(&mut file, qsa_raw, 0, raw_bytes, &mut chunk, progress)?;
read_buffer(&mut file, qsa_pooled, 0, pooled_bytes, &mut chunk, progress)?;
}
}
}
let mut trailing = [0];
@@ -1498,11 +1664,17 @@ impl QwenExecutor {
}
fn allocate_states(context: u32) -> Result<Vec<LayerState>, String> {
let token_capacity = context
.checked_add(MTP_TOKEN_RESERVE)
.ok_or_else(|| "Qwen attention capacity overflows".to_owned())?;
let block_capacity = qsa_block_capacity(context)?;
(0..LAYERS)
.map(|layer| {
if layer % 4 == 3 {
Ok(LayerState::Attention {
kv: Buffer::bytes(u64::from(context) * ATTN_KV_WIDTH as u64 * 4)?,
kv: Buffer::bytes(u64::from(token_capacity) * ATTN_KV_WIDTH as u64 * 4)?,
qsa_raw: Buffer::bytes(u64::from(token_capacity) * QSA_DIM as u64 * 2)?,
qsa_pooled: Buffer::bytes(u64::from(block_capacity) * QSA_DIM as u64 * 2)?,
})
} else {
let conv = Buffer::bytes(u64::from(GDN_QKV) * 3 * 2)?;
@@ -1519,6 +1691,46 @@ fn allocate_states(context: u32) -> Result<Vec<LayerState>, String> {
.collect()
}
fn qsa_block_capacity(context: u32) -> Result<u32, String> {
context
.checked_add(MTP_TOKEN_RESERVE)
.and_then(|tokens| tokens.checked_add(QSA_RATIO - 1))
.map(|tokens| tokens / QSA_RATIO)
.ok_or_else(|| "Qwen QSA block capacity overflows".to_owned())
}
fn attention_checkpoint_bytes(position: u32) -> [u64; 3] {
[
u64::from(position) * ATTN_KV_WIDTH as u64 * 4,
u64::from(position) * QSA_DIM as u64 * 2,
u64::from(position / QSA_RATIO) * QSA_DIM as u64 * 2,
]
}
#[cfg(test)]
fn reference_qsa_selection(scores: &[f32]) -> Vec<u32> {
let mut blocks = (0..scores.len() as u32).collect::<Vec<_>>();
blocks.sort_by(|&left, &right| {
scores[right as usize]
.total_cmp(&scores[left as usize])
.then(left.cmp(&right))
});
blocks.truncate(QSA_TOP_K.min(blocks.len() as u32) as usize);
blocks
}
#[cfg(test)]
fn reference_qsa_tokens(selected: &[u32], tokens: u32) -> Vec<u32> {
let complete = tokens / QSA_RATIO;
let mut visible = selected
.iter()
.flat_map(|&block| block * QSA_RATIO..(block + 1) * QSA_RATIO)
.collect::<Vec<_>>();
visible.extend(complete * QSA_RATIO..tokens);
visible.sort_unstable();
visible
}
fn allocate_ple_state() -> Result<PleState, String> {
let conv = Buffer::bytes(u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2)?;
conv.fill(0.0, u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 / 2)?;
@@ -1857,6 +2069,40 @@ mod tests {
);
}
#[test]
fn qsa_selection_and_capacity_cover_the_full_native_window() {
assert_eq!((QSA_HEADS, QSA_KV_HEADS, QSA_DIM), (4, 1, 128));
assert_eq!(qsa_block_capacity(1).unwrap(), 1);
assert_eq!(qsa_block_capacity(2_048).unwrap(), 513);
assert_eq!(qsa_block_capacity(262_144).unwrap(), 65_537);
let mut scores = (0..600).map(|index| index as f32).collect::<Vec<_>>();
scores[10] = 10_000.0;
scores[11] = 10_000.0;
let selected = reference_qsa_selection(&scores);
assert_eq!(selected.len(), QSA_TOP_K as usize);
assert_eq!(&selected[..2], &[10, 11]);
let visible = reference_qsa_tokens(&selected, 2_403);
assert_eq!(&visible[visible.len() - 3..], &[2_400, 2_401, 2_402]);
assert_eq!(visible.len(), QSA_TOP_K as usize * QSA_RATIO as usize + 3);
assert!(visible.windows(2).all(|pair| pair[0] < pair[1]));
for depth in [2_048, 16_384, 65_536, 131_072, 262_144] {
let complete = depth / QSA_RATIO;
let tail = depth % QSA_RATIO;
assert_eq!(complete * QSA_RATIO + tail, depth);
assert!(complete <= qsa_block_capacity(depth).unwrap());
assert_eq!(
attention_checkpoint_bytes(depth),
[
u64::from(depth) * ATTN_KV_WIDTH as u64 * 4,
u64::from(depth) * QSA_DIM as u64 * 2,
u64::from(complete) * QSA_DIM as u64 * 2,
]
);
}
}
#[test]
#[ignore = "requires Apple Metal"]
fn qwen_metal_primitives_match_reference_vectors() {
@@ -1881,6 +2127,8 @@ mod tests {
] {
bytes.extend_from_slice(&bf16(value).to_le_bytes());
}
let qsa_weight_offset = bytes.len() as u64;
bytes.extend_from_slice(&bf16(0.0).to_le_bytes().repeat(QSA_DIM as usize));
fs::write(&path, bytes).unwrap();
let file = File::open(&path).unwrap();
// SAFETY: this test owns the read-only file for the lifetime of the mapping.
@@ -2237,7 +2485,7 @@ mod tests {
None,
&[],
&dense,
1,
2,
1,
)
.unwrap();
@@ -2254,6 +2502,187 @@ mod tests {
probability * 3.0 + (1.0 - probability) * 7.0,
);
let qsa_projected = Buffer::floats(QSA_WIDTH.into()).unwrap();
let qsa_query = Buffer::floats((QSA_HEADS * QSA_DIM).into()).unwrap();
let qsa_raw = Buffer::bytes(u64::from(QSA_RATIO * QSA_DIM) * 2).unwrap();
let qsa_pooled = Buffer::bytes(u64::from(QSA_DIM) * 2).unwrap();
let mut projected = vec![0.0; QSA_WIDTH as usize];
for head in 0..QSA_HEADS {
projected[(head * QSA_DIM) as usize] = 1.0;
}
let mut qsa_store = args();
qsa_store.u[0] = QSA_DIM;
qsa_store.u[1] = QSA_HEADS;
for position in 0..QSA_RATIO {
projected[(QSA_HEADS * QSA_DIM) as usize] = position as f32 + 1.0;
qsa_projected.write_f32(&projected).unwrap();
qsa_store.u[2] = position;
dispatch_qwen(
c"kernel_qwen_qsa_store_raw",
&qsa_raw,
Some(&qsa_projected),
Some(&qsa_query),
None,
&[],
&qsa_store,
QSA_WIDTH,
1,
)
.unwrap();
}
let mut qsa_pool = args();
qsa_pool.u[0] = QSA_DIM;
qsa_pool.u[1] = 0;
qsa_pool.u[2] = QSA_RATIO;
qsa_pool.u[3] = 0;
qsa_pool.f[0] = 1.0e-6;
qsa_pool.f[1] = 10_000_000.0;
dispatch_qwen(
c"kernel_qwen_qsa_pool_key",
&qsa_pooled,
Some(&qsa_raw),
None,
None,
&[view(qsa_weight_offset, u64::from(QSA_DIM) * 2)],
&qsa_pool,
QSA_DIM,
1,
)
.unwrap();
let qsa_scores = Buffer::floats(1).unwrap();
let mut qsa_score = args();
qsa_score.u[0] = QSA_DIM;
qsa_score.u[1] = QSA_HEADS;
qsa_score.u[2] = 1;
qsa_score.f[0] = (QSA_DIM as f32).sqrt().recip();
dispatch_qwen(
c"kernel_qwen_qsa_scores",
&qsa_scores,
Some(&qsa_query),
Some(&qsa_pooled),
None,
&[],
&qsa_score,
1,
1,
)
.unwrap();
qsa_scores.read_f32(&mut scalar).unwrap();
close(scalar[0], 4.0);
let tie_scores = Buffer::floats(4).unwrap();
tie_scores.write_f32(&[5.0, 5.0, 4.0, 3.0]).unwrap();
let tie_selected = Buffer::bytes(8).unwrap();
call(
unsafe { ds4_gpu_indexer_topk_tensor(tie_selected.raw(), tie_scores.raw(), 4, 1, 2) },
"testing deterministic Qwen QSA ties",
)
.unwrap();
let mut tie_ids = [0; 2];
tie_selected.read_i32(&mut tie_ids).unwrap();
assert_eq!(tie_ids, [0, 1]);
let selected_blocks = Buffer::bytes(12).unwrap();
selected_blocks.write_i32(&[3, 1, 2]).unwrap();
let mut sort = args();
sort.u[0] = 3;
dispatch_qwen(
c"kernel_qwen_qsa_sort_blocks",
&selected_blocks,
None,
None,
None,
&[],
&sort,
1,
1,
)
.unwrap();
let mut sorted = [0; 3];
selected_blocks.read_i32(&mut sorted).unwrap();
assert_eq!(sorted, [1, 2, 3]);
let sparse_cache = Buffer::bytes(9 * 2 * 2).unwrap();
for token in 0..9 {
key.write_f32(&[0.0, 0.0]).unwrap();
value
.write_f32(&[token as f32 + 1.0, token as f32 + 11.0])
.unwrap();
store.u[1] = token;
dispatch_qwen(
c"kernel_qwen_store_kv_bf16",
&sparse_cache,
Some(&key),
Some(&value),
None,
&[],
&store,
2,
1,
)
.unwrap();
}
query.write_f32(&[0.0, 0.0]).unwrap();
selected_blocks.write_i32(&[1]).unwrap();
let mut sparse = args();
sparse.u[0] = 1;
sparse.u[1] = 1;
sparse.u[2] = 2;
sparse.u[3] = 9;
sparse.u[4] = 1;
sparse.u[5] = 4;
sparse.u[6] = 8;
sparse.u[7] = 1;
dispatch_qwen(
c"kernel_qwen_sparse_attention",
&attention,
Some(&query),
Some(&sparse_cache),
Some(&selected_blocks),
&[],
&sparse,
2,
1,
)
.unwrap();
attention.read_f32(&mut actual_attention).unwrap();
close(actual_attention[0], 7.0);
close(actual_attention[1], 17.0);
selected_blocks.write_i32(&[0, 1]).unwrap();
sparse.u[4] = 2;
dispatch_qwen(
c"kernel_qwen_sparse_attention",
&attention,
Some(&query),
Some(&sparse_cache),
Some(&selected_blocks),
&[],
&sparse,
2,
1,
)
.unwrap();
let dense_all = Buffer::floats(2).unwrap();
dense.u[3] = 9;
dispatch_qwen(
c"kernel_qwen_dense_attention",
&dense_all,
Some(&query),
Some(&sparse_cache),
None,
&[],
&dense,
2,
1,
)
.unwrap();
let mut sparse_all = [0.0; 2];
let mut dense_values = [0.0; 2];
attention.read_f32(&mut sparse_all).unwrap();
dense_all.read_f32(&mut dense_values).unwrap();
assert_eq!(sparse_all, dense_values);
let ple_packed = Buffer::bytes(80).unwrap();
ple_packed.write(0, &[0x33; 80]).unwrap();
let ple_scales = Buffer::bytes(10).unwrap();
@@ -2347,6 +2776,182 @@ mod tests {
fs::remove_file(path).unwrap();
}
#[test]
#[ignore = "requires Apple Metal and allocates the full QSA score window"]
fn qwen_qsa_selection_and_attention_remain_bounded_at_native_depths() {
use std::time::Instant;
configure_sources().unwrap();
let _context = Context::open_qwen(0).unwrap();
let max_blocks = 262_144 / QSA_RATIO;
let query = Buffer::floats((QSA_HEADS * QSA_DIM).into()).unwrap();
query
.write_f32(&vec![0.0; (QSA_HEADS * QSA_DIM) as usize])
.unwrap();
let pooled = Buffer::bytes(u64::from(max_blocks * QSA_DIM) * 2).unwrap();
pooled
.write(0, &vec![0; (max_blocks * QSA_DIM * 2) as usize])
.unwrap();
let scores = Buffer::floats(max_blocks.into()).unwrap();
let selected = Buffer::bytes(u64::from(QSA_TOP_K) * 4).unwrap();
for depth in [2_048, 16_384, 65_536, 131_072, 262_144] {
let blocks = depth / QSA_RATIO;
let started = Instant::now();
let mut score = args();
score.u[0] = QSA_DIM;
score.u[1] = QSA_HEADS;
score.u[2] = blocks;
score.f[0] = (QSA_DIM as f32).sqrt().recip();
dispatch_qwen(
c"kernel_qwen_qsa_scores",
&scores,
Some(&query),
Some(&pooled),
None,
&[],
&score,
blocks,
1,
)
.unwrap();
call(
unsafe {
ds4_gpu_indexer_topk_tensor(selected.raw(), scores.raw(), blocks, 1, QSA_TOP_K)
},
"testing full-window Qwen QSA selection",
)
.unwrap();
let mut order = args();
order.u[0] = QSA_TOP_K;
dispatch_qwen(
c"kernel_qwen_qsa_sort_blocks",
&selected,
None,
None,
None,
&[],
&order,
1,
1,
)
.unwrap();
let mut ids = vec![0; QSA_TOP_K as usize];
selected.read_i32(&mut ids).unwrap();
assert_eq!(ids, (0..QSA_TOP_K as i32).collect::<Vec<_>>());
eprintln!(
"Qwen QSA selection depth {depth}: {:.3} ms",
started.elapsed().as_secs_f64() * 1_000.0
);
}
let depth = 262_143_u32;
let tail_start = depth / QSA_RATIO * QSA_RATIO;
selected
.write_i32(&(0..QSA_TOP_K as i32).collect::<Vec<_>>())
.unwrap();
let cache = Buffer::bytes(u64::from(depth) * ATTN_KV_WIDTH as u64 * 4).unwrap();
cache
.write(
0,
&vec![0; (QSA_TOP_K * QSA_RATIO * ATTN_KV_WIDTH * 4) as usize],
)
.unwrap();
cache
.write(
u64::from(tail_start) * ATTN_KV_WIDTH as u64 * 4,
&vec![0; ((depth - tail_start) * ATTN_KV_WIDTH * 4) as usize],
)
.unwrap();
let main_query = Buffer::floats(ATTN_WIDTH.into()).unwrap();
main_query
.write_f32(&vec![0.0; ATTN_WIDTH as usize])
.unwrap();
let output = Buffer::floats(ATTN_WIDTH.into()).unwrap();
let mut sparse = args();
sparse.u[0] = ATTN_HEADS;
sparse.u[1] = ATTN_KV_HEADS;
sparse.u[2] = ATTN_DIM;
sparse.u[3] = depth;
sparse.u[4] = QSA_TOP_K;
sparse.u[5] = QSA_RATIO;
sparse.u[6] = tail_start;
sparse.u[7] = depth - tail_start;
let started = Instant::now();
dispatch_qwen(
c"kernel_qwen_sparse_attention",
&output,
Some(&main_query),
Some(&cache),
Some(&selected),
&[],
&sparse,
ATTN_DIM,
ATTN_HEADS,
)
.unwrap();
let mut values = vec![1.0; ATTN_WIDTH as usize];
output.read_f32(&mut values).unwrap();
assert!(values.into_iter().all(|value| value == 0.0));
eprintln!(
"Qwen sparse attention depth {depth}, {} visible tokens: {:.3} ms",
QSA_TOP_K * QSA_RATIO + depth - tail_start,
started.elapsed().as_secs_f64() * 1_000.0
);
}
#[test]
#[ignore = "requires the pinned 105 GB Qwen artifact set and Apple Metal"]
fn qwen_qsa_full_context_preallocation_succeeds() {
configure_sources().unwrap();
let root = std::env::var_os("DS4SERVER_QWEN38_SOURCE")
.map(PathBuf::from)
.expect("set DS4SERVER_QWEN38_SOURCE to the pinned artifact directory");
let model = QwenModel::open(&root, 262_144).unwrap();
assert_eq!(model.memory().kv_and_recurrent, 7_564_812_288);
let executor = QwenExecutor::open(model, 262_144).unwrap();
assert_eq!(executor.context(), 262_144);
assert_eq!(qsa_block_capacity(executor.context()).unwrap(), 65_537);
}
#[test]
#[ignore = "requires the pinned 105 GB Qwen artifact set and Apple Metal"]
fn qwen_qsa_real_layer_executes_the_first_sparse_token() {
configure_sources().unwrap();
let root = std::env::var_os("DS4SERVER_QWEN38_SOURCE")
.map(PathBuf::from)
.expect("set DS4SERVER_QWEN38_SOURCE to the pinned artifact directory");
let model = QwenModel::open(&root, DENSE_BUDGET + 1).unwrap();
let mut executor = QwenExecutor::open(model, DENSE_BUDGET + 1).unwrap();
executor.position = DENSE_BUDGET;
executor
.scratch
.block
.write_f32(&vec![0.0; HIDDEN as usize])
.unwrap();
let LayerState::Attention { kv, qsa_pooled, .. } = &executor.states[3] else {
unreachable!()
};
kv.write(0, &vec![0; (DENSE_BUDGET * ATTN_KV_WIDTH * 4) as usize])
.unwrap();
qsa_pooled
.write(0, &vec![0; (QSA_TOP_K * QSA_DIM * 2) as usize])
.unwrap();
executor
.attention("language_model.model.layers.3", 3)
.unwrap();
let mut hidden = vec![0.0; HIDDEN as usize];
executor.scratch.hidden.read_f32(&mut hidden).unwrap();
assert!(hidden.iter().all(|value| value.is_finite()));
let mut selected = vec![0; QSA_TOP_K as usize];
executor
.scratch
.qsa_selected
.read_i32(&mut selected)
.unwrap();
assert_eq!(selected, (0..QSA_TOP_K as i32).collect::<Vec<_>>());
}
#[test]
#[ignore = "requires the pinned 105 GB Qwen artifact set and Apple Metal"]
fn qwen_core_boundary_and_checkpoint_are_stable() {
@@ -2391,6 +2996,11 @@ mod tests {
.max_by(|a, b| a.1.total_cmp(b.1))
.unwrap()
.0 as i32;
let LayerState::Attention { qsa_raw, .. } = &executor.states[3] else {
unreachable!()
};
let mut saved_qsa_raw = vec![0; QSA_DIM as usize * 2];
qsa_raw.read(0, &mut saved_qsa_raw).unwrap();
for (actual, expected) in
reference
.into_iter()
@@ -2431,6 +3041,12 @@ mod tests {
assert_eq!(executor.tokens, [1]);
assert_eq!(executor.ple_state.history, [EOS_TOKEN, 1]);
assert_eq!(executor.checkpoint_tag, [7; 32]);
let LayerState::Attention { qsa_raw, .. } = &executor.states[3] else {
unreachable!()
};
let mut restored_qsa_raw = vec![0; saved_qsa_raw.len()];
qsa_raw.read(0, &mut restored_qsa_raw).unwrap();
assert_eq!(restored_qsa_raw, saved_qsa_raw);
assert_eq!(
[
executor.logits[0],
@@ -2488,10 +3104,34 @@ mod tests {
executor.eval(EOS_TOKEN).unwrap();
assert_eq!(executor.ple_state.history, [EOS_TOKEN; PLE_HISTORY]);
executor.position = DENSE_BUDGET;
executor.context = DENSE_BUDGET + 1;
let error = executor.eval(1).unwrap_err();
assert!(error.contains("issue #97"));
let LayerState::Attention {
qsa_raw,
qsa_pooled,
..
} = &executor.states[3]
else {
unreachable!()
};
let mut raw_key = vec![0; QSA_DIM as usize * 2];
qsa_raw.read(0, &mut raw_key).unwrap();
assert!(raw_key.iter().any(|&byte| byte != 0));
assert!(!qsa_pooled.raw().is_null());
let pooled_pattern = vec![0x5a; QSA_DIM as usize * 2];
qsa_pooled.write(0, &pooled_pattern).unwrap();
executor.position = 4;
executor.tokens = vec![EOS_TOKEN; 4];
executor
.save_checkpoint(&checkpoint, [8; 32], &mut |_| {})
.unwrap();
assert!(executor.load_checkpoint(&checkpoint, &mut |_| {}).unwrap());
let LayerState::Attention { qsa_pooled, .. } = &executor.states[3] else {
unreachable!()
};
let mut restored_pool = vec![0; pooled_pattern.len()];
qsa_pooled.read(0, &mut restored_pool).unwrap();
assert_eq!(restored_pool, pooled_pattern);
assert_eq!(executor.position, 4);
assert_eq!(executor.checkpoint_tag, [8; 32]);
fs::remove_file(checkpoint).unwrap();
}
}

View File

@@ -21,6 +21,12 @@ const CORE_BYTES: u64 = 71_742_682_599;
const PLE_BYTES: u64 = 32_000_154_008;
const MTP_BYTES: u64 = 1_672_575_532;
const KV_BYTES_PER_TOKEN: u64 = 24_576;
const QSA_RAW_BYTES_PER_TOKEN: u64 = 3_072;
const QSA_POOLED_BYTES_PER_BLOCK: u64 = 3_072;
const QSA_POOL_RATIO: u64 = 4;
const QSA_FIXED_SCRATCH_BYTES: u64 = 6_656;
const QSA_TOPK_SCRATCH_BYTES_PER_BLOCK: u64 = 8;
const MTP_TOKEN_RESERVE: u64 = 3;
const GDN_STATE_BYTES: u64 = 113_246_208;
const GDN_CONV_BYTES: u64 = 2_211_840;
const PLE_CONV_BYTES: u64 = 184_320;
@@ -634,14 +640,31 @@ pub(super) fn memory_plan(
if prefill_chunk == 0 {
return Err("Qwen prefill chunk must be positive".into());
}
let kv = KV_BYTES_PER_TOKEN
.checked_mul(u64::from(context))
let token_capacity = u64::from(context)
.checked_add(MTP_TOKEN_RESERVE)
.ok_or_else(|| "Qwen attention capacity overflows".to_owned())?;
let block_capacity = token_capacity.div_ceil(QSA_POOL_RATIO);
let topk_scratch = if context > 2_048 {
u64::from(context) / QSA_POOL_RATIO * QSA_TOPK_SCRATCH_BYTES_PER_BLOCK
} else {
0
};
let kv = (KV_BYTES_PER_TOKEN + QSA_RAW_BYTES_PER_TOKEN)
.checked_mul(token_capacity)
.and_then(|bytes| {
QSA_POOLED_BYTES_PER_BLOCK
.checked_mul(block_capacity)
.and_then(|pooled| bytes.checked_add(pooled))
})
.ok_or_else(|| "Qwen KV memory size overflows".to_owned())?;
let kv_and_recurrent = kv
.checked_add(GDN_STATE_BYTES + GDN_CONV_BYTES + PLE_CONV_BYTES)
.ok_or_else(|| "Qwen recurrent memory size overflows".to_owned())?;
let prefill_transient = u64::from(prefill_chunk)
.checked_mul((4 * 2_560 + 2_048 + 2_048 + 6_144 + 6_144) * 2)
.and_then(|bytes| bytes.checked_add(block_capacity * 4))
.and_then(|bytes| bytes.checked_add(QSA_FIXED_SCRATCH_BYTES))
.and_then(|bytes| bytes.checked_add(topk_scratch))
.ok_or_else(|| "Qwen prefill transient size overflows".to_owned())?;
let admitted_mtp = if enable_mtp { MTP_BYTES } else { 0 };
let admission = CORE_BYTES
@@ -681,12 +704,12 @@ mod tests {
assert_eq!(plan.resident_core, CORE_BYTES);
assert_eq!(plan.mapped_ple, PLE_BYTES);
assert_eq!(plan.optional_mtp, MTP_BYTES);
assert_eq!(plan.kv_and_recurrent, 6_558_093_312);
assert_eq!(plan.prefill_transient, 27_262_976);
assert_eq!(plan.admission, 80_000_614_419);
assert_eq!(plan.kv_and_recurrent, 7_564_812_288);
assert_eq!(plan.prefill_transient, 28_056_068);
assert_eq!(plan.admission, 81_008_126_487);
let without_mtp = memory_plan(262_144, false, 512).unwrap();
assert_eq!(without_mtp.optional_mtp, MTP_BYTES);
assert_eq!(without_mtp.admission, 78_328_038_887);
assert_eq!(without_mtp.admission, 79_335_550_955);
assert!(memory_plan(0, false, 512).is_err());
assert!(memory_plan(262_145, false, 512).is_err());
assert!(memory_plan(1, false, 0).is_err());