use onnx_runtime_ep_api::{EpError, Kernel, KernelFactory, Result, TensorMut, TensorView};
use onnx_runtime_ir::{DataType, Node};
use super::check_arity;
use super::sdpa::{
AttnBias, NoMask, QkCapture, QkCaptureStage, ScaleMode, SdpaConfig, SdpaTensors, sdpa_f32,
};
use crate::dtype::{to_dense_f32_widen, write_dense_f32_narrow};
pub struct AttentionKernel {
scale: Option<f32>,
is_causal: bool,
q_num_heads: Option<usize>,
kv_num_heads: Option<usize>,
qk_matmul_output_mode: i64,
softcap: f32,
since_version: u32,
}
pub struct AttentionFactory {
pub since_version: u32,
}
impl KernelFactory for AttentionFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let scale = node.attr("scale").and_then(|a| a.as_float());
let is_causal = node.attr("is_causal").and_then(|a| a.as_int()).unwrap_or(0) != 0;
let q_num_heads = node
.attr("q_num_heads")
.and_then(|a| a.as_int())
.map(|v| v as usize);
let kv_num_heads = node
.attr("kv_num_heads")
.and_then(|a| a.as_int())
.map(|v| v as usize);
let qk_matmul_output_mode = node
.attr("qk_matmul_output_mode")
.and_then(|a| a.as_int())
.unwrap_or(0);
let softcap = node
.attr("softcap")
.and_then(|a| a.as_float())
.unwrap_or(0.0);
if !(0..=3).contains(&qk_matmul_output_mode) {
return Err(EpError::KernelFailed(format!(
"Attention: qk_matmul_output_mode {qk_matmul_output_mode} is not supported \
(only 0, 1, 2, 3 are implemented)"
)));
}
Ok(Box::new(AttentionKernel {
scale,
is_causal,
q_num_heads,
kv_num_heads,
qk_matmul_output_mode,
softcap,
since_version: self.since_version,
}))
}
}
struct Bhsd {
data: Vec<f32>,
batch: usize,
heads: usize,
seq: usize,
dim: usize,
}
impl Bhsd {
#[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]
}
}
fn to_bhsd(view: &TensorView, name: &str, num_heads: Option<usize>) -> Result<Bhsd> {
let shape = view.shape;
match shape.len() {
4 => {
let (batch, heads, seq, dim) = (shape[0], shape[1], shape[2], shape[3]);
let data = to_dense_f32_widen("Attention", view)?.into_owned();
Ok(Bhsd {
data,
batch,
heads,
seq,
dim,
})
}
3 => {
let heads = num_heads.ok_or_else(|| {
EpError::KernelFailed(format!(
"Attention: 3D {name} input requires the corresponding \
q_num_heads/kv_num_heads attribute"
))
})?;
if heads == 0 {
return Err(EpError::KernelFailed(format!(
"Attention: {name} num_heads must be > 0"
)));
}
let (batch, seq, hidden) = (shape[0], shape[1], shape[2]);
if hidden % heads != 0 {
return Err(EpError::KernelFailed(format!(
"Attention: 3D {name} hidden size {hidden} is not divisible by num_heads \
{heads}"
)));
}
let dim = hidden / heads;
let src = to_dense_f32_widen("Attention", view)?;
let mut data = vec![0.0f32; batch * heads * seq * dim];
for b in 0..batch {
for s in 0..seq {
for h in 0..heads {
for d in 0..dim {
let src_i = ((b * seq + s) * heads + h) * dim + d;
let dst_i = ((b * heads + h) * seq + s) * dim + d;
data[dst_i] = src[src_i];
}
}
}
}
Ok(Bhsd {
data,
batch,
heads,
seq,
dim,
})
}
other => Err(EpError::KernelFailed(format!(
"Attention: {name} must be rank 3 or 4, got rank {other}"
))),
}
}
fn concat_cache(past: Option<&Bhsd>, cur: &Bhsd, name: &str) -> Result<Bhsd> {
let Some(past) = past else {
return Ok(Bhsd {
data: cur.data.clone(),
batch: cur.batch,
heads: cur.heads,
seq: cur.seq,
dim: cur.dim,
});
};
if past.batch != cur.batch || past.heads != cur.heads || past.dim != cur.dim {
return Err(EpError::KernelFailed(format!(
"Attention: past_{name} dims (b={},h={},d={}) incompatible with current \
(b={},h={},d={})",
past.batch, past.heads, past.dim, cur.batch, cur.heads, cur.dim
)));
}
let (batch, heads, dim) = (cur.batch, cur.heads, cur.dim);
let total = past.seq + cur.seq;
let mut data = vec![0.0f32; batch * heads * total * dim];
for b in 0..batch {
for h in 0..heads {
for d in 0..dim {
for j in 0..past.seq {
let dst = ((b * heads + h) * total + j) * dim + d;
data[dst] = past.at(b, h, j, d);
}
for j in 0..cur.seq {
let dst = ((b * heads + h) * total + past.seq + j) * dim + d;
data[dst] = cur.at(b, h, j, d);
}
}
}
}
Ok(Bhsd {
data,
batch,
heads,
seq: total,
dim,
})
}
enum Mask {
None,
Float {
data: Vec<f32>,
shape: Vec<usize>,
},
Bool {
data: Vec<bool>,
shape: Vec<usize>,
},
}
impl Mask {
fn bias(&self, b: usize, h: usize, i: usize, j: usize, total_seq: usize) -> f32 {
match self {
Mask::None => 0.0,
Mask::Float { data, shape } => Self::lookup_f32(data, shape, b, h, i, j, total_seq),
Mask::Bool { data, shape } => {
if !shape.is_empty() {
let last = shape[shape.len() - 1];
if j >= last && last < total_seq {
return f32::NEG_INFINITY;
}
}
if Self::lookup_bool(data, shape, b, h, i, j) {
0.0
} else {
f32::NEG_INFINITY
}
}
}
}
fn lookup_f32(
data: &[f32],
shape: &[usize],
b: usize,
h: usize,
i: usize,
j: usize,
total_seq: usize,
) -> f32 {
if !shape.is_empty() {
let last = shape[shape.len() - 1];
if j >= last && last < total_seq {
return f32::NEG_INFINITY;
}
}
data[Self::offset(shape, b, h, i, j)]
}
fn lookup_bool(data: &[bool], shape: &[usize], b: usize, h: usize, i: usize, j: usize) -> bool {
data[Self::offset(shape, b, h, i, j)]
}
fn offset(shape: &[usize], b: usize, h: usize, i: usize, j: usize) -> usize {
let full = [b, h, i, j];
let rank = shape.len();
let mut off = 0usize;
for (k, &dim) in shape.iter().enumerate() {
let logical = full[4 - rank + k];
let idx = if dim == 1 { 0 } else { logical };
off = off * dim + idx;
}
off
}
}
struct AttnMaskBias<'a> {
mask: &'a Mask,
is_causal: bool,
total_seq: usize,
offsets: &'a [i64],
pad_limits: &'a [Option<i64>],
}
impl AttnBias for AttnMaskBias<'_> {
#[inline]
fn at(&self, b: usize, head: usize, i: usize, j: usize) -> f32 {
if let Some(limit) = self.pad_limits[b]
&& (j as i64) >= limit
{
return f32::NEG_INFINITY;
}
if self.is_causal && (j as i64) > i as i64 + self.offsets[b] {
return f32::NEG_INFINITY;
}
self.mask.bias(b, head, i, j, self.total_seq)
}
}
impl Kernel for AttentionKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("Attention", inputs, outputs, 3, 7, 1)?;
let q_rank = inputs[0].shape.len();
let q = to_bhsd(&inputs[0], "Q", self.q_num_heads)?;
let k_cur = to_bhsd(&inputs[1], "K", self.kv_num_heads)?;
let v_cur = to_bhsd(&inputs[2], "V", self.kv_num_heads)?;
let has_past_key = inputs.len() > 4 && !inputs[4].is_absent();
let has_past_value = inputs.len() > 5 && !inputs[5].is_absent();
if has_past_key != has_past_value {
return Err(EpError::KernelFailed(
"Attention: past_key and past_value must be provided together".into(),
));
}
let past_key = if has_past_key {
Some(to_bhsd(&inputs[4], "past_key", self.kv_num_heads)?)
} else {
None
};
let past_value = if has_past_value {
Some(to_bhsd(&inputs[5], "past_value", self.kv_num_heads)?)
} else {
None
};
let past_seq = past_key.as_ref().map(|p| p.seq).unwrap_or(0);
let has_nonpad = inputs.len() > 6 && !inputs[6].is_absent();
if has_nonpad && self.since_version < 24 {
return Err(EpError::KernelFailed(
"Attention: the optional `nonpad_kv_seqlen` input was added in opset 24 and is \
not valid for opset 23"
.into(),
));
}
if has_nonpad && (has_past_key || has_past_value) {
return Err(EpError::KernelFailed(
"Attention: `nonpad_kv_seqlen` must not be used together with past_key/past_value \
(external vs. in-op KV cache)"
.into(),
));
}
let nonpad_kv_seqlen: Option<Vec<i64>> = if has_nonpad {
let seqlen = super::to_dense_i64(&inputs[6])?;
if seqlen.len() != q.batch {
return Err(EpError::KernelFailed(format!(
"Attention: nonpad_kv_seqlen length {} must equal batch_size {}",
seqlen.len(),
q.batch
)));
}
Some(seqlen)
} else {
None
};
let key = concat_cache(past_key.as_ref(), &k_cur, "key")?;
let value = concat_cache(past_value.as_ref(), &v_cur, "value")?;
let batch = q.batch;
let q_heads = q.heads;
let q_seq = q.seq;
let head_size = q.dim;
let kv_heads = key.heads;
let total_seq = key.seq;
let v_head_size = value.dim;
if key.dim != head_size {
return Err(EpError::KernelFailed(format!(
"Attention: Q head_size {head_size} != K head_size {}",
key.dim
)));
}
if value.seq != total_seq {
return Err(EpError::KernelFailed(format!(
"Attention: present_key seq {total_seq} != present_value seq {}",
value.seq
)));
}
if key.batch != batch || value.batch != batch {
return Err(EpError::KernelFailed(
"Attention: Q, K, V must share the batch dimension".into(),
));
}
if kv_heads == 0 || q_heads % kv_heads != 0 {
return Err(EpError::KernelFailed(format!(
"Attention: q_num_heads {q_heads} must be a positive multiple of kv_num_heads \
{kv_heads} (MHA/GQA/MQA)"
)));
}
crate::trace::record_kernel_metrics(inputs, outputs, || {
let score_elements = (batch as u64)
.saturating_mul(q_heads as u64)
.saturating_mul(q_seq as u64)
.saturating_mul(total_seq as u64);
let qk_flops = score_elements
.saturating_mul(head_size as u64)
.saturating_mul(2);
let pv_flops = score_elements
.saturating_mul(v_head_size as u64)
.saturating_mul(2);
let softmax_flops = score_elements.saturating_mul(4).saturating_add(
(batch as u64)
.saturating_mul(q_heads as u64)
.saturating_mul(q_seq as u64),
);
qk_flops
.saturating_add(pv_flops)
.saturating_add(softmax_flops)
});
let scale = self
.scale
.unwrap_or_else(|| 1.0 / (head_size as f32).sqrt());
let mask = if inputs.len() > 3 && !inputs[3].is_absent() {
let m = &inputs[3];
match m.dtype {
DataType::Bool => Mask::Bool {
data: super::to_dense_bytes(m)?.iter().map(|&b| b != 0).collect(),
shape: m.shape.to_vec(),
},
DataType::Float32 | DataType::Float16 | DataType::BFloat16 | DataType::Float64 => {
Mask::Float {
data: to_dense_f32_widen("Attention", m)?.into_owned(),
shape: m.shape.to_vec(),
}
}
other => {
return Err(EpError::KernelFailed(format!(
"Attention: attn_mask dtype {other:?} not supported (expected bool or floating-point)"
)));
}
}
} else {
Mask::None
};
let mut offsets = vec![0i64; batch];
let mut pad_limits: Vec<Option<i64>> = vec![None; batch];
for b in 0..batch {
offsets[b] = match &nonpad_kv_seqlen {
Some(seqlen) => seqlen[b] - q_seq as i64,
None => past_seq as i64,
};
pad_limits[b] = nonpad_kv_seqlen.as_ref().map(|seqlen| seqlen[b]);
}
let bias = AttnMaskBias {
mask: &mask,
is_causal: self.is_causal,
total_seq,
offsets: &offsets,
pad_limits: &pad_limits,
};
let cfg = SdpaConfig {
scale: ScaleMode::SplitSqrt(scale),
softcap: (self.softcap != 0.0).then_some(self.softcap),
causal: false,
past_seq: 0,
causal_fill: f32::NEG_INFINITY,
};
let tensors = SdpaTensors {
q: &q.data,
k: &key.data,
v: &value.data,
batch,
num_heads: q_heads,
num_kv_heads: kv_heads,
q_seq,
kv_seq: total_seq,
head_size,
v_head_size,
};
let mut y = vec![0.0f32; batch * q_heads * q_seq * v_head_size];
let want_qk = outputs.len() >= 4;
let mut qk_out = if want_qk {
vec![0.0f32; batch * q_heads * q_seq * total_seq]
} else {
Vec::new()
};
let qk_cap = want_qk.then(|| {
let stage = match self.qk_matmul_output_mode {
0 => QkCaptureStage::PostScale,
1 => QkCaptureStage::PostSoftcap,
2 => QkCaptureStage::PreSoftmax,
_ => QkCaptureStage::PostSoftmax,
};
QkCapture {
scores: &mut qk_out,
stage,
}
});
sdpa_f32(&tensors, &cfg, &bias, &NoMask, &mut y, qk_cap);
if q_rank == 3 {
let hidden = q_heads * v_head_size;
let mut y3 = vec![0.0f32; batch * q_seq * hidden];
for b in 0..batch {
for h in 0..q_heads {
for s in 0..q_seq {
for c in 0..v_head_size {
let src = ((b * q_heads + h) * q_seq + s) * v_head_size + c;
let dst = (b * q_seq + s) * hidden + h * v_head_size + c;
y3[dst] = y[src];
}
}
}
}
write_dense_f32_narrow("Attention", &mut outputs[0], &y3)?;
} else {
write_dense_f32_narrow("Attention", &mut outputs[0], &y)?;
}
if outputs.len() >= 2 {
write_dense_f32_narrow("Attention", &mut outputs[1], &key.data)?;
}
if outputs.len() >= 3 {
write_dense_f32_narrow("Attention", &mut outputs[2], &value.data)?;
}
if want_qk {
write_dense_f32_narrow("Attention", &mut outputs[3], &qk_out)?;
}
Ok(())
}
fn supports_strided_input(&self, _input_idx: usize) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
#[allow(clippy::too_many_arguments)]
fn reference(
q: &[f32],
k: &[f32],
v: &[f32],
batch: usize,
q_heads: usize,
kv_heads: usize,
q_seq: usize,
total_seq: usize,
head_size: usize,
v_head_size: usize,
scale: f32,
is_causal: bool,
causal_offset: i64,
bias: impl Fn(usize, usize, usize, usize) -> f32,
) -> Vec<f32> {
let group = q_heads / kv_heads;
let mut out = vec![0.0f32; batch * q_heads * q_seq * v_head_size];
for b in 0..batch {
for qh in 0..q_heads {
let kvh = qh / group;
for i in 0..q_seq {
let mut scores = vec![0.0f32; total_seq];
for (j, sc) in scores.iter_mut().enumerate() {
let mut acc = 0.0f32;
for p in 0..head_size {
let qi = ((b * q_heads + qh) * q_seq + i) * head_size + p;
let kj = ((b * kv_heads + kvh) * total_seq + j) * head_size + p;
acc += q[qi] * k[kj];
}
let mut s = acc * scale + bias(b, qh, i, j);
if is_causal && (j as i64) > i as i64 + causal_offset {
s = f32::NEG_INFINITY;
}
*sc = s;
}
let max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
if max == f32::NEG_INFINITY {
continue; }
let mut sum = 0.0f32;
for sc in scores.iter_mut() {
*sc = (*sc - max).exp();
sum += *sc;
}
for sc in scores.iter_mut() {
*sc /= sum;
}
for c in 0..v_head_size {
let mut acc = 0.0f32;
for (j, &p) in scores.iter().enumerate() {
let vj = ((b * kv_heads + kvh) * total_seq + j) * v_head_size + c;
acc += p * v[vj];
}
out[((b * q_heads + qh) * q_seq + i) * v_head_size + c] = acc;
}
}
}
}
out
}
fn approx(a: &[f32], b: &[f32], atol: f32) {
assert_eq!(
a.len(),
b.len(),
"length mismatch: {} vs {}",
a.len(),
b.len()
);
for (i, (x, y)) in a.iter().zip(b).enumerate() {
assert!(
(x - y).abs() < atol,
"element {i}: {x} vs {y}\n{a:?}\nvs\n{b:?}"
);
}
}
fn bsh_to_bhsd(src: &[f32], batch: usize, seq: usize, heads: usize, dim: usize) -> Vec<f32> {
let mut dst = 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 {
dst[((b * heads + h) * seq + s) * dim + d] =
src[((b * seq + s) * heads + h) * dim + d];
}
}
}
}
dst
}
fn bhsd_to_bsh(src: &[f32], batch: usize, heads: usize, seq: usize, dim: usize) -> Vec<f32> {
let mut dst = vec![0.0; src.len()];
for b in 0..batch {
for h in 0..heads {
for s in 0..seq {
for d in 0..dim {
dst[((b * seq + s) * heads + h) * dim + d] =
src[((b * heads + h) * seq + s) * dim + d];
}
}
}
}
dst
}
fn concat_bsh(
first: &[f32],
second: &[f32],
batch: usize,
first_seq: usize,
second_seq: usize,
hidden: usize,
) -> Vec<f32> {
let total_seq = first_seq + second_seq;
let mut dst = vec![0.0; batch * total_seq * hidden];
for b in 0..batch {
let first_src = b * first_seq * hidden;
let second_src = b * second_seq * hidden;
let dst_base = b * total_seq * hidden;
dst[dst_base..dst_base + first_seq * hidden]
.copy_from_slice(&first[first_src..first_src + first_seq * hidden]);
dst[dst_base + first_seq * hidden..dst_base + total_seq * hidden]
.copy_from_slice(&second[second_src..second_src + second_seq * hidden]);
}
dst
}
#[allow(clippy::too_many_arguments)]
fn concat_bhsd(
past: &[f32],
current: &[f32],
batch: usize,
heads: usize,
past_seq: usize,
current_seq: usize,
dim: usize,
) -> Vec<f32> {
let total_seq = past_seq + current_seq;
let mut dst = vec![0.0; batch * heads * total_seq * dim];
for b in 0..batch {
for h in 0..heads {
for s in 0..past_seq {
let src = ((b * heads + h) * past_seq + s) * dim;
let out = ((b * heads + h) * total_seq + s) * dim;
dst[out..out + dim].copy_from_slice(&past[src..src + dim]);
}
for s in 0..current_seq {
let src = ((b * heads + h) * current_seq + s) * dim;
let out = ((b * heads + h) * total_seq + past_seq + s) * dim;
dst[out..out + dim].copy_from_slice(¤t[src..src + dim]);
}
}
}
dst
}
fn conformance_values(len: usize, salt: usize) -> Vec<f32> {
(0..len)
.map(|i| (((i * 17 + salt * 29) % 97) as f32 - 48.0) / 128.0)
.collect()
}
fn kernel(
scale: Option<f32>,
is_causal: bool,
q_num_heads: Option<usize>,
kv_num_heads: Option<usize>,
qk_mode: i64,
softcap: f32,
) -> AttentionKernel {
kernel_v(
24,
scale,
is_causal,
q_num_heads,
kv_num_heads,
qk_mode,
softcap,
)
}
#[allow(clippy::too_many_arguments)]
fn kernel_v(
since_version: u32,
scale: Option<f32>,
is_causal: bool,
q_num_heads: Option<usize>,
kv_num_heads: Option<usize>,
qk_mode: i64,
softcap: f32,
) -> AttentionKernel {
AttentionKernel {
scale,
is_causal,
q_num_heads,
kv_num_heads,
qk_matmul_output_mode: qk_mode,
softcap,
since_version,
}
}
fn absent() -> TensorView<'static> {
TensorView::absent(DataType::Float32)
}
#[test]
fn mha_4d_no_mask_matches_reference() {
let (b, h, sq, sk, d, dv) = (2, 2, 3, 4, 5, 6);
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.1).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.07).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv)
.map(|i| (i as f32 * 0.03) - 0.5)
.collect();
let scale = 0.3f32;
let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
let qv = Owned::f32(&[b, h, sq, d], &q);
let kv = Owned::f32(&[b, h, sk, d], &k);
let vv = Owned::f32(&[b, h, sk, dv], &v);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(&[qv.view(), kv.view(), vv.view()], &mut [out.view_mut()])
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
}
#[test]
fn default_scale_is_inv_sqrt_head_size() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 4, 3);
let q: Vec<f32> = (0..b * h * sq * d).map(|i| i as f32 * 0.2).collect();
let k: Vec<f32> = (0..b * h * sk * d).map(|i| i as f32 * 0.1).collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.05).collect();
let scale = 1.0 / (d as f32).sqrt();
let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(None, false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-6);
}
#[test]
fn three_d_reshape_matches_four_d() {
let (b, h, sq, sk, d, dv) = (2, 2, 2, 3, 2, 2);
let scale = 0.5f32;
let q3: Vec<f32> = (0..b * sq * h * d)
.map(|i| (i as f32 * 0.13).sin())
.collect();
let k3: Vec<f32> = (0..b * sk * h * d)
.map(|i| (i as f32 * 0.09).cos())
.collect();
let v3: Vec<f32> = (0..b * sk * h * dv).map(|i| i as f32 * 0.02).collect();
let to4d = |src: &[f32], s: usize, dd: usize| {
let mut out = vec![0.0f32; b * h * s * dd];
for bb in 0..b {
for ss in 0..s {
for hh in 0..h {
for e in 0..dd {
let si = ((bb * s + ss) * h + hh) * dd + e;
let di = ((bb * h + hh) * s + ss) * dd + e;
out[di] = src[si];
}
}
}
}
out
};
let q4 = to4d(&q3, sq, d);
let k4 = to4d(&k3, sk, d);
let v4 = to4d(&v3, sk, dv);
let want4 = reference(
&q4,
&k4,
&v4,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
let mut want3 = vec![0.0f32; b * sq * h * dv];
for bb in 0..b {
for hh in 0..h {
for ss in 0..sq {
for c in 0..dv {
let si = ((bb * h + hh) * sq + ss) * dv + c;
let di = (bb * sq + ss) * (h * dv) + hh * dv + c;
want3[di] = want4[si];
}
}
}
}
let mut out = Owned::zeros_f32(&[b, sq, h * dv]);
kernel(Some(scale), false, Some(h), Some(h), 0, 0.0)
.execute(
&[
Owned::f32(&[b, sq, h * d], &q3).view(),
Owned::f32(&[b, sk, h * d], &k3).view(),
Owned::f32(&[b, sk, h * dv], &v3).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want3, 1e-5);
}
#[test]
fn mla_gqa_decode_with_asymmetric_head_dims_matches_hand_computed_result() {
let (batch, q_heads, kv_heads, seq, past_seq) = (1, 4, 2, 1, 2);
let (qk_head_dim, v_head_dim) = (2, 1);
let total_seq = past_seq + seq;
let q = [0.0; 8];
let k = [10.0, 20.0, 30.0, 40.0];
let v = [5.0, 8.0];
let past_key = [1.0, 0.0, 0.0, 1.0, 2.0, 0.0, 0.0, 2.0];
let past_value = [1.0, 3.0, 2.0, 4.0];
let mut y = Owned::zeros_f32(&[batch, seq, q_heads * v_head_dim]);
let mut present_key = Owned::zeros_f32(&[batch, kv_heads, total_seq, qk_head_dim]);
let mut present_value = Owned::zeros_f32(&[batch, kv_heads, total_seq, v_head_dim]);
kernel(None, true, Some(q_heads), Some(kv_heads), 0, 0.0)
.execute(
&[
Owned::f32(&[batch, seq, q_heads * qk_head_dim], &q).view(),
Owned::f32(&[batch, seq, kv_heads * qk_head_dim], &k).view(),
Owned::f32(&[batch, seq, kv_heads * v_head_dim], &v).view(),
absent(),
Owned::f32(&[batch, kv_heads, past_seq, qk_head_dim], &past_key).view(),
Owned::f32(&[batch, kv_heads, past_seq, v_head_dim], &past_value).view(),
],
&mut [
y.view_mut(),
present_key.view_mut(),
present_value.view_mut(),
],
)
.unwrap();
assert_eq!(y.shape, [batch, seq, q_heads * v_head_dim]);
assert_eq!(present_key.shape, [batch, kv_heads, total_seq, qk_head_dim]);
assert_eq!(
present_value.shape,
[batch, kv_heads, total_seq, v_head_dim]
);
approx(&y.to_f32(), &[3.0, 3.0, 14.0 / 3.0, 14.0 / 3.0], 1e-6);
assert_eq!(
present_key.to_f32(),
[
1.0, 0.0, 0.0, 1.0, 10.0, 20.0, 2.0, 0.0, 0.0, 2.0, 30.0, 40.0
]
);
assert_eq!(present_value.to_f32(), [1.0, 3.0, 5.0, 2.0, 4.0, 8.0]);
}
fn asymmetric_3d_prefill_decode_case(kv_heads: usize) {
let (batch, q_heads, prefill_seq, decode_seq) = (1usize, 4usize, 3usize, 1usize);
let (qk_head_dim, v_head_dim) = (192usize, 128usize);
assert_ne!(qk_head_dim, v_head_dim);
let q_prefill =
conformance_values(batch * prefill_seq * q_heads * qk_head_dim, 1 + kv_heads);
let k_prefill =
conformance_values(batch * prefill_seq * kv_heads * qk_head_dim, 3 + kv_heads);
let v_prefill =
conformance_values(batch * prefill_seq * kv_heads * v_head_dim, 5 + kv_heads);
let q_prefill_bhsd = bsh_to_bhsd(&q_prefill, batch, prefill_seq, q_heads, qk_head_dim);
let k_prefill_bhsd = bsh_to_bhsd(&k_prefill, batch, prefill_seq, kv_heads, qk_head_dim);
let v_prefill_bhsd = bsh_to_bhsd(&v_prefill, batch, prefill_seq, kv_heads, v_head_dim);
let scale = 1.0 / (qk_head_dim as f32).sqrt();
let prefill_oracle_bhsd = reference(
&q_prefill_bhsd,
&k_prefill_bhsd,
&v_prefill_bhsd,
batch,
q_heads,
kv_heads,
prefill_seq,
prefill_seq,
qk_head_dim,
v_head_dim,
scale,
true,
0,
|_, _, _, _| 0.0,
);
let prefill_oracle = bhsd_to_bsh(
&prefill_oracle_bhsd,
batch,
q_heads,
prefill_seq,
v_head_dim,
);
let mut prefill_y = Owned::zeros_f32(&[batch, prefill_seq, q_heads * v_head_dim]);
let mut present_key = Owned::zeros_f32(&[batch, kv_heads, prefill_seq, qk_head_dim]);
let mut present_value = Owned::zeros_f32(&[batch, kv_heads, prefill_seq, v_head_dim]);
kernel(None, true, Some(q_heads), Some(kv_heads), 0, 0.0)
.execute(
&[
Owned::f32(&[batch, prefill_seq, q_heads * qk_head_dim], &q_prefill).view(),
Owned::f32(&[batch, prefill_seq, kv_heads * qk_head_dim], &k_prefill).view(),
Owned::f32(&[batch, prefill_seq, kv_heads * v_head_dim], &v_prefill).view(),
],
&mut [
prefill_y.view_mut(),
present_key.view_mut(),
present_value.view_mut(),
],
)
.unwrap();
assert_eq!(
prefill_y.shape,
[batch, prefill_seq, q_heads * v_head_dim],
"3D Attention output hidden width must use V head width"
);
assert_eq!(
present_key.shape,
[batch, kv_heads, prefill_seq, qk_head_dim],
"present_key must preserve Q/K head width"
);
assert_eq!(
present_value.shape,
[batch, kv_heads, prefill_seq, v_head_dim],
"present_value must preserve V head width"
);
approx(&prefill_y.to_f32(), &prefill_oracle, 1e-5);
assert_eq!(present_key.to_f32(), k_prefill_bhsd);
assert_eq!(present_value.to_f32(), v_prefill_bhsd);
let q_decode = conformance_values(batch * decode_seq * q_heads * qk_head_dim, 7 + kv_heads);
let k_decode =
conformance_values(batch * decode_seq * kv_heads * qk_head_dim, 9 + kv_heads);
let v_decode =
conformance_values(batch * decode_seq * kv_heads * v_head_dim, 11 + kv_heads);
let q_decode_bhsd = bsh_to_bhsd(&q_decode, batch, decode_seq, q_heads, qk_head_dim);
let k_decode_bhsd = bsh_to_bhsd(&k_decode, batch, decode_seq, kv_heads, qk_head_dim);
let v_decode_bhsd = bsh_to_bhsd(&v_decode, batch, decode_seq, kv_heads, v_head_dim);
let full_key = concat_bhsd(
&k_prefill_bhsd,
&k_decode_bhsd,
batch,
kv_heads,
prefill_seq,
decode_seq,
qk_head_dim,
);
let full_value = concat_bhsd(
&v_prefill_bhsd,
&v_decode_bhsd,
batch,
kv_heads,
prefill_seq,
decode_seq,
v_head_dim,
);
let total_seq = prefill_seq + decode_seq;
let decode_oracle_bhsd = reference(
&q_decode_bhsd,
&full_key,
&full_value,
batch,
q_heads,
kv_heads,
decode_seq,
total_seq,
qk_head_dim,
v_head_dim,
scale,
true,
prefill_seq as i64,
|_, _, _, _| 0.0,
);
let decode_oracle =
bhsd_to_bsh(&decode_oracle_bhsd, batch, q_heads, decode_seq, v_head_dim);
let mut decode_y = Owned::zeros_f32(&[batch, decode_seq, q_heads * v_head_dim]);
let mut decode_present_key = Owned::zeros_f32(&[batch, kv_heads, total_seq, qk_head_dim]);
let mut decode_present_value = Owned::zeros_f32(&[batch, kv_heads, total_seq, v_head_dim]);
kernel(None, true, Some(q_heads), Some(kv_heads), 0, 0.0)
.execute(
&[
Owned::f32(&[batch, decode_seq, q_heads * qk_head_dim], &q_decode).view(),
Owned::f32(&[batch, decode_seq, kv_heads * qk_head_dim], &k_decode).view(),
Owned::f32(&[batch, decode_seq, kv_heads * v_head_dim], &v_decode).view(),
absent(),
present_key.view(),
present_value.view(),
],
&mut [
decode_y.view_mut(),
decode_present_key.view_mut(),
decode_present_value.view_mut(),
],
)
.unwrap();
assert_eq!(decode_y.shape, [batch, decode_seq, q_heads * v_head_dim]);
assert_eq!(
decode_present_key.shape,
[batch, kv_heads, total_seq, qk_head_dim]
);
assert_eq!(
decode_present_value.shape,
[batch, kv_heads, total_seq, v_head_dim]
);
approx(&decode_y.to_f32(), &decode_oracle, 1e-5);
assert_eq!(decode_present_key.to_f32(), full_key);
assert_eq!(decode_present_value.to_f32(), full_value);
let full_q_bsh = concat_bsh(
&q_prefill,
&q_decode,
batch,
prefill_seq,
decode_seq,
q_heads * qk_head_dim,
);
let full_k_bsh = concat_bsh(
&k_prefill,
&k_decode,
batch,
prefill_seq,
decode_seq,
kv_heads * qk_head_dim,
);
let full_v_bsh = concat_bsh(
&v_prefill,
&v_decode,
batch,
prefill_seq,
decode_seq,
kv_heads * v_head_dim,
);
let mut full_y = Owned::zeros_f32(&[batch, total_seq, q_heads * v_head_dim]);
let mut full_present_key = Owned::zeros_f32(&[batch, kv_heads, total_seq, qk_head_dim]);
let mut full_present_value = Owned::zeros_f32(&[batch, kv_heads, total_seq, v_head_dim]);
kernel(None, true, Some(q_heads), Some(kv_heads), 0, 0.0)
.execute(
&[
Owned::f32(&[batch, total_seq, q_heads * qk_head_dim], &full_q_bsh).view(),
Owned::f32(&[batch, total_seq, kv_heads * qk_head_dim], &full_k_bsh).view(),
Owned::f32(&[batch, total_seq, kv_heads * v_head_dim], &full_v_bsh).view(),
],
&mut [
full_y.view_mut(),
full_present_key.view_mut(),
full_present_value.view_mut(),
],
)
.unwrap();
let expected_full_y = concat_bsh(
&prefill_oracle,
&decode_oracle,
batch,
prefill_seq,
decode_seq,
q_heads * v_head_dim,
);
approx(&full_y.to_f32(), &expected_full_y, 1e-5);
assert_eq!(full_present_key.to_f32(), decode_present_key.to_f32());
assert_eq!(full_present_value.to_f32(), decode_present_value.to_f32());
}
#[test]
fn asymmetric_3d_prefill_decode_gqa_matches_scalar_oracle() {
asymmetric_3d_prefill_decode_case(2);
}
#[test]
fn asymmetric_3d_prefill_decode_mqa_matches_scalar_oracle() {
asymmetric_3d_prefill_decode_case(1);
}
#[test]
fn gqa_head_sharing() {
let (b, qh, kvh, sq, sk, d, dv) = (1, 4, 2, 2, 3, 3, 2);
let scale = 0.4f32;
let q: Vec<f32> = (0..b * qh * sq * d)
.map(|i| (i as f32 * 0.11).sin())
.collect();
let k: Vec<f32> = (0..b * kvh * sk * d)
.map(|i| (i as f32 * 0.08).cos())
.collect();
let v: Vec<f32> = (0..b * kvh * sk * dv)
.map(|i| i as f32 * 0.04 - 1.0)
.collect();
let want = reference(
&q,
&k,
&v,
b,
qh,
kvh,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, qh, sq, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, qh, sq, d], &q).view(),
Owned::f32(&[b, kvh, sk, d], &k).view(),
Owned::f32(&[b, kvh, sk, dv], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
}
#[test]
fn causal_masking_blocks_future() {
let (b, h, s, d, dv) = (1, 1, 4, 3, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * s * d)
.map(|i| (i as f32 * 0.17).sin())
.collect();
let k: Vec<f32> = (0..b * h * s * d)
.map(|i| (i as f32 * 0.05).cos())
.collect();
let v: Vec<f32> = (0..b * h * s * dv).map(|i| i as f32 * 0.1).collect();
let want = reference(
&q,
&k,
&v,
b,
h,
h,
s,
s,
d,
dv,
scale,
true,
0,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, h, s, dv]);
kernel(Some(scale), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, s, d], &q).view(),
Owned::f32(&[b, h, s, d], &k).view(),
Owned::f32(&[b, h, s, dv], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
let got = out.to_f32();
approx(&got[0..dv], &v[0..dv], 1e-5);
}
#[test]
fn float_additive_mask() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let scale = 0.6f32;
let q = [1.0f32, 2.0, -1.0, 0.5];
let k = [1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0];
let v = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let mask = [0.0f32, -1e4, 0.0, 0.0, 0.0, -1e4]; let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, i, j| mask[i * sk + j],
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::f32(&[sq, sk], &mask).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-4);
}
#[test]
fn bool_mask_matches_neg_inf_bias() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let scale = 0.6f32;
let q = [1.0f32, 2.0, -1.0, 0.5];
let k = [1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0];
let v = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let keep = [true, false, true, true, true, false];
let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, i, j| {
if keep[i * sk + j] {
0.0
} else {
f32::NEG_INFINITY
}
},
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::bool_(&[sq, sk], &keep).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
}
#[test]
fn fully_masked_row_is_zero_not_nan() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let q = [1.0f32, 2.0, 3.0, 4.0];
let k = [1.0f32, 0.0, 0.0, 1.0];
let v = [5.0f32, 6.0, 7.0, 8.0];
let keep = [false, false, true, true];
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(0.5), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::bool_(&[sq, sk], &keep).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let got = out.to_f32();
assert!(got.iter().all(|x| x.is_finite()), "no NaN/inf: {got:?}");
approx(&got[0..dv], &[0.0, 0.0], 1e-6);
}
#[test]
fn kv_cache_concat_and_present_outputs() {
let (b, h, sq, d, dv) = (1, 1, 1, 2, 2);
let past_seq = 2usize;
let kv_seq = 1usize;
let total = past_seq + kv_seq;
let scale = 0.5f32;
let q = [0.5f32, -0.5];
let past_k = [1.0f32, 0.0, 0.0, 1.0]; let cur_k = [1.0f32, 1.0]; let past_v = [1.0f32, 2.0, 3.0, 4.0];
let cur_v = [5.0f32, 6.0];
let mut full_k = past_k.to_vec();
full_k.extend_from_slice(&cur_k);
let mut full_v = past_v.to_vec();
full_v.extend_from_slice(&cur_v);
let want = reference(
&q,
&full_k,
&full_v,
b,
h,
h,
sq,
total,
d,
dv,
scale,
false,
past_seq as i64,
|_, _, _, _| 0.0,
);
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, total, d]);
let mut pv = Owned::zeros_f32(&[b, h, total, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, kv_seq, d], &cur_k).view(),
Owned::f32(&[b, h, kv_seq, dv], &cur_v).view(),
absent(), Owned::f32(&[b, h, past_seq, d], &past_k).view(),
Owned::f32(&[b, h, past_seq, dv], &past_v).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut()],
)
.unwrap();
approx(&y.to_f32(), &want, 1e-5);
approx(&pk.to_f32(), &full_k, 1e-6);
approx(&pv.to_f32(), &full_v, 1e-6);
}
#[test]
fn causal_kv_cache_offset() {
let (b, h, sq, d, dv) = (1, 1, 2, 2, 2);
let past_seq = 2usize;
let kv_seq = 2usize;
let total = past_seq + kv_seq;
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let cur_k: Vec<f32> = (0..kv_seq * d).map(|i| (i as f32 * 0.2).cos()).collect();
let cur_v: Vec<f32> = (0..kv_seq * dv).map(|i| i as f32 * 0.5).collect();
let past_k: Vec<f32> = (0..past_seq * d).map(|i| i as f32 * 0.1).collect();
let past_v: Vec<f32> = (0..past_seq * dv).map(|i| i as f32 * 0.3).collect();
let mut full_k = past_k.clone();
full_k.extend_from_slice(&cur_k);
let mut full_v = past_v.clone();
full_v.extend_from_slice(&cur_v);
let want = reference(
&q,
&full_k,
&full_v,
b,
h,
h,
sq,
total,
d,
dv,
scale,
true,
past_seq as i64,
|_, _, _, _| 0.0,
);
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, kv_seq, d], &cur_k).view(),
Owned::f32(&[b, h, kv_seq, dv], &cur_v).view(),
absent(),
Owned::f32(&[b, h, past_seq, d], &past_k).view(),
Owned::f32(&[b, h, past_seq, dv], &past_v).view(),
],
&mut [y.view_mut()],
)
.unwrap();
approx(&y.to_f32(), &want, 1e-5);
}
#[test]
fn softcap_changes_output() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let scale = 1.0f32;
let q = [3.0f32, 4.0, -2.0, 1.0];
let k = [2.0f32, 1.0, -1.0, 3.0];
let v = [1.0f32, 0.0, 0.0, 1.0];
let softcap = 2.0f32;
let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
)
.into_iter()
.collect::<Vec<_>>();
let want_capped = {
let mut out = vec![0.0f32; sq * dv];
for i in 0..sq {
let mut scores = [0.0f32; 2];
for (j, sc) in scores.iter_mut().enumerate() {
let mut acc = 0.0f32;
for p in 0..d {
acc += q[i * d + p] * k[j * d + p];
}
let s = acc * scale;
*sc = softcap * (s / softcap).tanh();
}
let max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for sc in scores.iter_mut() {
*sc = (*sc - max).exp();
sum += *sc;
}
for sc in scores.iter_mut() {
*sc /= sum;
}
for c in 0..dv {
let mut acc = 0.0f32;
for (j, &p) in scores.iter().enumerate() {
acc += p * v[j * dv + c];
}
out[i * dv + c] = acc;
}
}
out
};
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, softcap)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want_capped, 1e-5);
assert!(
out.to_f32()
.iter()
.zip(&want)
.any(|(a, b)| (a - b).abs() > 1e-3),
"softcap should change the output"
);
}
#[test]
fn qk_matmul_output_mode0_is_scaled_scores() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let scale = 0.5f32;
let q = [1.0f32, 2.0, 3.0, 4.0];
let k = [1.0f32, 1.0, 2.0, 0.0];
let v = [1.0f32, 0.0, 0.0, 1.0];
let mut expected = [0.0f32; 4];
for i in 0..sq {
for j in 0..sk {
let mut acc = 0.0f32;
for p in 0..d {
acc += q[i * d + p] * k[j * d + p];
}
expected[i * sk + j] = acc * scale;
}
}
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
approx(&qk.to_f32(), &expected, 1e-6);
}
#[test]
fn qk_matmul_output_mode3_is_softmax() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let scale = 0.5f32;
let q = [1.0f32, 0.0, 0.0, 1.0];
let k = [1.0f32, 3.0, 2.0, 4.0];
let v = [1.0f32, 0.0, 0.0, 1.0];
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel(Some(scale), false, None, None, 3, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
let probs = qk.to_f32();
assert!((probs[0] + probs[1] - 1.0).abs() < 1e-6);
assert!((probs[2] + probs[3] - 1.0).abs() < 1e-6);
}
#[test]
fn large_logits_are_numerically_stable() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let scale = 1.0f32;
let q = [100.0f32, 100.0, -100.0, -100.0];
let k = [100.0f32, 100.0, -100.0, -100.0, 50.0, 50.0];
let v = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
assert!(out.to_f32().iter().all(|x| x.is_finite()));
}
#[test]
fn factory_rejects_unsupported_qk_mode() {
use onnx_runtime_ir::{Attribute, NodeId};
let mut node = Node::new(NodeId(0), "Attention", vec![], vec![]);
node.attributes
.insert("qk_matmul_output_mode".to_string(), Attribute::Int(5));
let err = AttentionFactory { since_version: 24 }.create(&node, &[]);
assert!(err.is_err(), "qk_matmul_output_mode=5 must be rejected");
}
#[test]
fn factory_accepts_valid_attributes() {
use onnx_runtime_ir::{Attribute, NodeId};
let mut node = Node::new(NodeId(0), "Attention", vec![], vec![]);
node.attributes
.insert("is_causal".to_string(), Attribute::Int(1));
node.attributes
.insert("scale".to_string(), Attribute::Float(0.25));
node.attributes
.insert("qk_matmul_output_mode".to_string(), Attribute::Int(3));
assert!(
AttentionFactory { since_version: 24 }
.create(&node, &[])
.is_ok()
);
}
#[test]
fn nonpad_kv_seqlen_rejected_for_opset23() {
let (b, h, s, d, dv) = (1, 1, 2, 2, 2);
let q = vec![0.1f32; b * h * s * d];
let k = vec![0.1f32; b * h * s * d];
let v = vec![0.1f32; b * h * s * dv];
let seqlen = [2i64];
let mut out = Owned::zeros_f32(&[b, h, s, dv]);
let err = kernel_v(23, Some(0.5), false, None, None, 0, 0.0).execute(
&[
Owned::f32(&[b, h, s, d], &q).view(),
Owned::f32(&[b, h, s, d], &k).view(),
Owned::f32(&[b, h, s, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &seqlen).view(),
],
&mut [out.view_mut()],
);
assert!(err.is_err(), "nonpad_kv_seqlen must error for opset 23");
}
#[test]
fn nonpad_kv_seqlen_rejected_with_past_cache() {
let (b, h, s, d, dv) = (1, 1, 2, 2, 2);
let q = vec![0.1f32; b * h * s * d];
let k = vec![0.1f32; b * h * s * d];
let v = vec![0.1f32; b * h * s * dv];
let past_k = vec![0.1f32; b * h * s * d];
let past_v = vec![0.1f32; b * h * s * dv];
let seqlen = [2i64];
let mut out = Owned::zeros_f32(&[b, h, s, dv]);
let err = kernel_v(24, Some(0.5), false, None, None, 0, 0.0).execute(
&[
Owned::f32(&[b, h, s, d], &q).view(),
Owned::f32(&[b, h, s, d], &k).view(),
Owned::f32(&[b, h, s, dv], &v).view(),
absent(),
Owned::f32(&[b, h, s, d], &past_k).view(),
Owned::f32(&[b, h, s, dv], &past_v).view(),
Owned::i64(&[b], &seqlen).view(),
],
&mut [out.view_mut()],
);
assert!(
err.is_err(),
"nonpad_kv_seqlen with past_key/past_value must error"
);
}
#[test]
fn non_divisible_gqa_errors() {
let (b, qh, kvh, s, d, dv) = (1, 3, 2, 2, 2, 2);
let q = vec![0.1f32; b * qh * s * d];
let k = vec![0.1f32; b * kvh * s * d];
let v = vec![0.1f32; b * kvh * s * dv];
let mut out = Owned::zeros_f32(&[b, qh, s, dv]);
let err = kernel(Some(0.5), false, None, None, 0, 0.0).execute(
&[
Owned::f32(&[b, qh, s, d], &q).view(),
Owned::f32(&[b, kvh, s, d], &k).view(),
Owned::f32(&[b, kvh, s, dv], &v).view(),
],
&mut [out.view_mut()],
);
assert!(err.is_err(), "non-divisible GQA must error");
}
#[test]
fn sqrt_scale_avoids_overflow() {
let (b, h, sq, sk, d, dv) = (1, 1, 1, 1, 1, 1);
let q = [1e30f32];
let k = [1e30f32];
let v = [7.0f32];
let scale = 1e-30f32;
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
assert!(!(1e30f32 * 1e30f32).is_finite());
let score = qk.to_f32()[0];
assert!(score.is_finite(), "score must be finite, got {score}");
assert!((score - 1e30).abs() < 1e26, "score ~1e30, got {score}");
approx(&y.to_f32(), &v, 1e-3);
}
#[test]
fn scalar_bool_false_mask_zeros_all_rows() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let q = [1.0f32, 2.0, 3.0, 4.0];
let k = [1.0f32, 0.0, 0.0, 1.0];
let v = [5.0f32, 6.0, 7.0, 8.0];
let mask = [false];
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(0.5), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::bool_(&[], &mask).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let got = out.to_f32();
assert!(
got.iter().all(|x| *x == 0.0),
"scalar false → zero: {got:?}"
);
}
#[test]
fn scalar_bool_true_mask_is_noop() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.2).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.1).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.3).collect();
let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
let mask = [true];
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::bool_(&[], &mask).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
}
#[test]
fn scalar_float_mask_adds_bias() {
let (b, h, sq, sk, d, dv) = (1, 1, 1, 2, 2, 2);
let scale = 0.5f32;
let q = [1.0f32, 2.0];
let k = [1.0f32, 0.0, 0.0, 1.0];
let v = [1.0f32, 0.0, 0.0, 1.0];
let bias = [3.0f32];
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel(Some(scale), false, None, None, 2, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::f32(&[], &bias).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
let mut expected = [0.0f32; 2];
for (j, e) in expected.iter_mut().enumerate() {
let mut acc = 0.0f32;
for p in 0..d {
acc += q[p] * k[j * d + p];
}
*e = acc * scale + 3.0;
}
approx(&qk.to_f32(), &expected, 1e-4);
}
#[test]
fn negative_softcap_is_applied() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let scale = 1.0f32;
let q = [3.0f32, 4.0, -2.0, 1.0];
let k = [2.0f32, 1.0, -1.0, 3.0];
let v = [1.0f32, 0.0, 0.0, 1.0];
let softcap = -2.0f32;
let want = {
let mut out = vec![0.0f32; sq * dv];
for i in 0..sq {
let mut scores = [0.0f32; 2];
for (j, sc) in scores.iter_mut().enumerate() {
let mut acc = 0.0f32;
for p in 0..d {
acc += q[i * d + p] * k[j * d + p];
}
let s = acc * scale;
*sc = softcap * (s / softcap).tanh();
}
let max = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for sc in scores.iter_mut() {
*sc = (*sc - max).exp();
sum += *sc;
}
for sc in scores.iter_mut() {
*sc /= sum;
}
for c in 0..dv {
let mut acc = 0.0f32;
for (j, &p) in scores.iter().enumerate() {
acc += p * v[j * dv + c];
}
out[i * dv + c] = acc;
}
}
out
};
let want_nocap = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(scale), false, None, None, 0, softcap)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
assert!(
out.to_f32()
.iter()
.zip(&want_nocap)
.any(|(a, b)| (a - b).abs() > 1e-3),
"negative softcap must change the output"
);
}
#[test]
fn qk_modes_are_version_independent() {
let (b, h, sq, sk, d, dv) = (1, 1, 1, 2, 2, 2);
let scale = 0.5f32;
let softcap = 3.0f32;
let q = [1.0f32, 2.0];
let k = [1.0f32, 1.0, 2.0, 0.0];
let v = [1.0f32, 0.0, 0.0, 1.0];
let maskv = [0.5f32, -1.0];
let scaled: Vec<f32> = (0..sk)
.map(|j| {
let mut acc = 0.0f32;
for p in 0..d {
acc += q[p] * k[j * d + p];
}
acc * scale
})
.collect();
let capped: Vec<f32> = scaled
.iter()
.map(|s| softcap * (s / softcap).tanh())
.collect();
let masked: Vec<f32> = capped.iter().zip(&maskv).map(|(s, m)| s + m).collect();
let softmaxed = {
let max = masked.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut e: Vec<f32> = masked.iter().map(|s| (s - max).exp()).collect();
let sum: f32 = e.iter().sum();
for x in e.iter_mut() {
*x /= sum;
}
e
};
let run = |ver: u32, mode: i64| {
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel_v(ver, Some(scale), false, None, None, mode, softcap)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::f32(&[sq, sk], &maskv).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
qk.to_f32()
};
approx(&run(24, 0), &scaled, 1e-5);
approx(&run(23, 0), &scaled, 1e-5);
approx(&run(24, 3), &softmaxed, 1e-5);
approx(&run(23, 3), &softmaxed, 1e-5);
approx(&run(24, 1), &capped, 1e-5);
approx(&run(24, 2), &masked, 1e-5);
approx(&run(23, 1), &capped, 1e-5);
approx(&run(23, 2), &masked, 1e-5);
}
#[test]
fn nonpad_kv_seqlen_offset_positive() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 4, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.1).collect();
let nonpad = [4i64];
let offset = nonpad[0] - sq as i64; let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
true,
offset,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel_v(24, Some(scale), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
}
#[test]
fn nonpad_kv_seqlen_offset_zero_is_lower_triangular() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.1).collect();
let nonpad = [2i64];
let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
true,
0,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel_v(24, Some(scale), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-5);
}
#[test]
fn nonpad_kv_seqlen_offset_negative_zeros_leading_rows() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 2, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3 + 0.1).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.1 + 0.2).collect();
let nonpad = [1i64];
let offset = nonpad[0] - sq as i64; let want = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
true,
offset,
|_, _, _, _| 0.0,
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel_v(24, Some(scale), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let got = out.to_f32();
assert!(got.iter().all(|x| x.is_finite()), "no NaN/inf: {got:?}");
approx(&got[0..dv], &[0.0, 0.0], 1e-6);
approx(&got, &want, 1e-5);
}
#[test]
fn nonpad_kv_seqlen_masks_padding_for_noncausal() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 4, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.1).collect();
let nonpad = [2i64];
let pad_bias = |_b: usize, _h: usize, _i: usize, j: usize| {
if (j as i64) >= nonpad[0] {
f32::NEG_INFINITY
} else {
0.0
}
};
let want = reference(
&q, &k, &v, b, h, h, sq, sk, d, dv, scale, false, 0, pad_bias,
);
let no_mask = reference(
&q,
&k,
&v,
b,
h,
h,
sq,
sk,
d,
dv,
scale,
false,
0,
|_, _, _, _| 0.0,
);
assert!(
want.iter().zip(&no_mask).any(|(a, c)| (a - c).abs() > 1e-4),
"padding mask must differ from the unmasked result"
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel_v(24, Some(scale), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let got = out.to_f32();
assert!(got.iter().all(|x| x.is_finite()), "no NaN/inf: {got:?}");
approx(&got, &want, 1e-5);
}
#[test]
fn nonpad_kv_seqlen_causal_and_padding_compose() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 4, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.1).collect();
let nonpad = [3i64];
let offset = nonpad[0] - sq as i64; let pad_bias = |_b: usize, _h: usize, _i: usize, j: usize| {
if (j as i64) >= nonpad[0] {
f32::NEG_INFINITY
} else {
0.0
}
};
let want = reference(
&q, &k, &v, b, h, h, sq, sk, d, dv, scale, true, offset, pad_bias,
);
let mut out = Owned::zeros_f32(&[b, h, sq, dv]);
kernel_v(24, Some(scale), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [out.view_mut()],
)
.unwrap();
let got = out.to_f32();
assert!(got.iter().all(|x| x.is_finite()), "no NaN/inf: {got:?}");
approx(&got, &want, 1e-5);
}
#[test]
fn nonpad_kv_seqlen_qk_output_reflects_padding_mask() {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let scale = 0.5f32;
let q: Vec<f32> = (0..b * h * sq * d)
.map(|i| (i as f32 * 0.3).sin())
.collect();
let k: Vec<f32> = (0..b * h * sk * d)
.map(|i| (i as f32 * 0.2).cos())
.collect();
let v: Vec<f32> = (0..b * h * sk * dv).map(|i| i as f32 * 0.1).collect();
let nonpad = [2i64];
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel_v(24, Some(scale), false, None, None, 2, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
let masked = qk.to_f32();
for i in 0..sq {
assert_eq!(
masked[i * sk + 2],
f32::NEG_INFINITY,
"mode-2 padding column must be -inf"
);
assert!(masked[i * sk].is_finite() && masked[i * sk + 1].is_finite());
}
let mut y3 = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk3 = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv3 = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk3 = Owned::zeros_f32(&[b, h, sq, sk]);
kernel_v(24, Some(scale), false, None, None, 3, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &nonpad).view(),
],
&mut [
y3.view_mut(),
pk3.view_mut(),
pv3.view_mut(),
qk3.view_mut(),
],
)
.unwrap();
let probs = qk3.to_f32();
for i in 0..sq {
assert_eq!(probs[i * sk + 2], 0.0, "mode-3 padding prob must be 0");
let s = probs[i * sk] + probs[i * sk + 1];
assert!((s - 1.0).abs() < 1e-6, "valid probs must sum to 1: {s}");
}
}
#[test]
#[ignore]
fn byte_identical_capture() {
use std::fmt::Write as _;
let path = std::env::var("ATTN_GOLDEN_DUMP")
.unwrap_or_else(|_| "attn_golden_dump.txt".to_string());
let mut out = String::new();
let mut dump = |name: &str, outs: &[Owned]| {
let _ = write!(out, "CASE {name}");
for o in outs {
let _ = write!(out, " |");
for x in o.to_f32() {
let _ = write!(out, " {:08x}", x.to_bits());
}
}
out.push('\n');
};
let mut idx = 0usize;
for &rank4 in &[true, false] {
for &causal in &[false, true] {
for &softcap in &[0.0f32, 20.0, -15.0] {
for qk_mode in 0..=3i64 {
for &(qh, kvh) in &[(2usize, 2usize), (4, 2), (2, 1)] {
idx += 1;
let (b, sq, sk, d, dv) = (2usize, 3usize, 3usize, 4usize, 4usize);
let q = conformance_values(b * qh * sq * d, idx);
let k = conformance_values(b * kvh * sk * d, idx + 1);
let v = conformance_values(b * kvh * sk * dv, idx + 2);
let want_qk = true;
let ker =
kernel(Some(0.3), causal, Some(qh), Some(kvh), qk_mode, softcap);
let mut y;
let mut pk = Owned::zeros_f32(&[b, kvh, sk, d]);
let mut pv = Owned::zeros_f32(&[b, kvh, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, qh, sq, sk]);
let (qi, ki, vi);
if rank4 {
y = Owned::zeros_f32(&[b, qh, sq, dv]);
qi = Owned::f32(&[b, qh, sq, d], &q);
ki = Owned::f32(&[b, kvh, sk, d], &k);
vi = Owned::f32(&[b, kvh, sk, dv], &v);
} else {
y = Owned::zeros_f32(&[b, sq, qh * dv]);
let q3 = bhsd_to_bsh(&q, b, qh, sq, d);
let k3 = bhsd_to_bsh(&k, b, kvh, sk, d);
let v3 = bhsd_to_bsh(&v, b, kvh, sk, dv);
qi = Owned::f32(&[b, sq, qh * d], &q3);
ki = Owned::f32(&[b, sk, kvh * d], &k3);
vi = Owned::f32(&[b, sk, kvh * dv], &v3);
}
let _ = want_qk;
ker.execute(
&[qi.view(), ki.view(), vi.view()],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
dump(
&format!(
"sweep r4={rank4} causal={causal} softcap={softcap} qk={qk_mode} qh={qh} kvh={kvh}"
),
&[y, pk, pv, qk],
);
}
}
}
}
}
{
let (b, h, sq, sk, d, dv) = (1, 2, 2, 3, 3, 2);
let q = conformance_values(b * h * sq * d, 7);
let k = conformance_values(b * h * sk * d, 8);
let v = conformance_values(b * h * sk * dv, 9);
let mask = conformance_values(sq * sk, 11);
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(0.4), false, None, None, 2, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::f32(&[1, 1, sq, sk], &mask).view(),
],
&mut [y.view_mut()],
)
.unwrap();
dump("float_mask", &[y]);
}
{
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let q = conformance_values(b * h * sq * d, 13);
let k = conformance_values(b * h * sk * d, 14);
let v = conformance_values(b * h * sk * dv, 15);
let mask = [true, false, false, false, false, false];
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
kernel(Some(0.5), false, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
Owned::bool_(&[b, sq, sk], &mask).view(),
],
&mut [y.view_mut()],
)
.unwrap();
dump("bool_mask_fully_masked_row", &[y]);
}
{
let (b, h, sq, past, cur, d, dv) = (1, 2, 2, 2, 2, 3, 3);
let q = conformance_values(b * h * sq * d, 21);
let pk = conformance_values(b * h * past * d, 22);
let pv = conformance_values(b * h * past * dv, 23);
let ck = conformance_values(b * h * cur * d, 24);
let cv = conformance_values(b * h * cur * dv, 25);
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut prk = Owned::zeros_f32(&[b, h, past + cur, d]);
let mut prv = Owned::zeros_f32(&[b, h, past + cur, dv]);
kernel(Some(0.3), true, None, None, 0, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, cur, d], &ck).view(),
Owned::f32(&[b, h, cur, dv], &cv).view(),
absent(),
Owned::f32(&[b, h, past, d], &pk).view(),
Owned::f32(&[b, h, past, dv], &pv).view(),
],
&mut [y.view_mut(), prk.view_mut(), prv.view_mut()],
)
.unwrap();
dump("kv_cache", &[y, prk, prv]);
}
for &(off_name, nonpad) in &[("pos", 3i64), ("zero", 2), ("neg", 1)] {
for qk_mode in [2i64, 3] {
let (b, h, sq, sk, d, dv) = (1, 1, 2, 3, 2, 2);
let q = conformance_values(b * h * sq * d, 31);
let k = conformance_values(b * h * sk * d, 32);
let v = conformance_values(b * h * sk * dv, 33);
let seq = [nonpad];
let mut y = Owned::zeros_f32(&[b, h, sq, dv]);
let mut pk = Owned::zeros_f32(&[b, h, sk, d]);
let mut pv = Owned::zeros_f32(&[b, h, sk, dv]);
let mut qk = Owned::zeros_f32(&[b, h, sq, sk]);
kernel_v(24, Some(0.5), true, None, None, qk_mode, 0.0)
.execute(
&[
Owned::f32(&[b, h, sq, d], &q).view(),
Owned::f32(&[b, h, sk, d], &k).view(),
Owned::f32(&[b, h, sk, dv], &v).view(),
absent(),
absent(),
absent(),
Owned::i64(&[b], &seq).view(),
],
&mut [y.view_mut(), pk.view_mut(), pv.view_mut(), qk.view_mut()],
)
.unwrap();
dump(&format!("nonpad_{off_name}_qk{qk_mode}"), &[y, pk, pv, qk]);
}
}
std::fs::write(&path, out).unwrap();
eprintln!("wrote golden dump to {path}");
}
#[test]
fn attention_bf16_matches_widened_f32_reference() {
let q = Owned::bf16(&[1, 1, 2, 2], &[1., -1., 0., 2.]);
let k = Owned::bf16(&[1, 1, 2, 2], &[1., 0., 0., 1.]);
let v = Owned::bf16(&[1, 1, 2, 2], &[2., -1., -2., 3.]);
let mut out = Owned::zeros(DataType::BFloat16, &[1, 1, 2, 2]);
kernel(Some(1.), false, Some(1), Some(1), 0, 0.)
.execute(&[q.view(), k.view(), v.view()], &mut [out.view_mut()])
.unwrap();
let expected = reference(
&q.to_bf16_as_f32(),
&k.to_bf16_as_f32(),
&v.to_bf16_as_f32(),
1,
1,
1,
2,
2,
2,
2,
1.,
false,
0,
|_, _, _, _| 0.,
);
let expected: Vec<_> = expected
.into_iter()
.map(half::bf16::from_f32)
.map(half::bf16::to_f32)
.collect();
assert_eq!(out.to_bf16_as_f32(), expected);
}
}