pub struct SdpaTensors<'a> {
pub q: &'a [f32],
pub k: &'a [f32],
pub v: &'a [f32],
pub batch: usize,
pub num_heads: usize,
pub num_kv_heads: usize,
pub q_seq: usize,
pub kv_seq: usize,
pub head_size: usize,
pub v_head_size: usize,
}
#[derive(Clone, Copy, Debug)]
pub enum ScaleMode {
PostDot(f32),
SplitSqrt(f32),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SoftmaxExp {
F32,
F64Intermediate,
}
impl SoftmaxExp {
#[inline]
fn exp(self, value: f32) -> f32 {
match self {
Self::F32 => value.exp(),
Self::F64Intermediate => (value as f64).exp() as f32,
}
}
}
pub struct SdpaConfig {
pub scale: ScaleMode,
pub softcap: Option<f32>,
pub causal: bool,
pub past_seq: usize,
pub causal_fill: f32,
}
pub trait AttnBias: Sync {
fn at(&self, b: usize, head: usize, i: usize, j: usize) -> f32;
}
pub trait KeyMask: Sync {
fn at(&self, b: usize, i: usize, j: usize) -> f32;
}
pub struct NoBias;
impl AttnBias for NoBias {
#[inline]
fn at(&self, _b: usize, _head: usize, _i: usize, _j: usize) -> f32 {
0.0
}
}
pub struct NoMask;
impl KeyMask for NoMask {
#[inline]
fn at(&self, _b: usize, _i: usize, _j: usize) -> f32 {
0.0
}
}
pub struct BroadcastBias<'a> {
data: &'a [f32],
dims: [usize; 4],
}
impl<'a> BroadcastBias<'a> {
pub fn new(data: &'a [f32], dims: [usize; 4]) -> Self {
Self { data, dims }
}
}
impl AttnBias for BroadcastBias<'_> {
#[inline]
fn at(&self, b: usize, head: usize, i: usize, j: usize) -> f32 {
let b0 = if self.dims[0] == 1 { 0 } else { b };
let n0 = if self.dims[1] == 1 { 0 } else { head };
let off = (((b0 * self.dims[1] + n0) * self.dims[2] + i) * self.dims[3]) + j;
self.data[off]
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum QkCaptureStage {
PostScale,
PostSoftcap,
PreSoftmax,
PostSoftmax,
}
pub struct QkCapture<'a> {
pub scores: &'a mut [f32],
pub stage: QkCaptureStage,
}
pub fn sdpa_f32(
t: &SdpaTensors,
cfg: &SdpaConfig,
bias: &dyn AttnBias,
mask: &dyn KeyMask,
y: &mut [f32],
qk: Option<QkCapture>,
) {
#[cfg(feature = "mlas")]
{
let non_empty = t.batch > 0
&& t.num_heads > 0
&& t.q_seq > 0
&& t.kv_seq > 0
&& t.head_size > 0
&& t.v_head_size > 0;
if qk.is_none() && non_empty {
sdpa_f32_fast(t, cfg, bias, mask, y);
return;
}
}
sdpa_f32_scalar(t, cfg, bias, mask, y, qk);
}
pub fn sdpa_decode_row(
q: &[f32],
k: &[f32],
v: &[f32],
kv_seq: usize,
lo: usize,
hi: usize,
scale: f32,
softcap: Option<f32>,
exp: SoftmaxExp,
output: &mut [f32],
) {
debug_assert!(lo <= hi && hi <= kv_seq);
debug_assert_eq!(k.len(), kv_seq * q.len());
debug_assert_eq!(v.len(), kv_seq * output.len());
let mut scores = vec![0.0f32; hi - lo];
for (i, ks) in (lo..hi).enumerate() {
let k_base = ks * q.len();
let mut score = dot_f32(q, &k[k_base..k_base + q.len()]);
score *= scale;
if let Some(softcap) = softcap {
score = softcap * (score / softcap).tanh();
}
scores[i] = score;
}
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for score in &mut scores {
*score = exp.exp(*score - max);
sum += *score;
}
if sum > 0.0 {
for score in &mut scores {
*score /= sum;
}
}
output.fill(0.0);
for (i, ks) in (lo..hi).enumerate() {
let probability = scores[i];
if probability == 0.0 {
continue;
}
let v_base = ks * output.len();
axpy_f32(output, probability, &v[v_base..v_base + output.len()]);
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct DecodePartial {
pub max: f64,
pub sum: f64,
}
pub fn sdpa_decode_partial(
q: &[f32],
k: &[f32],
v: &[f32],
kv_seq: usize,
lo: usize,
hi: usize,
scale: f32,
softcap: Option<f32>,
partial_output: &mut [f64],
) -> DecodePartial {
debug_assert!(lo <= hi && hi <= kv_seq);
debug_assert_eq!(k.len(), kv_seq * q.len());
debug_assert_eq!(v.len(), kv_seq * partial_output.len());
partial_output.fill(0.0);
if hi <= lo {
return DecodePartial {
max: f64::NEG_INFINITY,
sum: 0.0,
};
}
let mut scores = vec![0.0f32; hi - lo];
for (i, ks) in (lo..hi).enumerate() {
let k_base = ks * q.len();
let mut score = dot_f32(q, &k[k_base..k_base + q.len()]);
score *= scale;
if let Some(softcap) = softcap {
score = softcap * (score / softcap).tanh();
}
scores[i] = score;
}
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let max_f64 = max as f64;
let mut sum = 0.0f64;
for (i, ks) in (lo..hi).enumerate() {
let weight = ((scores[i] as f64) - max_f64).exp();
sum += weight;
let v_base = ks * partial_output.len();
let v_row = &v[v_base..v_base + partial_output.len()];
for (o, &value) in partial_output.iter_mut().zip(v_row) {
*o += weight * value as f64;
}
}
DecodePartial { max: max_f64, sum }
}
pub fn combine_decode_partials(
partials: &[DecodePartial],
partial_outputs: &[f64],
v_head_size: usize,
output: &mut [f32],
) {
debug_assert_eq!(output.len(), v_head_size);
debug_assert_eq!(partial_outputs.len(), partials.len() * v_head_size);
let global_max = partials
.iter()
.map(|partial| partial.max)
.fold(f64::NEG_INFINITY, f64::max);
if global_max == f64::NEG_INFINITY {
output.fill(0.0);
return;
}
let mut denominator = 0.0f64;
let mut accumulator = vec![0.0f64; v_head_size];
for (chunk, partial) in partials.iter().enumerate() {
if partial.max == f64::NEG_INFINITY {
continue;
}
let rescale = (partial.max - global_max).exp();
denominator += rescale * partial.sum;
let base = chunk * v_head_size;
let chunk_output = &partial_outputs[base..base + v_head_size];
for (acc, &value) in accumulator.iter_mut().zip(chunk_output) {
*acc += rescale * value;
}
}
if denominator > 0.0 {
let inverse = 1.0 / denominator;
for (out, &acc) in output.iter_mut().zip(&accumulator) {
*out = (acc * inverse) as f32;
}
} else {
output.fill(0.0);
}
}
#[inline]
fn softmax_in_place(scores: &mut [f32], exp: SoftmaxExp) {
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
if max == f32::NEG_INFINITY {
scores.fill(0.0);
return;
}
let mut sum = 0.0f32;
for score in scores.iter_mut() {
let e = exp.exp(*score - max);
*score = e;
sum += e;
}
let inv = 1.0 / sum;
for score in scores.iter_mut() {
*score *= inv;
}
}
#[inline(always)]
fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
if crate::backend::has_simd_x86() {
return unsafe { dot_avx2_fma(a, b) };
}
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
#[inline(always)]
fn axpy_f32(dst: &mut [f32], scalar: f32, src: &[f32]) {
debug_assert_eq!(dst.len(), src.len());
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
if crate::backend::has_simd_x86() {
unsafe { axpy_avx2_fma(dst, scalar, src) };
return;
}
for (d, s) in dst.iter_mut().zip(src) {
*d += scalar * s;
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_avx2_fma(a: &[f32], b: &[f32]) -> f32 {
#[cfg(target_arch = "x86")]
use std::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
let n = a.len();
let a_ptr = a.as_ptr();
let b_ptr = b.as_ptr();
unsafe {
let mut acc0 = _mm256_setzero_ps();
let mut acc1 = _mm256_setzero_ps();
let chunks16 = n / 16;
for i in 0..chunks16 {
let av0 = _mm256_loadu_ps(a_ptr.add(i * 16));
let bv0 = _mm256_loadu_ps(b_ptr.add(i * 16));
acc0 = _mm256_fmadd_ps(av0, bv0, acc0);
let av1 = _mm256_loadu_ps(a_ptr.add(i * 16 + 8));
let bv1 = _mm256_loadu_ps(b_ptr.add(i * 16 + 8));
acc1 = _mm256_fmadd_ps(av1, bv1, acc1);
}
let mut tail = chunks16 * 16;
if tail + 8 <= n {
let av = _mm256_loadu_ps(a_ptr.add(tail));
let bv = _mm256_loadu_ps(b_ptr.add(tail));
acc0 = _mm256_fmadd_ps(av, bv, acc0);
tail += 8;
}
let acc = _mm256_add_ps(acc0, acc1);
let lo = _mm256_extractf128_ps(acc, 0);
let hi = _mm256_extractf128_ps(acc, 1);
let v4 = _mm_add_ps(lo, hi);
let shuf = _mm_movehdup_ps(v4);
let v2 = _mm_add_ps(v4, shuf);
let shuf2 = _mm_movehl_ps(shuf, v2);
let v1 = _mm_add_ss(v2, shuf2);
let mut result = _mm_cvtss_f32(v1);
for i in tail..n {
result += *a_ptr.add(i) * *b_ptr.add(i);
}
result
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2,fma")]
unsafe fn axpy_avx2_fma(dst: &mut [f32], scalar: f32, src: &[f32]) {
#[cfg(target_arch = "x86")]
use std::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
let n = dst.len();
let s = _mm256_set1_ps(scalar);
let dst_ptr = dst.as_mut_ptr();
let src_ptr = src.as_ptr();
unsafe {
let mut i = 0;
while i + 8 <= n {
let d = _mm256_loadu_ps(dst_ptr.add(i));
let x = _mm256_loadu_ps(src_ptr.add(i));
_mm256_storeu_ps(dst_ptr.add(i), _mm256_fmadd_ps(s, x, d));
i += 8;
}
while i < n {
*dst_ptr.add(i) += scalar * *src_ptr.add(i);
i += 1;
}
}
}
pub fn sdpa_f32_scalar(
t: &SdpaTensors,
cfg: &SdpaConfig,
bias: &dyn AttnBias,
mask: &dyn KeyMask,
y: &mut [f32],
mut qk: Option<QkCapture>,
) {
let SdpaTensors {
q,
k,
v,
batch,
num_heads,
num_kv_heads,
q_seq,
kv_seq,
head_size,
v_head_size,
} = *t;
debug_assert_eq!(q.len(), batch * num_heads * q_seq * head_size);
debug_assert_eq!(k.len(), batch * num_kv_heads * kv_seq * head_size);
debug_assert_eq!(v.len(), batch * num_kv_heads * kv_seq * v_head_size);
debug_assert_eq!(y.len(), batch * num_heads * q_seq * v_head_size);
debug_assert!(num_kv_heads > 0 && num_heads.is_multiple_of(num_kv_heads));
let heads_per_kv = num_heads / num_kv_heads;
let (post_scale, operand_scale) = match cfg.scale {
ScaleMode::PostDot(s) => (s, 1.0f32),
ScaleMode::SplitSqrt(s) => (1.0f32, s.sqrt()),
};
let mut scores = vec![0.0f32; kv_seq];
for b in 0..batch {
for n in 0..num_heads {
let kv_n = n / heads_per_kv;
for i in 0..q_seq {
let q_base = ((b * num_heads + n) * q_seq + i) * head_size;
let cap_base = ((b * num_heads + n) * q_seq + i) * kv_seq;
for (j, sc) in scores.iter_mut().enumerate() {
let k_base = ((b * num_kv_heads + kv_n) * kv_seq + j) * head_size;
let mut acc = 0.0f32;
for p in 0..head_size {
acc += (q[q_base + p] * operand_scale) * (k[k_base + p] * operand_scale);
}
let mut s = acc * post_scale;
if let Some(cap) = qk.as_mut()
&& cap.stage == QkCaptureStage::PostScale
{
cap.scores[cap_base + j] = s;
}
if let Some(softcap) = cfg.softcap {
s = softcap * (s / softcap).tanh();
}
if let Some(cap) = qk.as_mut()
&& cap.stage == QkCaptureStage::PostSoftcap
{
cap.scores[cap_base + j] = s;
}
s += bias.at(b, n, i, j);
s += mask.at(b, i, j);
if cfg.causal && (j as i64) > cfg.past_seq as i64 + i as i64 {
s = cfg.causal_fill;
}
*sc = s;
}
if let Some(cap) = qk.as_mut()
&& cap.stage == QkCaptureStage::PreSoftmax
{
cap.scores[cap_base..cap_base + kv_seq].copy_from_slice(&scores);
}
softmax_in_place(&mut scores, SoftmaxExp::F32);
if let Some(cap) = qk.as_mut()
&& cap.stage == QkCaptureStage::PostSoftmax
{
cap.scores[cap_base..cap_base + kv_seq].copy_from_slice(&scores);
}
let y_base = ((b * num_heads + n) * q_seq + i) * v_head_size;
for c in 0..v_head_size {
let mut acc = 0.0f32;
for (j, &p) in scores.iter().enumerate() {
let v_idx = ((b * num_kv_heads + kv_n) * kv_seq + j) * v_head_size + c;
acc += p * v[v_idx];
}
y[y_base + c] = acc;
}
}
}
}
}
#[cfg(feature = "mlas")]
fn sdpa_f32_fast(
t: &SdpaTensors,
cfg: &SdpaConfig,
bias: &dyn AttnBias,
mask: &dyn KeyMask,
y: &mut [f32],
) {
use rayon::prelude::*;
let SdpaTensors {
q,
k,
v,
batch,
num_heads,
num_kv_heads,
q_seq,
kv_seq,
head_size,
v_head_size,
} = *t;
debug_assert_eq!(q.len(), batch * num_heads * q_seq * head_size);
debug_assert_eq!(k.len(), batch * num_kv_heads * kv_seq * head_size);
debug_assert_eq!(v.len(), batch * num_kv_heads * kv_seq * v_head_size);
debug_assert_eq!(y.len(), batch * num_heads * q_seq * v_head_size);
debug_assert!(num_kv_heads > 0 && num_heads.is_multiple_of(num_kv_heads));
let heads_per_kv = num_heads / num_kv_heads;
let alpha = match cfg.scale {
ScaleMode::PostDot(s) => s,
ScaleMode::SplitSqrt(s) => s,
};
let tile_v = q_seq * v_head_size;
y.par_chunks_mut(tile_v)
.enumerate()
.for_each(|(bh, y_tile)| {
let b = bh / num_heads;
let n = bh % num_heads;
let kv_n = n / heads_per_kv;
let q_off = ((b * num_heads + n) * q_seq) * head_size;
let k_off = ((b * num_kv_heads + kv_n) * kv_seq) * head_size;
let v_off = ((b * num_kv_heads + kv_n) * kv_seq) * v_head_size;
let q_tile = &q[q_off..q_off + q_seq * head_size];
let k_tile = &k[k_off..k_off + kv_seq * head_size];
let v_tile = &v[v_off..v_off + kv_seq * v_head_size];
let mut logits = vec![0.0f32; q_seq * kv_seq];
mlas_sys::sgemm(
false,
true,
q_seq,
kv_seq,
head_size,
alpha,
q_tile,
head_size,
k_tile,
head_size,
0.0,
&mut logits,
kv_seq,
);
for i in 0..q_seq {
let row = &mut logits[i * kv_seq..i * kv_seq + kv_seq];
for (j, s) in row.iter_mut().enumerate() {
let mut val = *s;
if let Some(softcap) = cfg.softcap {
val = softcap * (val / softcap).tanh();
}
val += bias.at(b, n, i, j);
val += mask.at(b, i, j);
if cfg.causal && (j as i64) > cfg.past_seq as i64 + i as i64 {
val = cfg.causal_fill;
}
*s = val;
}
softmax_in_place(row, SoftmaxExp::F32);
}
mlas_sys::sgemm(
false,
false,
q_seq,
v_head_size,
kv_seq,
1.0,
&logits,
kv_seq,
v_tile,
v_head_size,
0.0,
y_tile,
v_head_size,
);
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decode_row_f64_intermediate_is_bit_exact_with_gqa_reference() {
let (kv_seq, dh, dv) = (23usize, 133usize, 17usize);
let (lo, hi) = (5usize, 21usize);
let scale = 1.0 / (dh as f32).sqrt();
let softcap = 7.5f32;
let q: Vec<f32> = (0..dh)
.map(|i| ((i * 17 % 101) as f32 - 50.0) / 37.0)
.collect();
let k: Vec<f32> = (0..kv_seq * dh)
.map(|i| ((i * 29 % 211) as f32 - 105.0) / 61.0)
.collect();
let v: Vec<f32> = (0..kv_seq * dv)
.map(|i| ((i * 43 % 157) as f32 - 78.0) / 53.0)
.collect();
let mut scores = vec![0.0f32; hi - lo];
for (i, ks) in (lo..hi).enumerate() {
let k_base = ks * dh;
let mut score = dot_f32(&q, &k[k_base..k_base + dh]);
score *= scale;
score = softcap * (score / softcap).tanh();
scores[i] = score;
}
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for score in &mut scores {
*score = ((*score - max) as f64).exp() as f32;
sum += *score;
}
if sum > 0.0 {
for score in &mut scores {
*score /= sum;
}
}
let mut expected = vec![0.0f32; dv];
for (i, ks) in (lo..hi).enumerate() {
let probability = scores[i];
if probability == 0.0 {
continue;
}
axpy_f32(&mut expected, probability, &v[ks * dv..(ks + 1) * dv]);
}
let mut actual = vec![f32::NAN; dv];
sdpa_decode_row(
&q,
&k,
&v,
kv_seq,
lo,
hi,
scale,
Some(softcap),
SoftmaxExp::F64Intermediate,
&mut actual,
);
assert_eq!(
actual.iter().map(|x| x.to_bits()).collect::<Vec<_>>(),
expected.iter().map(|x| x.to_bits()).collect::<Vec<_>>()
);
}
#[test]
fn split_decode_matches_sequential_reference_within_tolerance() {
const TOLERANCE: f32 = 1e-6;
fn run_case(kv_seq: usize, dh: usize, dv: usize, lo: usize, hi: usize, split_count: usize) {
let scale = 1.0 / (dh.max(1) as f32).sqrt();
let softcap = Some(6.25f32);
let q: Vec<f32> = (0..dh)
.map(|i| ((i * 13 % 97) as f32 - 48.0) / 29.0)
.collect();
let k: Vec<f32> = (0..kv_seq * dh)
.map(|i| ((i * 31 % 199) as f32 - 99.0) / 57.0)
.collect();
let v: Vec<f32> = (0..kv_seq * dv)
.map(|i| ((i * 37 % 173) as f32 - 86.0) / 47.0)
.collect();
let mut reference = vec![f32::NAN; dv];
sdpa_decode_row(
&q,
&k,
&v,
kv_seq,
lo,
hi,
scale,
softcap,
SoftmaxExp::F64Intermediate,
&mut reference,
);
let length = hi - lo;
let base = length / split_count;
let remainder = length % split_count;
let mut partials = Vec::with_capacity(split_count);
let mut partial_outputs = vec![0.0f64; split_count * dv];
for chunk in 0..split_count {
let chunk_lo = lo + chunk * base + chunk.min(remainder);
let chunk_hi = chunk_lo + base + usize::from(chunk < remainder);
let slot = &mut partial_outputs[chunk * dv..(chunk + 1) * dv];
partials.push(sdpa_decode_partial(
&q, &k, &v, kv_seq, chunk_lo, chunk_hi, scale, softcap, slot,
));
}
let mut combined = vec![f32::NAN; dv];
combine_decode_partials(&partials, &partial_outputs, dv, &mut combined);
let mut max_abs_error = 0.0f32;
for (&reference_value, &combined_value) in reference.iter().zip(&combined) {
max_abs_error = max_abs_error.max((reference_value - combined_value).abs());
}
assert!(
max_abs_error <= TOLERANCE,
"kv_seq={kv_seq} dh={dh} dv={dv} lo={lo} hi={hi} split_count={split_count}: \
max abs error {max_abs_error} exceeds {TOLERANCE}"
);
}
for &(dh, dv) in &[(64usize, 64usize), (128, 128), (96, 40), (133, 17)] {
for &kv_seq in &[1usize, 2, 7, 64, 200, 1024] {
for &split_count in &[1usize, 2, 3, 4, 8, 16] {
run_case(kv_seq, dh, dv, 0, kv_seq, split_count);
}
}
}
run_case(200, 128, 128, 37, 200, 4);
run_case(200, 128, 128, 37, 200, 7);
run_case(1024, 96, 40, 511, 1024, 5);
run_case(5, 64, 64, 0, 5, 8);
run_case(5, 64, 64, 2, 5, 16);
run_case(16, 64, 64, 8, 8, 4);
}
fn reference(
q: &[f32],
k: &[f32],
v: &[f32],
s: usize,
dh: usize,
dv: usize,
scale: f32,
) -> Vec<f32> {
let mut out = vec![0.0f32; s * dv];
for i in 0..s {
let mut scores = vec![0.0f32; s];
for (j, sc) in scores.iter_mut().enumerate() {
let mut acc = 0.0f32;
for p in 0..dh {
acc += q[i * dh + p] * k[j * dh + p];
}
*sc = acc * scale;
}
let m = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let sum: f32 = scores.iter().map(|x| (x - m).exp()).sum();
for c in 0..dv {
let mut acc = 0.0f32;
for (j, sc) in scores.iter().enumerate() {
acc += ((sc - m).exp() / sum) * v[j * dv + c];
}
out[i * dv + c] = acc;
}
}
out
}
#[test]
fn postdot_matches_reference() {
let (s, dh, dv) = (3usize, 4usize, 2usize);
let q: Vec<f32> = (0..s * dh).map(|x| (x as f32) * 0.1 - 0.5).collect();
let k: Vec<f32> = (0..s * dh).map(|x| (x as f32) * 0.05).collect();
let v: Vec<f32> = (0..s * dv).map(|x| (x as f32) * 0.2).collect();
let scale = 1.0 / (dh as f32).sqrt();
let t = SdpaTensors {
q: &q,
k: &k,
v: &v,
batch: 1,
num_heads: 1,
num_kv_heads: 1,
q_seq: s,
kv_seq: s,
head_size: dh,
v_head_size: dv,
};
let cfg = SdpaConfig {
scale: ScaleMode::PostDot(scale),
softcap: None,
causal: false,
past_seq: 0,
causal_fill: f32::MIN,
};
let mut y = vec![0.0f32; s * dv];
sdpa_f32_scalar(&t, &cfg, &NoBias, &NoMask, &mut y, None);
let want = reference(&q, &k, &v, s, dh, dv, scale);
for (a, b) in y.iter().zip(want.iter()) {
assert!((a - b).abs() < 1e-6, "got {y:?} want {want:?}");
}
}
#[test]
fn causal_masks_future_keys() {
let (s, dh, dv) = (2usize, 2usize, 2usize);
let q = vec![1.0f32, 0.0, 0.0, 1.0];
let k = vec![1.0f32, 0.0, 0.0, 1.0];
let v = vec![10.0f32, 20.0, 30.0, 40.0];
let t = SdpaTensors {
q: &q,
k: &k,
v: &v,
batch: 1,
num_heads: 1,
num_kv_heads: 1,
q_seq: s,
kv_seq: s,
head_size: dh,
v_head_size: dv,
};
let cfg = SdpaConfig {
scale: ScaleMode::PostDot(1.0),
softcap: None,
causal: true,
past_seq: 0,
causal_fill: f32::MIN,
};
let mut y = vec![0.0f32; s * dv];
sdpa_f32_scalar(&t, &cfg, &NoBias, &NoMask, &mut y, None);
assert!((y[0] - 10.0).abs() < 1e-6 && (y[1] - 20.0).abs() < 1e-6);
}
#[test]
fn gqa_head_sharing_reads_grouped_kv() {
let (s, dh, dv) = (1usize, 2usize, 2usize);
let q = vec![1.0f32, 0.0, 0.0, 1.0];
let k = vec![1.0f32, 1.0]; let v = vec![5.0f32, 7.0];
let t = SdpaTensors {
q: &q,
k: &k,
v: &v,
batch: 1,
num_heads: 2,
num_kv_heads: 1,
q_seq: s,
kv_seq: s,
head_size: dh,
v_head_size: dv,
};
let cfg = SdpaConfig {
scale: ScaleMode::PostDot(1.0),
softcap: None,
causal: false,
past_seq: 0,
causal_fill: f32::MIN,
};
let mut y = vec![0.0f32; 2 * s * dv];
sdpa_f32_scalar(&t, &cfg, &NoBias, &NoMask, &mut y, None);
for h in 0..2 {
assert!((y[h * dv] - 5.0).abs() < 1e-6 && (y[h * dv + 1] - 7.0).abs() < 1e-6);
}
}
#[test]
fn splitsqrt_scale_equivalent_to_postdot_for_moderate_values() {
let (s, dh, dv) = (2usize, 3usize, 2usize);
let q: Vec<f32> = (0..s * dh).map(|x| (x as f32) * 0.3).collect();
let k: Vec<f32> = (0..s * dh).map(|x| (x as f32) * 0.2 - 0.1).collect();
let v: Vec<f32> = (0..s * dv).map(|x| (x as f32) * 0.5).collect();
let scale = 1.0 / (dh as f32).sqrt();
let base = SdpaTensors {
q: &q,
k: &k,
v: &v,
batch: 1,
num_heads: 1,
num_kv_heads: 1,
q_seq: s,
kv_seq: s,
head_size: dh,
v_head_size: dv,
};
let mut y_post = vec![0.0f32; s * dv];
sdpa_f32_scalar(
&base,
&SdpaConfig {
scale: ScaleMode::PostDot(scale),
softcap: None,
causal: false,
past_seq: 0,
causal_fill: f32::MIN,
},
&NoBias,
&NoMask,
&mut y_post,
None,
);
let mut y_split = vec![0.0f32; s * dv];
sdpa_f32_scalar(
&base,
&SdpaConfig {
scale: ScaleMode::SplitSqrt(scale),
softcap: None,
causal: false,
past_seq: 0,
causal_fill: f32::MIN,
},
&NoBias,
&NoMask,
&mut y_split,
None,
);
for (a, b) in y_post.iter().zip(y_split.iter()) {
assert!((a - b).abs() < 1e-5, "post {y_post:?} split {y_split:?}");
}
}
#[cfg(feature = "mlas")]
fn fill(n: usize, seed: u64) -> Vec<f32> {
let mut s = seed.wrapping_add(0x9E37_79B9_7F4A_7C15);
(0..n)
.map(|_| {
s ^= s >> 30;
s = s.wrapping_mul(0xBF58_476D_1CE4_E5B9);
s ^= s >> 27;
((s >> 40) as f32 / (1u64 << 24) as f32) * 2.0 - 1.0
})
.collect()
}
#[cfg(feature = "mlas")]
struct DenseKeyMask<'a> {
data: &'a [f32],
q_seq: usize,
kv_seq: usize,
}
#[cfg(feature = "mlas")]
impl KeyMask for DenseKeyMask<'_> {
fn at(&self, b: usize, i: usize, j: usize) -> f32 {
self.data[(b * self.q_seq + i) * self.kv_seq + j]
}
}
#[cfg(feature = "mlas")]
#[test]
fn fast_path_matches_scalar_reference() {
struct Shape {
name: &'static str,
batch: usize,
nq: usize,
nkv: usize,
sq: usize,
tk: usize,
dh: usize,
dv: usize,
causal: bool,
past: usize,
softcap: Option<f32>,
split_sqrt: bool,
with_bias: bool,
with_mask: bool,
}
let shapes = [
Shape {
name: "mha-prefill",
batch: 2,
nq: 4,
nkv: 4,
sq: 7,
tk: 7,
dh: 8,
dv: 8,
causal: false,
past: 0,
softcap: None,
split_sqrt: false,
with_bias: false,
with_mask: false,
},
Shape {
name: "mha-causal",
batch: 1,
nq: 3,
nkv: 3,
sq: 6,
tk: 6,
dh: 5,
dv: 5,
causal: true,
past: 0,
softcap: None,
split_sqrt: false,
with_bias: false,
with_mask: false,
},
Shape {
name: "gqa",
batch: 2,
nq: 8,
nkv: 2,
sq: 5,
tk: 5,
dh: 4,
dv: 4,
causal: false,
past: 0,
softcap: None,
split_sqrt: false,
with_bias: false,
with_mask: false,
},
Shape {
name: "mqa-decode",
batch: 2,
nq: 6,
nkv: 1,
sq: 1,
tk: 9,
dh: 8,
dv: 8,
causal: false,
past: 8,
softcap: None,
split_sqrt: false,
with_bias: false,
with_mask: false,
},
Shape {
name: "cross-diff-dv",
batch: 1,
nq: 2,
nkv: 2,
sq: 4,
tk: 6,
dh: 5,
dv: 3,
causal: false,
past: 0,
softcap: None,
split_sqrt: false,
with_bias: true,
with_mask: false,
},
Shape {
name: "softcap",
batch: 1,
nq: 2,
nkv: 2,
sq: 5,
tk: 5,
dh: 6,
dv: 6,
causal: false,
past: 0,
softcap: Some(30.0),
split_sqrt: false,
with_bias: false,
with_mask: false,
},
Shape {
name: "split-sqrt-mask",
batch: 2,
nq: 3,
nkv: 3,
sq: 4,
tk: 5,
dh: 7,
dv: 7,
causal: false,
past: 0,
softcap: None,
split_sqrt: true,
with_bias: false,
with_mask: true,
},
Shape {
name: "causal-past-decode",
batch: 1,
nq: 4,
nkv: 4,
sq: 1,
tk: 12,
dh: 8,
dv: 8,
causal: true,
past: 11,
softcap: None,
split_sqrt: false,
with_bias: false,
with_mask: false,
},
];
for sh in &shapes {
let q = fill(sh.batch * sh.nq * sh.sq * sh.dh, 1 + sh.sq as u64);
let k = fill(sh.batch * sh.nkv * sh.tk * sh.dh, 2 + sh.tk as u64);
let v = fill(sh.batch * sh.nkv * sh.tk * sh.dv, 3 + sh.dv as u64);
let scale = 1.0 / (sh.dh as f32).sqrt();
let t = SdpaTensors {
q: &q,
k: &k,
v: &v,
batch: sh.batch,
num_heads: sh.nq,
num_kv_heads: sh.nkv,
q_seq: sh.sq,
kv_seq: sh.tk,
head_size: sh.dh,
v_head_size: sh.dv,
};
let cfg = SdpaConfig {
scale: if sh.split_sqrt {
ScaleMode::SplitSqrt(scale)
} else {
ScaleMode::PostDot(scale)
},
softcap: sh.softcap,
causal: sh.causal,
past_seq: sh.past,
causal_fill: f32::MIN,
};
let bias_data = fill(sh.batch * sh.nq * sh.sq * sh.tk, 7);
let mask_data: Vec<f32> = fill(sh.batch * sh.sq * sh.tk, 9)
.into_iter()
.map(|x| if x < -0.5 { -1.0e9 } else { 0.0 })
.collect();
let no_bias = NoBias;
let bc_bias = BroadcastBias::new(&bias_data, [sh.batch, sh.nq, sh.sq, sh.tk]);
let bias: &dyn AttnBias = if sh.with_bias { &bc_bias } else { &no_bias };
let no_mask = NoMask;
let dm = DenseKeyMask {
data: &mask_data,
q_seq: sh.sq,
kv_seq: sh.tk,
};
let mask: &dyn KeyMask = if sh.with_mask { &dm } else { &no_mask };
let out_len = sh.batch * sh.nq * sh.sq * sh.dv;
let mut y_scalar = vec![0.0f32; out_len];
sdpa_f32_scalar(&t, &cfg, bias, mask, &mut y_scalar, None);
let mut y_fast = vec![0.0f32; out_len];
sdpa_f32_fast(&t, &cfg, bias, mask, &mut y_fast);
let mut max_abs = 0.0f32;
let mut worst = 0.0f32;
for (a, b) in y_fast.iter().zip(y_scalar.iter()) {
let abs = (a - b).abs();
max_abs = max_abs.max(abs);
worst = worst.max(abs - (1e-5 + 1e-4 * b.abs()));
}
assert!(
worst <= 0.0,
"shape {}: fast vs scalar exceeds atol+rtol (max_abs={max_abs:e})",
sh.name
);
}
}
#[cfg(feature = "mlas")]
#[test]
#[ignore = "provisional microbench; shared host — run manually with --nocapture"]
fn sdpa_fast_provisional_bench() {
use std::time::Instant;
fn run(name: &str, batch: usize, nq: usize, nkv: usize, sq: usize, tk: usize, dh: usize) {
let q = fill(batch * nq * sq * dh, 11);
let k = fill(batch * nkv * tk * dh, 22);
let v = fill(batch * nkv * tk * dh, 33);
let t = SdpaTensors {
q: &q,
k: &k,
v: &v,
batch,
num_heads: nq,
num_kv_heads: nkv,
q_seq: sq,
kv_seq: tk,
head_size: dh,
v_head_size: dh,
};
let cfg = SdpaConfig {
scale: ScaleMode::PostDot(1.0 / (dh as f32).sqrt()),
softcap: None,
causal: sq > 1,
past_seq: tk - sq,
causal_fill: f32::MIN,
};
let out_len = batch * nq * sq * dh;
let mut y = vec![0.0f32; out_len];
let iters = 20;
sdpa_f32_scalar(&t, &cfg, &NoBias, &NoMask, &mut y, None);
let t0 = Instant::now();
for _ in 0..iters {
sdpa_f32_scalar(&t, &cfg, &NoBias, &NoMask, &mut y, None);
}
let scalar = t0.elapsed().as_secs_f64() / iters as f64;
sdpa_f32_fast(&t, &cfg, &NoBias, &NoMask, &mut y);
let t1 = Instant::now();
for _ in 0..iters {
sdpa_f32_fast(&t, &cfg, &NoBias, &NoMask, &mut y);
}
let fast = t1.elapsed().as_secs_f64() / iters as f64;
println!(
"[sdpa-bench PROVISIONAL] {name:>16}: scalar {:>9.3} ms fast {:>9.3} ms speedup {:>5.2}x",
scalar * 1e3,
fast * 1e3,
scalar / fast
);
}
println!("[sdpa-bench] PROVISIONAL numbers — shared host, treat as indicative only");
run("prefill", 1, 32, 32, 512, 512, 128);
run("decode", 1, 32, 32, 1, 513, 128);
run("gqa-prefill", 1, 32, 8, 512, 512, 128);
}
}