Skip to main content

scaled_dot_product_attention/
lib.rs

1use {
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        // Simple case with 2 tokens and 3 dimensions
24        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        // Expected output should be the same as v since q and k are identity-like
31        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        // Larger dimension to test scaling
40        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        // Should still match v despite higher dimension due to scaling
47        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        // Test with f32 dtype
56        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}