use super::activation::softmax;
use super::matrix::{matmul, transpose};
use super::vector::{axpy, dot};
pub fn scaled_dot_product_attention(
query: &[f32],
key: &[f32],
value: &[f32],
seq_len: usize,
d_head: usize,
mask: Option<&[f32]>,
) -> Vec<f32> {
if seq_len == 1 {
return scaled_dot_product_attention_single(query, key, value, d_head, mask);
}
let kv_len = key.len() / d_head;
let scale = 1.0 / (d_head as f32).sqrt();
let key_t = transpose(key, kv_len, d_head);
let mut scores = matmul(query, &key_t, seq_len, d_head, kv_len);
for s in &mut scores {
*s *= scale;
}
if let Some(m) = mask {
assert_eq!(m.len(), seq_len * kv_len, "mask dimensions mismatch");
for (i, &mask_val) in m.iter().enumerate() {
scores[i] += mask_val;
}
}
let mut weights = Vec::with_capacity(scores.len());
for i in 0..seq_len {
let start = i * kv_len;
let end = start + kv_len;
let row_softmax = softmax(&scores[start..end]);
weights.extend(row_softmax);
}
matmul(&weights, value, seq_len, kv_len, d_head)
}
#[must_use]
pub fn scaled_dot_product_attention_single(
query: &[f32],
key: &[f32],
value: &[f32],
d_head: usize,
mask: Option<&[f32]>,
) -> Vec<f32> {
let kv_len = key.len() / d_head;
let scale = 1.0 / (d_head as f32).sqrt();
let mut scores = Vec::with_capacity(kv_len);
for pos in 0..kv_len {
let k_start = pos * d_head;
let score = dot(query, &key[k_start..k_start + d_head]) * scale;
scores.push(score);
}
if let Some(m) = mask {
for (i, &mask_val) in m.iter().take(kv_len).enumerate() {
scores[i] += mask_val;
}
}
let weights = softmax(&scores);
let mut output = vec![0.0_f32; d_head];
for (pos, &weight) in weights.iter().enumerate() {
let v_start = pos * d_head;
axpy(weight, &value[v_start..v_start + d_head], &mut output);
}
output
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_scaled_dot_product_attention() {
let query = vec![1.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0]; let key = query.clone();
let value = query.clone();
let result = scaled_dot_product_attention(&query, &key, &value, 2, 4, None);
assert_eq!(result.len(), 8);
assert!(result.iter().all(|&v| v.is_finite()));
}
#[test]
fn test_scaled_dot_product_attention_with_mask() {
let query = vec![1.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0]; let key = query.clone();
let value = query.clone();
let mask = vec![1.0, 0.0, 1.0, 1.0];
let result = scaled_dot_product_attention(&query, &key, &value, 2, 4, Some(&mask));
assert_eq!(result.len(), 8);
assert!(result.iter().all(|&v| v.is_finite()));
}
#[test]
fn test_scaled_dot_product_attention_single() {
let query = vec![1.0, 0.0, 0.0, 1.0]; let key = vec![1.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0]; let value = key.clone();
let result = scaled_dot_product_attention_single(&query, &key, &value, 4, None);
assert_eq!(result.len(), 4);
assert!(result.iter().all(|&v| v.is_finite()));
}
#[test]
fn test_scaled_dot_product_attention_single_with_mask() {
let query = vec![1.0, 0.0, 0.0, 1.0]; let key = vec![1.0, 0.0, 0.0, 1.0, 0.0, 1.0, 1.0, 0.0]; let value = key.clone();
let mask = vec![0.0, f32::NEG_INFINITY];
let result = scaled_dot_product_attention_single(&query, &key, &value, 4, Some(&mask));
assert_eq!(result.len(), 4);
assert!(result.iter().all(|&v| v.is_finite()));
}
}