use onnx_runtime_ep_api::{
DeviceId, EpError, Kernel, KernelFactory, Result, TensorMut, TensorView,
};
use onnx_runtime_ir::{DataType, Node, broadcast_shapes, compute_contiguous_strides};
use super::add::broadcast_apply;
use super::matmul::matmul_dense;
use super::softmax::softmax_slices;
use super::{check_arity, to_dense_f32, write_dense_f32};
use crate::strided::numel;
pub struct FusedAttentionKernel {
scale: f32,
k_transposed: bool,
}
pub struct FusedAttentionFactory;
impl KernelFactory for FusedAttentionFactory {
fn create(&self, node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
let scale = node
.attr("scale")
.and_then(|a| a.as_float())
.ok_or_else(|| {
EpError::KernelFailed("FusedAttention: missing f32 `scale` attribute".into())
})?;
let k_transposed = node
.attr("k_transposed")
.and_then(|a| a.as_int())
.unwrap_or(0)
!= 0;
Ok(Box::new(FusedAttentionKernel {
scale,
k_transposed,
}))
}
}
fn matmul_result_shape(a: &[usize], b: &[usize], stage: &str) -> Result<Vec<usize>> {
if a.len() < 2 || b.len() < 2 {
return Err(EpError::KernelFailed(format!(
"FusedAttention: {stage} operands must be rank ≥ 2 (got {a:?}, {b:?})"
)));
}
let (m, ka) = (a[a.len() - 2], a[a.len() - 1]);
let (kb, n) = (b[b.len() - 2], b[b.len() - 1]);
if ka != kb {
return Err(EpError::KernelFailed(format!(
"FusedAttention: {stage} contraction mismatch ({ka} vs {kb})"
)));
}
let mut shape = broadcast_shapes(&a[..a.len() - 2], &b[..b.len() - 2]).map_err(EpError::Ir)?;
shape.push(m);
shape.push(n);
Ok(shape)
}
impl Kernel for FusedAttentionKernel {
fn execute(&self, inputs: &[TensorView], outputs: &mut [TensorMut]) -> Result<()> {
check_arity("FusedAttention", inputs, outputs, 3, 4, 1)?;
let q = &inputs[0];
let k = &inputs[1];
let v = &inputs[2];
let has_mask = inputs.len() == 4;
if q.shape.len() < 2 || k.shape.len() < 2 || v.shape.len() < 2 {
return Err(EpError::KernelFailed(
"FusedAttention: Q, K, V must each be rank ≥ 2".into(),
));
}
let (mut scores, scores_shape) = if self.k_transposed {
let s = matmul_dense(q, k)?;
let shape = matmul_result_shape(q.shape, k.shape, "Q·K")?;
(s, shape)
} else {
let rank = k.shape.len();
let mut kt_shape = k.shape.to_vec();
kt_shape.swap(rank - 2, rank - 1);
let mut kt_strides = k.strides.to_vec();
kt_strides.swap(rank - 2, rank - 1);
let kt = TensorView::new(k.data, k.dtype, &kt_shape, &kt_strides, k.device)
.with_byte_offset(k.byte_offset);
let s = matmul_dense(q, &kt)?;
let shape = matmul_result_shape(q.shape, &kt_shape, "Q·Kᵀ")?;
(s, shape)
};
for s in &mut scores {
*s *= self.scale;
}
if has_mask {
let mask = to_dense_f32(&inputs[3])?;
let mask_shape = inputs[3].shape;
broadcast_apply(&mask, mask_shape, &scores_shape, |i, val| scores[i] += val)?;
}
let seq_k = *scores_shape.last().unwrap();
let n = numel(&scores_shape);
let outer = n.checked_div(seq_k).unwrap_or(0);
let mut probs = vec![0.0f32; n];
softmax_slices(&scores, &mut probs, outer, seq_k, 1);
let probs_strides = compute_contiguous_strides(&scores_shape);
let probs_view = TensorView::new(
onnx_runtime_ep_api::DevicePtr(probs.as_ptr() as *const std::ffi::c_void),
DataType::Float32,
&scores_shape,
&probs_strides,
DeviceId::cpu(),
);
let out = matmul_dense(&probs_view, v)?;
write_dense_f32(&mut outputs[0], &out)
}
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],
sq: usize,
d: usize,
k: &[f32],
sk: usize,
v: &[f32],
dv: usize,
scale: f32,
mask: Option<&[f32]>,
) -> Vec<f32> {
let mut scores = vec![0.0f32; sq * sk];
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];
}
scores[i * sk + j] = acc * scale + mask.map_or(0.0, |m| m[i * sk + j]);
}
}
for i in 0..sq {
let row = &mut scores[i * sk..i * sk + sk];
let max = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for x in row.iter_mut() {
*x = (*x - max).exp();
sum += *x;
}
for x in row.iter_mut() {
*x /= sum;
}
}
let mut out = vec![0.0f32; sq * dv];
for i in 0..sq {
for c in 0..dv {
let mut acc = 0.0f32;
for j in 0..sk {
acc += scores[i * sk + j] * v[j * dv + c];
}
out[i * dv + c] = acc;
}
}
out
}
fn approx(a: &[f32], b: &[f32], atol: f32) {
assert_eq!(a.len(), b.len(), "length mismatch");
for (i, (x, y)) in a.iter().zip(b).enumerate() {
assert!(
(x - y).abs() < atol,
"element {i}: {x} vs {y} ({a:?} vs {b:?})"
);
}
}
#[test]
fn sdpa_unmasked_pretransposed_k_matches_reference() {
let q = [1.0f32, 0.0, -1.0, 0.5, 2.0, 1.0];
let k_natural = [1.0f32, 2.0, 0.0, -1.0, 1.0, 3.0]; let v = [1.0f32, 0.0, 0.0, 2.0]; let scale = 0.5f32;
let mut kt = vec![0.0f32; 3 * 2];
for j in 0..2 {
for p in 0..3 {
kt[p * 2 + j] = k_natural[j * 3 + p];
}
}
let want = reference(&q, 2, 3, &k_natural, 2, &v, 2, scale, None);
let qv = Owned::f32(&[2, 3], &q);
let kv = Owned::f32(&[3, 2], &kt);
let vv = Owned::f32(&[2, 2], &v);
let mut out = Owned::zeros_f32(&[2, 2]);
FusedAttentionKernel {
scale,
k_transposed: true,
}
.execute(&[qv.view(), kv.view(), vv.view()], &mut [out.view_mut()])
.unwrap();
approx(&out.to_f32(), &want, 1e-6);
}
#[test]
fn sdpa_unmasked_internal_transpose_k_matches_reference() {
let q = [1.0f32, 0.0, -1.0, 0.5, 2.0, 1.0];
let k_natural = [1.0f32, 2.0, 0.0, -1.0, 1.0, 3.0]; let v = [1.0f32, 0.0, 0.0, 2.0];
let scale = 0.5f32;
let want = reference(&q, 2, 3, &k_natural, 2, &v, 2, scale, None);
let qv = Owned::f32(&[2, 3], &q);
let kv = Owned::f32(&[2, 3], &k_natural);
let vv = Owned::f32(&[2, 2], &v);
let mut out = Owned::zeros_f32(&[2, 2]);
FusedAttentionKernel {
scale,
k_transposed: false,
}
.execute(&[qv.view(), kv.view(), vv.view()], &mut [out.view_mut()])
.unwrap();
approx(&out.to_f32(), &want, 1e-6);
}
#[test]
fn sdpa_masked_matches_reference() {
let q = [1.0f32, 2.0, -1.0, 0.5];
let k_natural = [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 scale = 0.7f32;
let want = reference(&q, 2, 2, &k_natural, 3, &v, 2, scale, Some(&mask));
let mut kt = vec![0.0f32; 2 * 3];
for j in 0..3 {
for p in 0..2 {
kt[p * 3 + j] = k_natural[j * 2 + p];
}
}
let qv = Owned::f32(&[2, 2], &q);
let kv = Owned::f32(&[2, 3], &kt);
let vv = Owned::f32(&[3, 2], &v);
let mv = Owned::f32(&[2, 3], &mask);
let mut out = Owned::zeros_f32(&[2, 2]);
FusedAttentionKernel {
scale,
k_transposed: true,
}
.execute(
&[qv.view(), kv.view(), vv.view(), mv.view()],
&mut [out.view_mut()],
)
.unwrap();
approx(&out.to_f32(), &want, 1e-6);
}
#[test]
fn sdpa_batched_leading_dims() {
let sq = 2;
let d = 2;
let sk = 2;
let dv = 2;
let scale = 0.25f32;
let q = [
1.0f32, 0.0, 0.0, 1.0, 2.0, 1.0, -1.0, 0.5, ];
let k_nat = [
1.0f32, 1.0, 0.0, 2.0, 0.5, -1.0, 1.0, 1.0, ];
let v = [
1.0f32, 0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 5.0, ];
let mut kt = vec![0.0f32; q.len()];
for b in 0..2 {
for j in 0..sk {
for p in 0..d {
kt[b * d * sk + p * sk + j] = k_nat[b * sk * d + j * d + p];
}
}
}
let qv = Owned::f32(&[2, 1, sq, d], &q);
let kv = Owned::f32(&[2, 1, d, sk], &kt);
let vv = Owned::f32(&[2, 1, sk, dv], &v);
let mut out = Owned::zeros_f32(&[2, 1, sq, dv]);
FusedAttentionKernel {
scale,
k_transposed: true,
}
.execute(&[qv.view(), kv.view(), vv.view()], &mut [out.view_mut()])
.unwrap();
let got = out.to_f32();
for b in 0..2usize {
let want = reference(
&q[b * sq * d..(b + 1) * sq * d],
sq,
d,
&k_nat[b * sk * d..(b + 1) * sk * d],
sk,
&v[b * sk * dv..(b + 1) * sk * dv],
dv,
scale,
None,
);
approx(&got[b * sq * dv..(b + 1) * sq * dv], &want, 1e-6);
}
}
#[test]
fn sdpa_softmax_stage_matches_row_softmax() {
let q = [1.0f32, 0.0, 0.0, 1.0]; let kt = [1.0f32, 3.0, 2.0, 4.0]; let v = [1.0f32, 0.0, 0.0, 1.0]; let mut out = Owned::zeros_f32(&[2, 2]);
FusedAttentionKernel {
scale: 1.0,
k_transposed: true,
}
.execute(
&[
Owned::f32(&[2, 2], &q).view(),
Owned::f32(&[2, 2], &kt).view(),
Owned::f32(&[2, 2], &v).view(),
],
&mut [out.view_mut()],
)
.unwrap();
approx(
&out.to_f32(),
&[0.119_202_92, 0.880_797_1, 0.119_202_92, 0.880_797_1],
1e-6,
);
let r = out.to_f32();
assert!((r[0] + r[1] - 1.0).abs() < 1e-6);
assert!((r[2] + r[3] - 1.0).abs() < 1e-6);
}
}