use std::borrow::Cow;
use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::Node;
use rayon::prelude::*;
use super::{check_arity, to_dense_i64};
use crate::dtype::{to_dense_f32_widen, write_dense_f32_narrow};
const MIN_PARALLEL_ATTENTION_WORK: usize = 160 * 1024;
pub struct GroupQueryAttentionKernel {
num_heads: usize,
kv_num_heads: usize,
scale: Option<f32>,
do_rotary: bool,
rotary_interleaved: bool,
local_window_size: i64,
softcap: f32,
}
pub struct GroupQueryAttentionFactory;
impl KernelFactory for GroupQueryAttentionFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let required_heads = |name: &str| -> Result<usize> {
let value = node.attr(name).and_then(|a| a.as_int()).ok_or_else(|| {
EpError::KernelFailed(format!(
"GroupQueryAttention: missing required `{name}` attribute"
))
})?;
usize::try_from(value)
.ok()
.filter(|&v| v > 0)
.ok_or_else(|| {
EpError::KernelFailed(format!("GroupQueryAttention: `{name}` must be > 0"))
})
};
let num_heads = required_heads("num_heads")?;
let kv_num_heads = required_heads("kv_num_heads")?;
if !num_heads.is_multiple_of(kv_num_heads) {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: num_heads {num_heads} must be a multiple of kv_num_heads {kv_num_heads}"
)));
}
for name in ["k_quant_type", "v_quant_type"] {
if let Some(value) = node.attr(name)
&& value.as_str() != Some("NONE")
{
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: `{name}` other than NONE is not yet supported by the f32 CPU kernel"
)));
}
}
if node
.attr("kv_cache_bit_width")
.and_then(|a| a.as_int())
.unwrap_or(0)
!= 0
{
return Err(EpError::KernelFailed(
"GroupQueryAttention: quantized KV cache is not yet supported".into(),
));
}
if node.attr("qk_output").and_then(|a| a.as_int()).unwrap_or(0) != 0 {
return Err(EpError::KernelFailed(
"GroupQueryAttention: qk_output is not yet supported".into(),
));
}
if node
.attr("smooth_softmax")
.and_then(|a| a.as_int())
.unwrap_or(0)
!= 0
{
return Err(EpError::KernelFailed(
"GroupQueryAttention: smooth_softmax is not yet supported".into(),
));
}
let softcap = node
.attr("softcap")
.and_then(|a| a.as_float())
.unwrap_or(0.0);
if softcap < 0.0 {
return Err(EpError::KernelFailed(
"GroupQueryAttention: softcap must be non-negative".into(),
));
}
Ok(Box::new(GroupQueryAttentionKernel {
num_heads,
kv_num_heads,
scale: node.attr("scale").and_then(|a| a.as_float()),
do_rotary: node.attr("do_rotary").and_then(|a| a.as_int()).unwrap_or(0) != 0,
rotary_interleaved: node
.attr("rotary_interleaved")
.and_then(|a| a.as_int())
.unwrap_or(0)
!= 0,
local_window_size: node
.attr("local_window_size")
.and_then(|a| a.as_int())
.unwrap_or(-1),
softcap,
}))
}
}
struct Bhsd {
data: Vec<f32>,
batch: usize,
heads: usize,
seq: usize,
dim: usize,
}
impl Bhsd {
fn from_bsh(view: &TensorView, heads: usize, name: &str) -> Result<Self> {
if view.shape.len() != 3 {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: unpacked {name} must be rank 3 [B,S,H*D], got {:?}",
view.shape
)));
}
let (batch, seq, hidden) = (view.shape[0], view.shape[1], view.shape[2]);
if !hidden.is_multiple_of(heads) {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: {name} hidden size {hidden} is not divisible by {heads} heads"
)));
}
let dim = hidden / heads;
let src = to_dense_f32_widen("GroupQueryAttention", view)?;
let mut data = vec![0.0; src.len()];
for b in 0..batch {
for s in 0..seq {
for h in 0..heads {
for d in 0..dim {
data[((b * heads + h) * seq + s) * dim + d] =
src[((b * seq + s) * heads + h) * dim + d];
}
}
}
}
Ok(Self {
data,
batch,
heads,
seq,
dim,
})
}
fn from_packed_qkv(
view: &TensorView,
num_heads: usize,
kv_num_heads: usize,
) -> Result<(Self, Self, Self)> {
if view.shape.len() != 3 {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: packed query must be rank 3 [B,S,(N+2*Nk)*D], got {:?}",
view.shape
)));
}
let (batch, seq, hidden) = (view.shape[0], view.shape[1], view.shape[2]);
let packed_heads = num_heads + 2 * kv_num_heads;
if !hidden.is_multiple_of(packed_heads) {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: packed QKV hidden size {hidden} is not divisible by num_heads + 2*kv_num_heads ({packed_heads})"
)));
}
let dim = hidden / packed_heads;
if dim == 0 {
return Err(EpError::KernelFailed(
"GroupQueryAttention: packed QKV head size must be positive".into(),
));
}
let src = to_dense_f32_widen("GroupQueryAttention", view)?;
let q_hidden = num_heads * dim;
let kv_hidden = kv_num_heads * dim;
let mut q = vec![0.0; batch * num_heads * seq * dim];
let mut k = vec![0.0; batch * kv_num_heads * seq * dim];
let mut v = vec![0.0; k.len()];
for b in 0..batch {
for s in 0..seq {
let src_base = (b * seq + s) * hidden;
for h in 0..num_heads {
for d in 0..dim {
q[((b * num_heads + h) * seq + s) * dim + d] = src[src_base + h * dim + d];
}
}
for h in 0..kv_num_heads {
for d in 0..dim {
let dst = ((b * kv_num_heads + h) * seq + s) * dim + d;
k[dst] = src[src_base + q_hidden + h * dim + d];
v[dst] = src[src_base + q_hidden + kv_hidden + h * dim + d];
}
}
}
}
Ok((
Self {
data: q,
batch,
heads: num_heads,
seq,
dim,
},
Self {
data: k,
batch,
heads: kv_num_heads,
seq,
dim,
},
Self {
data: v,
batch,
heads: kv_num_heads,
seq,
dim,
},
))
}
#[inline]
fn at(&self, b: usize, h: usize, s: usize, d: usize) -> f32 {
self.data[((b * self.heads + h) * self.seq + s) * self.dim + d]
}
}
struct PastCache<'a> {
data: Cow<'a, [f32]>,
seq: usize,
dim: usize,
batch: usize,
}
impl<'a> PastCache<'a> {
fn from_cache(view: &'a TensorView<'a>, heads: usize, name: &str) -> Result<Self> {
if view.shape.len() != 4 || view.shape[1] != heads {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: {name} must use BNSH layout [B,{heads},S,D], got {:?}",
view.shape
)));
}
Ok(Self {
data: to_dense_f32_widen("GroupQueryAttention", view)?,
seq: view.shape[2],
dim: view.shape[3],
batch: view.shape[0],
})
}
}
fn scalar_i64(view: &TensorView, name: &str) -> Result<usize> {
let values = to_dense_i64(view)?;
if values.len() != 1 || values[0] < 0 {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: {name} must be one non-negative int32 scalar"
)));
}
Ok(values[0] as usize)
}
fn rotate(
tensor: &mut Bhsd,
cos: &[f32],
sin: &[f32],
cache_rows: usize,
positions: &[usize],
interleaved: bool,
) -> Result<()> {
if !tensor.dim.is_multiple_of(2) {
return Err(EpError::KernelFailed(
"GroupQueryAttention: do_rotary requires an even head_size".into(),
));
}
let half = tensor.dim / 2;
if cos.len() != cache_rows * half || sin.len() != cache_rows * half {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: cos_cache/sin_cache must have shape [max_sequence_length,{half}]"
)));
}
for b in 0..tensor.batch {
for s in 0..tensor.seq {
let pos = positions[b * tensor.seq + s];
if pos >= cache_rows {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: rotary position {pos} exceeds cache rows {cache_rows}"
)));
}
for h in 0..tensor.heads {
for k in 0..half {
let (d0, d1) = if interleaved {
(2 * k, 2 * k + 1)
} else {
(k, k + half)
};
let i0 = ((b * tensor.heads + h) * tensor.seq + s) * tensor.dim + d0;
let i1 = ((b * tensor.heads + h) * tensor.seq + s) * tensor.dim + d1;
let (x0, x1) = (tensor.data[i0], tensor.data[i1]);
let (c, sn) = (cos[pos * half + k], sin[pos * half + k]);
tensor.data[i0] = c * x0 - sn * x1;
tensor.data[i1] = sn * x0 + c * x1;
}
}
}
}
Ok(())
}
fn write_decode_output(out: &mut TensorMut, data: &[f32]) -> Result<()> {
if out.dtype != onnx_runtime_ir::DataType::Float32 || !out.is_contiguous() {
return write_dense_f32_narrow("GroupQueryAttention", out, data);
}
out.validate()?;
if out.numel() != data.len() {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: output element count {} does not match produced {}",
out.numel(),
data.len()
)));
}
if data.is_empty() {
return Ok(());
}
let dst = unsafe { std::slice::from_raw_parts_mut(out.data_ptr_mut::<f32>(), data.len()) };
dst.copy_from_slice(data);
Ok(())
}
impl Kernel for GroupQueryAttentionKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("GroupQueryAttention", inputs, outputs, 7, 14, 1)?;
if outputs.len() > 3 {
return Err(EpError::KernelFailed(
"GroupQueryAttention: output_qk is not yet supported".into(),
));
}
let packed_qkv = inputs[1].is_absent() && inputs[2].is_absent();
if inputs[1].is_absent() != inputs[2].is_absent() {
return Err(EpError::KernelFailed(
"GroupQueryAttention: key and value must both be present for unpacked Q/K/V or both absent for packed QKV".into(),
));
}
if inputs.get(10).is_some_and(|v| !v.is_absent()) {
return Err(EpError::KernelFailed(
"GroupQueryAttention: attention_bias is not yet supported".into(),
));
}
if inputs.get(11).is_some_and(|v| !v.is_absent()) {
return Err(EpError::KernelFailed(
"GroupQueryAttention: head_sink is not yet supported".into(),
));
}
if inputs.get(12).is_some_and(|v| !v.is_absent())
|| inputs.get(13).is_some_and(|v| !v.is_absent())
{
return Err(EpError::KernelFailed(
"GroupQueryAttention: quantized-cache k_scale/v_scale inputs are not yet supported"
.into(),
));
}
if self.local_window_size == 0 || self.local_window_size < -1 {
return Err(EpError::KernelFailed(
"GroupQueryAttention: local_window_size must be -1 or a positive integer".into(),
));
}
let (mut q, mut k, v) = if packed_qkv {
Bhsd::from_packed_qkv(&inputs[0], self.num_heads, self.kv_num_heads)?
} else {
(
Bhsd::from_bsh(&inputs[0], self.num_heads, "query")?,
Bhsd::from_bsh(&inputs[1], self.kv_num_heads, "key")?,
Bhsd::from_bsh(&inputs[2], self.kv_num_heads, "value")?,
)
};
if q.batch != k.batch
|| q.batch != v.batch
|| k.seq != v.seq
|| k.dim != q.dim
|| v.dim != q.dim
{
return Err(EpError::KernelFailed(
"GroupQueryAttention: incompatible query/key/value batch, sequence, or head dimensions".into(),
));
}
let has_past_key = !inputs[3].is_absent();
let has_past_value = !inputs[4].is_absent();
if has_past_key != has_past_value {
return Err(EpError::KernelFailed(
"GroupQueryAttention: past_key and past_value must be provided together".into(),
));
}
let past_key = has_past_key
.then(|| PastCache::from_cache(&inputs[3], self.kv_num_heads, "past_key"))
.transpose()?;
let past_value = has_past_value
.then(|| PastCache::from_cache(&inputs[4], self.kv_num_heads, "past_value"))
.transpose()?;
if let (Some(pk), Some(pv)) = (&past_key, &past_value)
&& (pk.batch != q.batch
|| pv.batch != q.batch
|| pk.seq != pv.seq
|| pk.dim != q.dim
|| pv.dim != q.dim)
{
return Err(EpError::KernelFailed(
"GroupQueryAttention: past_key/past_value dimensions are incompatible with Q/K/V"
.into(),
));
}
let seqlens = to_dense_i64(&inputs[5])?;
if seqlens.len() != q.batch || seqlens.iter().any(|&x| x < 0) {
return Err(EpError::KernelFailed(
"GroupQueryAttention: seqlens_k must be non-negative int32 [batch_size]".into(),
));
}
let total_sequence_length = scalar_i64(&inputs[6], "total_sequence_length")?;
let totals: Vec<usize> = seqlens.iter().map(|&x| x as usize + 1).collect();
if totals.iter().copied().max().unwrap_or(0) != total_sequence_length {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: total_sequence_length {total_sequence_length} must equal max(seqlens_k + 1)"
)));
}
let mut past_lengths = Vec::with_capacity(q.batch);
let mut query_starts = Vec::with_capacity(q.batch);
for &total in &totals {
let past = total.checked_sub(k.seq).ok_or_else(|| {
EpError::KernelFailed(
"GroupQueryAttention: seqlens_k + 1 is shorter than current key sequence"
.into(),
)
})?;
if past > past_key.as_ref().map_or(0, |cache| cache.seq) {
return Err(EpError::KernelFailed(
"GroupQueryAttention: effective past length exceeds past cache extent".into(),
));
}
past_lengths.push(past);
query_starts.push(total.checked_sub(q.seq).ok_or_else(|| {
EpError::KernelFailed(
"GroupQueryAttention: total sequence is shorter than query sequence".into(),
)
})?);
}
if self.do_rotary {
let cos_view = inputs.get(7).filter(|v| !v.is_absent()).ok_or_else(|| {
EpError::KernelFailed("GroupQueryAttention: do_rotary=1 requires cos_cache".into())
})?;
let sin_view = inputs.get(8).filter(|v| !v.is_absent()).ok_or_else(|| {
EpError::KernelFailed("GroupQueryAttention: do_rotary=1 requires sin_cache".into())
})?;
if cos_view.shape.len() != 2 || sin_view.shape != cos_view.shape {
return Err(EpError::KernelFailed(
"GroupQueryAttention: cos_cache and sin_cache must have equal rank-2 shapes"
.into(),
));
}
if cos_view.shape[1] != q.dim / 2 {
return Err(EpError::KernelFailed(format!(
"GroupQueryAttention: cos_cache/sin_cache second dimension must be head_size/2 ({})",
q.dim / 2
)));
}
let explicit_position_ids = inputs.get(9).filter(|v| !v.is_absent());
let query_positions = if let Some(position_ids) = explicit_position_ids {
let ids = to_dense_i64(position_ids)?;
if position_ids.shape != [q.batch, q.seq] || ids.iter().any(|&x| x < 0) {
return Err(EpError::KernelFailed(
"GroupQueryAttention: position_ids must be non-negative int64 [batch_size, sequence_length]".into(),
));
}
ids.into_iter().map(|x| x as usize).collect()
} else {
let mut ids = Vec::with_capacity(q.batch * q.seq);
for &total in &totals {
let start = total.checked_sub(q.seq).ok_or_else(|| {
EpError::KernelFailed(
"GroupQueryAttention: total sequence is shorter than query sequence"
.into(),
)
})?;
ids.extend((0..q.seq).map(|s| start + s));
}
ids
};
let key_positions = if explicit_position_ids.is_some() && k.seq == q.seq {
query_positions.clone()
} else {
let mut ids = Vec::with_capacity(k.batch * k.seq);
for &total in &totals {
let start = total.checked_sub(k.seq).ok_or_else(|| {
EpError::KernelFailed(
"GroupQueryAttention: total sequence is shorter than key sequence"
.into(),
)
})?;
ids.extend((0..k.seq).map(|s| start + s));
}
ids
};
let cos = to_dense_f32_widen("GroupQueryAttention", cos_view)?;
let sin = to_dense_f32_widen("GroupQueryAttention", sin_view)?;
rotate(
&mut q,
&cos,
&sin,
cos_view.shape[0],
&query_positions,
self.rotary_interleaved,
)?;
rotate(
&mut k,
&cos,
&sin,
cos_view.shape[0],
&key_positions,
self.rotary_interleaved,
)?;
}
let cache_dim = q.dim;
let present_sequence_length = past_key.as_ref().map_or(total_sequence_length, |cache| {
cache.seq.max(total_sequence_length)
});
let mut present_k =
vec![0.0; q.batch * self.kv_num_heads * present_sequence_length * cache_dim];
let mut present_v = vec![0.0; present_k.len()];
for (b, &past_len) in past_lengths.iter().enumerate() {
for h in 0..self.kv_num_heads {
let head = b * self.kv_num_heads + h;
let dst_base = head * present_sequence_length * cache_dim;
if past_len > 0 {
let copy = past_len * cache_dim;
let pk = past_key.as_ref().unwrap();
let pv = past_value.as_ref().unwrap();
let src = head * pk.seq * cache_dim;
present_k[dst_base..dst_base + copy].copy_from_slice(&pk.data[src..src + copy]);
present_v[dst_base..dst_base + copy].copy_from_slice(&pv.data[src..src + copy]);
}
let cur = k.seq * cache_dim;
let dst_cur = dst_base + past_len * cache_dim;
let src_cur = head * k.seq * cache_dim;
present_k[dst_cur..dst_cur + cur].copy_from_slice(&k.data[src_cur..src_cur + cur]);
present_v[dst_cur..dst_cur + cur].copy_from_slice(&v.data[src_cur..src_cur + cur]);
}
}
let scale = self
.scale
.filter(|&scale| scale != 0.0)
.unwrap_or_else(|| 1.0 / (cache_dim as f32).sqrt());
let group = self.num_heads / self.kv_num_heads;
let attention_rows = q.batch * q.seq * self.num_heads;
let mut y_bhsd = vec![0.0; attention_rows * v.dim];
let compute_row = |b: usize, qh: usize, qs: usize, output_row: &mut [f32]| {
let kvh = qh / group;
let causal_limit = query_starts[b] + qs;
let local_start = if self.local_window_size > 0 {
(causal_limit + 1).saturating_sub(self.local_window_size as usize)
} else {
0
};
let mut scores = vec![f32::NEG_INFINITY; total_sequence_length];
for (ks, score_slot) in scores
.iter_mut()
.enumerate()
.take(causal_limit + 1)
.skip(local_start)
{
let mut score = 0.0;
for d in 0..cache_dim {
let ki = ((b * self.kv_num_heads + kvh) * present_sequence_length + ks)
* cache_dim
+ d;
score += q.at(b, qh, qs, d) * present_k[ki];
}
score *= scale;
if self.softcap != 0.0 {
score = self.softcap * (score / self.softcap).tanh();
}
*score_slot = score;
}
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for score in &mut scores {
if score.is_finite() {
*score = ((*score - max) as f64).exp() as f32;
sum += *score;
} else {
*score = 0.0;
}
}
for (d, output) in output_row.iter_mut().enumerate() {
let mut value = 0.0;
for (ks, probability) in scores.iter().enumerate() {
let vi =
((b * self.kv_num_heads + kvh) * present_sequence_length + ks) * v.dim + d;
value += (*probability / sum) * present_v[vi];
}
*output = value;
}
};
let attention_work = attention_rows
.saturating_mul(total_sequence_length)
.saturating_mul(cache_dim);
if attention_rows > 1 && attention_work >= MIN_PARALLEL_ATTENTION_WORK {
y_bhsd
.par_chunks_mut(v.dim)
.enumerate()
.for_each(|(row_index, output_row)| {
let b = row_index / (self.num_heads * q.seq);
let row_in_batch = row_index % (self.num_heads * q.seq);
let qh = row_in_batch / q.seq;
let qs = row_in_batch % q.seq;
compute_row(b, qh, qs, output_row);
});
} else {
for b in 0..q.batch {
for qh in 0..self.num_heads {
for qs in 0..q.seq {
let row_index = (b * self.num_heads + qh) * q.seq + qs;
compute_row(
b,
qh,
qs,
&mut y_bhsd[row_index * v.dim..(row_index + 1) * v.dim],
);
}
}
}
}
let mut output = vec![0.0; y_bhsd.len()];
for b in 0..q.batch {
for s in 0..q.seq {
for h in 0..self.num_heads {
for d in 0..v.dim {
output[((b * q.seq + s) * self.num_heads + h) * v.dim + d] =
y_bhsd[((b * self.num_heads + h) * q.seq + s) * v.dim + d];
}
}
}
}
let decode_fast_write = q.seq == 1 && k.seq == 1;
if decode_fast_write {
write_decode_output(&mut outputs[0], &output)?;
} else {
write_dense_f32_narrow("GroupQueryAttention", &mut outputs[0], &output)?;
}
if outputs.len() >= 2 {
if decode_fast_write {
write_decode_output(&mut outputs[1], &present_k)?;
} else {
write_dense_f32_narrow("GroupQueryAttention", &mut outputs[1], &present_k)?;
}
}
if outputs.len() >= 3 {
if decode_fast_write {
write_decode_output(&mut outputs[2], &present_v)?;
} else {
write_dense_f32_narrow("GroupQueryAttention", &mut outputs[2], &present_v)?;
}
}
Ok(())
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::CpuExecutionProvider;
use crate::kernels::testutil::Owned;
use onnx_runtime_ep_api::{ExecutionProvider, TensorView};
use onnx_runtime_ir::{Attribute, DataType, Graph, Node, NodeId, static_shape};
use onnx_runtime_loader::Model;
fn absent() -> TensorView<'static> {
TensorView::absent(DataType::Float32)
}
fn kernel(attrs: &[(&str, Attribute)]) -> Box<dyn Kernel> {
let mut graph = Graph::new();
graph.opset_imports.insert("com.microsoft".into(), 1);
let inputs = [
("query", DataType::Float32, vec![1, 1, 8]),
("key", DataType::Float32, vec![1, 1, 4]),
("value", DataType::Float32, vec![1, 1, 4]),
("past_key", DataType::Float32, vec![1, 2, 0, 2]),
("past_value", DataType::Float32, vec![1, 2, 0, 2]),
("seqlens_k", DataType::Int32, vec![1]),
("total_sequence_length", DataType::Int32, vec![]),
]
.into_iter()
.map(|(name, dtype, shape)| {
let value = graph.create_named_value(name, dtype, static_shape(shape));
graph.add_input(value);
Some(value)
})
.collect();
let output = graph.create_named_value("output", DataType::Float32, static_shape([1, 1, 8]));
let mut node = Node::new(NodeId(0), "GroupQueryAttention", inputs, vec![output]);
node.domain = "com.microsoft".into();
for (name, value) in attrs {
node.attributes.insert((*name).into(), value.clone());
}
let id = graph.insert_node(node);
graph.add_output(output);
let model = Model::new(&graph);
CpuExecutionProvider::new()
.get_kernel(model.graph.node(id), &[], 1)
.unwrap()
}
fn gqa_kernel(extra: &[(&str, Attribute)]) -> Box<dyn Kernel> {
let mut attrs = vec![
("num_heads", Attribute::Int(4)),
("kv_num_heads", Attribute::Int(2)),
];
attrs.extend_from_slice(extra);
kernel(&attrs)
}
fn reference(
q: &[f32],
k: &[f32],
v: &[f32],
q_seq: usize,
total: usize,
past: usize,
) -> Vec<f32> {
let (qh, kvh, d) = (4, 2, 2);
let mut out = vec![0.0; q_seq * qh * d];
for s in 0..q_seq {
for h in 0..qh {
let kh = h / (qh / kvh);
let mut scores = vec![0.0; past + s + 1];
for j in 0..scores.len() {
scores[j] = (0..d)
.map(|x| q[(s * qh + h) * d + x] * k[(kh * total + j) * d + x])
.sum::<f32>()
/ (d as f32).sqrt();
}
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let sum: f32 = scores
.iter_mut()
.map(|x| {
*x = ((*x - max) as f64).exp() as f32;
*x
})
.sum();
for x in &mut scores {
*x /= sum;
}
for x in 0..d {
out[(s * qh + h) * d + x] = scores
.iter()
.enumerate()
.map(|(j, p)| p * v[(kh * total + j) * d + x])
.sum();
}
}
}
out
}
fn close(got: &[f32], want: &[f32]) {
assert_eq!(got.len(), want.len());
for (i, (a, b)) in got.iter().zip(want).enumerate() {
assert!((a - b).abs() < 1e-5, "{i}: {a} != {b}");
}
}
fn reference_rope_bsh(
input: &[f32],
seq: usize,
heads: usize,
positions: &[usize],
cos: &[f32],
sin: &[f32],
) -> Vec<f32> {
let mut output = input.to_vec();
for s in 0..seq {
for h in 0..heads {
let base = (s * heads + h) * 2;
let (x0, x1) = (input[base], input[base + 1]);
output[base] = cos[positions[s]] * x0 - sin[positions[s]] * x1;
output[base + 1] = sin[positions[s]] * x0 + cos[positions[s]] * x1;
}
}
output
}
fn bsh_to_bnsh(input: &[f32], seq: usize, heads: usize) -> Vec<f32> {
let mut output = vec![0.0; input.len()];
for s in 0..seq {
for h in 0..heads {
output[(h * seq + s) * 2..(h * seq + s + 1) * 2]
.copy_from_slice(&input[(s * heads + h) * 2..(s * heads + h + 1) * 2]);
}
}
output
}
#[test]
fn prefill_gqa_grouping_and_causal_match_reference() {
let q = vec![
1., 0., 1., 0., 0., 1., 0., 1., 0., 1., 0., 1., 1., 0., 1., 0., 1., 1., 1., 1., -1.,
1., -1., 1.,
];
let k_bsh = vec![1., 0., 0., 1., 0., 1., 1., 0., 1., 1., -1., 1.];
let v_bsh = vec![1., 2., 10., 20., 3., 4., 30., 40., 5., 6., 50., 60.];
let mut k_bnsh = vec![0.0; 12];
let mut v_bnsh = vec![0.0; 12];
for h in 0..2 {
for s in 0..3 {
for d in 0..2 {
k_bnsh[(h * 3 + s) * 2 + d] = k_bsh[(s * 2 + h) * 2 + d];
v_bnsh[(h * 3 + s) * 2 + d] = v_bsh[(s * 2 + h) * 2 + d];
}
}
}
let mut out = Owned::zeros_f32(&[1, 3, 8]);
let mut pk = Owned::zeros_f32(&[1, 2, 3, 2]);
let mut pv = Owned::zeros_f32(&[1, 2, 3, 2]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, 3, 8], &q).view(),
Owned::f32(&[1, 3, 4], &k_bsh).view(),
Owned::f32(&[1, 3, 4], &v_bsh).view(),
absent(),
absent(),
Owned::i32(&[1], &[2]).view(),
Owned::i32(&[], &[3]).view(),
],
&mut [out.view_mut(), pk.view_mut(), pv.view_mut()],
)
.unwrap();
close(&out.to_f32(), &reference(&q, &k_bnsh, &v_bnsh, 3, 3, 0));
assert_eq!(pk.shape, vec![1, 2, 3, 2]);
close(&pk.to_f32(), &k_bnsh);
close(&pv.to_f32(), &v_bnsh);
}
#[test]
fn large_prefill_parallel_path_matches_reference() {
let seq = 160;
let q = (0..seq * 8)
.map(|i| ((i % 17) as f32 - 8.0) / 8.0)
.collect::<Vec<_>>();
let k_bsh = (0..seq * 4)
.map(|i| ((i % 13) as f32 - 6.0) / 7.0)
.collect::<Vec<_>>();
let v_bsh = (0..seq * 4)
.map(|i| ((i % 19) as f32 - 9.0) / 9.0)
.collect::<Vec<_>>();
let k_bnsh = bsh_to_bnsh(&k_bsh, seq, 2);
let v_bnsh = bsh_to_bnsh(&v_bsh, seq, 2);
let mut out = Owned::zeros_f32(&[1, seq, 8]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, seq, 8], &q).view(),
Owned::f32(&[1, seq, 4], &k_bsh).view(),
Owned::f32(&[1, seq, 4], &v_bsh).view(),
absent(),
absent(),
Owned::i32(&[1], &[(seq - 1) as i32]).view(),
Owned::i32(&[], &[seq as i32]).view(),
],
&mut [out.view_mut()],
)
.unwrap();
close(&out.to_f32(), &reference(&q, &k_bnsh, &v_bnsh, seq, seq, 0));
}
#[test]
fn packed_qkv_matches_unpacked_and_independent_reference() {
let q = vec![
1., 0., 1., 0., 0., 1., 0., 1., 0., 1., 0., 1., 1., 0., 1., 0.,
];
let k_bsh = vec![1., 0., 0., 1., 0., 1., 1., 0.];
let v_bsh = vec![1., 2., 10., 20., 3., 4., 30., 40.];
let mut packed = Vec::with_capacity(q.len() + k_bsh.len() + v_bsh.len());
for s in 0..2 {
packed.extend_from_slice(&q[s * 8..(s + 1) * 8]);
packed.extend_from_slice(&k_bsh[s * 4..(s + 1) * 4]);
packed.extend_from_slice(&v_bsh[s * 4..(s + 1) * 4]);
}
let k_bnsh = bsh_to_bnsh(&k_bsh, 2, 2);
let v_bnsh = bsh_to_bnsh(&v_bsh, 2, 2);
let want = reference(&q, &k_bnsh, &v_bnsh, 2, 2, 0);
let mut unpacked_out = Owned::zeros_f32(&[1, 2, 8]);
let mut packed_out = Owned::zeros_f32(&[1, 2, 8]);
let mut unpacked_k = Owned::zeros_f32(&[1, 2, 2, 2]);
let mut unpacked_v = Owned::zeros_f32(&[1, 2, 2, 2]);
let mut packed_k = Owned::zeros_f32(&[1, 2, 2, 2]);
let mut packed_v = Owned::zeros_f32(&[1, 2, 2, 2]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, 2, 8], &q).view(),
Owned::f32(&[1, 2, 4], &k_bsh).view(),
Owned::f32(&[1, 2, 4], &v_bsh).view(),
absent(),
absent(),
Owned::i32(&[1], &[1]).view(),
Owned::i32(&[], &[2]).view(),
],
&mut [
unpacked_out.view_mut(),
unpacked_k.view_mut(),
unpacked_v.view_mut(),
],
)
.unwrap();
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, 2, 16], &packed).view(),
absent(),
absent(),
absent(),
absent(),
Owned::i32(&[1], &[1]).view(),
Owned::i32(&[], &[2]).view(),
],
&mut [
packed_out.view_mut(),
packed_k.view_mut(),
packed_v.view_mut(),
],
)
.unwrap();
close(&unpacked_out.to_f32(), &want);
close(&packed_out.to_f32(), &want);
close(&packed_out.to_f32(), &unpacked_out.to_f32());
close(&packed_k.to_f32(), &unpacked_k.to_f32());
close(&packed_v.to_f32(), &unpacked_v.to_f32());
}
#[test]
fn decode_appends_past_and_matches_reference() {
let q = vec![1., 0., 1., 0., 0., 1., 0., 1.];
let past_k = vec![1., 0., 0., 1., 10., 0., 0., 10.];
let past_v = vec![1., 2., 3., 4., 10., 20., 30., 40.];
let cur_k = vec![1., 1., 10., 10.];
let cur_v = vec![5., 6., 50., 60.];
let mut all_k = vec![0.0; 12];
let mut all_v = vec![0.0; 12];
for h in 0..2 {
all_k[h * 6..h * 6 + 4].copy_from_slice(&past_k[h * 4..h * 4 + 4]);
all_v[h * 6..h * 6 + 4].copy_from_slice(&past_v[h * 4..h * 4 + 4]);
all_k[h * 6 + 4..h * 6 + 6].copy_from_slice(&cur_k[h * 2..h * 2 + 2]);
all_v[h * 6 + 4..h * 6 + 6].copy_from_slice(&cur_v[h * 2..h * 2 + 2]);
}
let mut out = Owned::zeros_f32(&[1, 1, 8]);
let mut pk = Owned::zeros_f32(&[1, 2, 3, 2]);
let mut pv = Owned::zeros_f32(&[1, 2, 3, 2]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, 1, 8], &q).view(),
Owned::f32(&[1, 1, 4], &cur_k).view(),
Owned::f32(&[1, 1, 4], &cur_v).view(),
Owned::f32(&[1, 2, 2, 2], &past_k).view(),
Owned::f32(&[1, 2, 2, 2], &past_v).view(),
Owned::i32(&[1], &[2]).view(),
Owned::i32(&[], &[3]).view(),
],
&mut [out.view_mut(), pk.view_mut(), pv.view_mut()],
)
.unwrap();
close(&pk.to_f32(), &all_k);
close(&pv.to_f32(), &all_v);
close(&out.to_f32(), &reference(&q, &all_k, &all_v, 1, 3, 2));
}
#[test]
fn decode_widens_f16_past_cache_before_materializing_present_cache() {
let q = vec![1., 0., 1., 0., 0., 1., 0., 1.];
let past_k = vec![1., 0., 0., 1., 10., 0., 0., 10.];
let past_v = vec![1., 2., 3., 4., 10., 20., 30., 40.];
let cur_k = vec![1., 1., 10., 10.];
let cur_v = vec![5., 6., 50., 60.];
let expected_k = vec![1., 0., 0., 1., 1., 1., 10., 0., 0., 10., 10., 10.];
let expected_v = vec![1., 2., 3., 4., 5., 6., 10., 20., 30., 40., 50., 60.];
let mut out = Owned::zeros_f32(&[1, 1, 8]);
let mut pk = Owned::zeros_f32(&[1, 2, 3, 2]);
let mut pv = Owned::zeros_f32(&[1, 2, 3, 2]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, 1, 8], &q).view(),
Owned::f32(&[1, 1, 4], &cur_k).view(),
Owned::f32(&[1, 1, 4], &cur_v).view(),
Owned::f16(&[1, 2, 2, 2], &past_k).view(),
Owned::f16(&[1, 2, 2, 2], &past_v).view(),
Owned::i32(&[1], &[2]).view(),
Owned::i32(&[], &[3]).view(),
],
&mut [out.view_mut(), pk.view_mut(), pv.view_mut()],
)
.unwrap();
close(&pk.to_f32(), &expected_k);
close(&pv.to_f32(), &expected_v);
close(
&out.to_f32(),
&reference(&q, &expected_k, &expected_v, 1, 3, 2),
);
}
#[test]
fn decode_batch_ragged_past_lengths_materialize_independently() {
let q = vec![
1., 0., 1., 0., 0., 1., 0., 1., 1., 1., 1., -1., -1., 1., -1., -1.,
];
let past_k = vec![
1., 0., 91., 92., 93., 94., 0., 1., 95., 96., 97., 98., 2., 0., 3., 0., 4., 0., 5., 0.,
6., 0., 7., 0.,
];
let past_v = vec![
1., 2., 71., 72., 73., 74., 3., 4., 75., 76., 77., 78., 10., 20., 30., 40., 50., 60.,
70., 80., 90., 100., 110., 120.,
];
let cur_k = vec![1., 1., 10., 10., 8., 0., 9., 0.];
let cur_v = vec![5., 6., 50., 60., 130., 140., 150., 160.];
let expected_k = vec![
1., 0., 1., 1., 0., 0., 0., 0., 0., 1., 10., 10., 0., 0., 0., 0., 2., 0., 3., 0., 4.,
0., 8., 0., 5., 0., 6., 0., 7., 0., 9., 0.,
];
let expected_v = vec![
1., 2., 5., 6., 0., 0., 0., 0., 3., 4., 50., 60., 0., 0., 0., 0., 10., 20., 30., 40.,
50., 60., 130., 140., 70., 80., 90., 100., 110., 120., 150., 160.,
];
let mut out = Owned::zeros_f32(&[2, 1, 8]);
let mut pk = Owned::zeros_f32(&[2, 2, 4, 2]);
let mut pv = Owned::zeros_f32(&[2, 2, 4, 2]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[2, 1, 8], &q).view(),
Owned::f32(&[2, 1, 4], &cur_k).view(),
Owned::f32(&[2, 1, 4], &cur_v).view(),
Owned::f32(&[2, 2, 3, 2], &past_k).view(),
Owned::f32(&[2, 2, 3, 2], &past_v).view(),
Owned::i32(&[2], &[1, 3]).view(),
Owned::i32(&[], &[4]).view(),
],
&mut [out.view_mut(), pk.view_mut(), pv.view_mut()],
)
.unwrap();
close(&pk.to_f32(), &expected_k);
close(&pv.to_f32(), &expected_v);
let mut want = reference(
&q[..8],
&[1., 0., 1., 1., 0., 1., 10., 10.],
&[1., 2., 5., 6., 3., 4., 50., 60.],
1,
2,
1,
);
want.extend(reference(
&q[8..],
&expected_k[16..],
&expected_v[16..],
1,
4,
3,
));
close(&out.to_f32(), &want);
}
#[test]
fn decode_preserves_fixed_cache_capacity_and_appends_at_logical_length() {
let q = vec![1., 0., 1., 0., 0., 1., 0., 1.];
let past_k = vec![
1., 0., 0., 1., 91., 92., 93., 94., 95., 96., 10., 0., 0., 10., 81., 82., 83., 84.,
85., 86.,
];
let past_v = vec![
1., 2., 3., 4., 71., 72., 73., 74., 75., 76., 10., 20., 30., 40., 61., 62., 63., 64.,
65., 66.,
];
let cur_k = vec![1., 1., 10., 10.];
let cur_v = vec![5., 6., 50., 60.];
let expected_k = vec![
1., 0., 0., 1., 1., 1., 0., 0., 0., 0., 10., 0., 0., 10., 10., 10., 0., 0., 0., 0.,
];
let expected_v = vec![
1., 2., 3., 4., 5., 6., 0., 0., 0., 0., 10., 20., 30., 40., 50., 60., 0., 0., 0., 0.,
];
let mut out = Owned::zeros_f32(&[1, 1, 8]);
let mut pk = Owned::zeros_f32(&[1, 2, 5, 2]);
let mut pv = Owned::zeros_f32(&[1, 2, 5, 2]);
gqa_kernel(&[])
.execute(
&[
Owned::f32(&[1, 1, 8], &q).view(),
Owned::f32(&[1, 1, 4], &cur_k).view(),
Owned::f32(&[1, 1, 4], &cur_v).view(),
Owned::f32(&[1, 2, 5, 2], &past_k).view(),
Owned::f32(&[1, 2, 5, 2], &past_v).view(),
Owned::i32(&[1], &[2]).view(),
Owned::i32(&[], &[3]).view(),
],
&mut [out.view_mut(), pk.view_mut(), pv.view_mut()],
)
.unwrap();
assert_eq!(pk.shape, vec![1, 2, 5, 2]);
assert_eq!(pv.shape, vec![1, 2, 5, 2]);
close(&pk.to_f32(), &expected_k);
close(&pv.to_f32(), &expected_v);
close(
&out.to_f32(),
&reference(&q, &expected_k, &expected_v, 1, 5, 2),
);
}
#[test]
fn rotary_path_matches_rotated_reference() {
let q = vec![1., 2., 3., 4., 5., 6., 7., 8.];
let k = vec![1., 2., 3., 4.];
let v = vec![1., 2., 3., 4.];
let cos = vec![0.0];
let sin = vec![1.0];
let q_rot = vec![-2., 1., -4., 3., -6., 5., -8., 7.];
let k_rot_bsh = vec![-2., 1., -4., 3.];
let k_rot_bnsh = vec![-2., 1., -4., 3.];
let mut out = Owned::zeros_f32(&[1, 1, 8]);
gqa_kernel(&[("do_rotary", Attribute::Int(1))])
.execute(
&[
Owned::f32(&[1, 1, 8], &q).view(),
Owned::f32(&[1, 1, 4], &k).view(),
Owned::f32(&[1, 1, 4], &v).view(),
absent(),
absent(),
Owned::i32(&[1], &[0]).view(),
Owned::i32(&[], &[1]).view(),
Owned::f32(&[1, 1], &cos).view(),
Owned::f32(&[1, 1], &sin).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let _ = k_rot_bsh;
close(&out.to_f32(), &reference(&q_rot, &k_rot_bnsh, &v, 1, 1, 0));
}
#[test]
fn rotary_explicit_position_ids_apply_to_query_and_key() {
let q = vec![
1., 2., 2., -1., -1., 3., 4., 2., 3., -2., 1., 4., -3., 1., 2., 5.,
];
let k = vec![2., 1., -1., 3., 4., -2., 2., 5.];
let v = vec![1., 2., 10., 20., 3., 4., 30., 40.];
let angles = [0.0_f32, 0.2, 0.7, 1.1, 1.6];
let cos: Vec<f32> = angles.iter().map(|angle| angle.cos()).collect();
let sin: Vec<f32> = angles.iter().map(|angle| angle.sin()).collect();
let positions = [2_usize, 4];
let q_rot = reference_rope_bsh(&q, 2, 4, &positions, &cos, &sin);
let k_rot_bsh = reference_rope_bsh(&k, 2, 2, &positions, &cos, &sin);
let k_rot_bnsh = bsh_to_bnsh(&k_rot_bsh, 2, 2);
let v_bnsh = bsh_to_bnsh(&v, 2, 2);
let mut out = Owned::zeros_f32(&[1, 2, 8]);
let mut present_k = Owned::zeros_f32(&[1, 2, 2, 2]);
gqa_kernel(&[("do_rotary", Attribute::Int(1))])
.execute(
&[
Owned::f32(&[1, 2, 8], &q).view(),
Owned::f32(&[1, 2, 4], &k).view(),
Owned::f32(&[1, 2, 4], &v).view(),
absent(),
absent(),
Owned::i32(&[1], &[1]).view(),
Owned::i32(&[], &[2]).view(),
Owned::f32(&[5, 1], &cos).view(),
Owned::f32(&[5, 1], &sin).view(),
Owned::i64(&[1, 2], &[2, 4]).view(),
],
&mut [out.view_mut(), present_k.view_mut()],
)
.unwrap();
close(&present_k.to_f32(), &k_rot_bnsh);
close(
&out.to_f32(),
&reference(&q_rot, &k_rot_bnsh, &v_bnsh, 2, 2, 0),
);
}
#[test]
fn local_window_masks_older_cache_tokens() {
let q = [0.0; 8];
let past_k = [0.0; 8];
let past_v = [1., 1., 2., 2., 10., 10., 20., 20.];
let cur_k = [0.0; 4];
let cur_v = [9., 9., 90., 90.];
let mut out = Owned::zeros_f32(&[1, 1, 8]);
gqa_kernel(&[("local_window_size", Attribute::Int(1))])
.execute(
&[
Owned::f32(&[1, 1, 8], &q).view(),
Owned::f32(&[1, 1, 4], &cur_k).view(),
Owned::f32(&[1, 1, 4], &cur_v).view(),
Owned::f32(&[1, 2, 2, 2], &past_k).view(),
Owned::f32(&[1, 2, 2, 2], &past_v).view(),
Owned::i32(&[1], &[2]).view(),
Owned::i32(&[], &[3]).view(),
],
&mut [out.view_mut()],
)
.unwrap();
close(&out.to_f32(), &[9., 9., 9., 9., 90., 90., 90., 90.]);
}
#[test]
fn softcap_matches_independent_score_transform() {
let q = [
2., 0., 2., 0., 2., 0., 2., 0., 2., 0., 2., 0., 2., 0., 2., 0.,
];
let k = [1., 0., 1., 0., 4., 0., 4., 0.];
let v = [1., 0., 10., 0., 3., 0., 30., 0.];
let mut out = Owned::zeros_f32(&[1, 2, 8]);
gqa_kernel(&[("softcap", Attribute::Float(1.5))])
.execute(
&[
Owned::f32(&[1, 2, 8], &q).view(),
Owned::f32(&[1, 2, 4], &k).view(),
Owned::f32(&[1, 2, 4], &v).view(),
absent(),
absent(),
Owned::i32(&[1], &[1]).view(),
Owned::i32(&[], &[2]).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let s0 = 1.5 * ((2.0 / 2.0_f32.sqrt()) / 1.5_f32).tanh();
let s1 = 1.5 * ((8.0 / 2.0_f32.sqrt()) / 1.5_f32).tanh();
let p1 = (s1 - s0).exp() / (1.0 + (s1 - s0).exp());
let expected_second = 1.0 * (1.0 - p1) + 3.0 * p1;
let expected = [
1.,
0.,
1.,
0.,
10.,
0.,
10.,
0.,
expected_second,
0.,
expected_second,
0.,
expected_second * 10.0,
0.,
expected_second * 10.0,
0.,
];
close(&out.to_f32(), &expected);
}
#[test]
fn explicit_zero_scale_matches_default_scale() {
let q = [
1., 0., 1., 0., 1., 0., 1., 0., 1., 0., 1., 0., 1., 0., 1., 0.,
];
let k = [0., 0., 0., 0., 4., 0., 4., 0.];
let v = [1., 0., 1., 0., 9., 0., 9., 0.];
let run = |attrs: &[(&str, Attribute)]| {
let mut out = Owned::zeros_f32(&[1, 2, 8]);
gqa_kernel(attrs)
.execute(
&[
Owned::f32(&[1, 2, 8], &q).view(),
Owned::f32(&[1, 2, 4], &k).view(),
Owned::f32(&[1, 2, 4], &v).view(),
absent(),
absent(),
Owned::i32(&[1], &[1]).view(),
Owned::i32(&[], &[2]).view(),
],
&mut [out.view_mut()],
)
.unwrap();
out.to_f32()
};
let default = run(&[]);
let zero = run(&[("scale", Attribute::Float(0.0))]);
close(&zero, &default);
assert!(zero[8] > 8.0, "zero scale produced uniform attention");
}
}