scaled_dot_product_attention/
lib.rs1use {
2 candle_core::{D, Result, Tensor},
3 candle_nn::ops::softmax_last_dim,
4};
5
6pub fn scaled_dot_product_attention(q: &Tensor, k: &Tensor, v: &Tensor) -> Result<Tensor> {
7 let dim = q.dim(D::Minus1)?;
8 let scale_factor = 1.0 / (dim as f64).sqrt();
9 let attn_weights = (q.matmul(&k.t()?)? * scale_factor)?;
10 softmax_last_dim(&attn_weights)?.matmul(v)
11}
12
13#[cfg(test)]
14mod tests {
15 use super::*;
16 use all_close::TensorAllClose;
17 use candle_core::{DType, Device};
18 use candle_core::{Result, Tensor};
19
20 #[test]
21 fn test_basic_attention() -> Result<()> {
22 let device = &Device::Cpu;
23 let q = Tensor::new(&[[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], device)?;
25 let k = Tensor::new(&[[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], device)?;
26 let v = Tensor::new(&[[1.0, 2.0], [3.0, 4.0]], device)?;
27
28 let output = scaled_dot_product_attention(&q, &k, &v)?;
29
30 let expected = Tensor::new(&[[1.7191, 2.7191], [2.2809, 3.2809]], device)?;
32 assert!(output.all_close(&expected, 1e-1)?);
33 Ok(())
34 }
35
36 #[test]
37 fn test_attention_with_scale() -> Result<()> {
38 let device = &Device::Cpu;
39 let q = Tensor::new(&[[1.0, 0.0, 0.0, 0.0]], device)?;
41 let k = Tensor::new(&[[1.0, 0.0, 0.0, 0.0]], device)?;
42 let v = Tensor::new(&[[1.0, 2.0]], device)?;
43
44 let output = scaled_dot_product_attention(&q, &k, &v)?;
45
46 let expected = Tensor::new(&[[1.0, 2.0]], device)?;
48 assert!(output.all_close(&expected, 1e-1)?);
49 Ok(())
50 }
51
52 #[test]
53 fn test_dtype_consistency() -> Result<()> {
54 let device = &Device::Cpu;
55 let q = Tensor::new(&[[1.0f32, 0.0], [0.0, 1.0]], device)?;
57 let k = Tensor::new(&[[1.0f32, 0.0], [0.0, 1.0]], device)?;
58 let v = Tensor::new(&[[1.0f32, 2.0], [3.0, 4.0]], device)?;
59
60 let output = scaled_dot_product_attention(&q, &k, &v)?;
61
62 assert_eq!(output.dtype(), DType::F32);
63 Ok(())
64 }
65}