MCPcopy Create free account
hub / github.com/evilsocket/cake / attention

Method attention

cake-core/src/backends/cpu/mod.rs:44–85  ·  view source on GitHub ↗
(
        &self,
        q: &Tensor,
        k: &Tensor,
        v: &Tensor,
        scale: f32,
        causal: bool,
    )

Source from the content-addressed store, hash-verified

42 }
43
44 fn attention(
45 &self,
46 q: &Tensor,
47 k: &Tensor,
48 v: &Tensor,
49 scale: f32,
50 causal: bool,
51 ) -> Result<Tensor> {
52 // Manual SDPA: Q·K^T * scale → softmax → ·V
53 // Skip dtype conversion when already F32 (avoids alloc check overhead)
54 let q_f32 = if q.dtype() == DType::F32 { q.clone() } else { q.to_dtype(DType::F32)? };
55 let k_f32 = if k.dtype() == DType::F32 { k.clone() } else { k.to_dtype(DType::F32)? };
56 let v_f32 = if v.dtype() == DType::F32 { v.clone() } else { v.to_dtype(DType::F32)? };
57 let attn = q_f32.matmul(&k_f32.t()?)?;
58 let attn = (attn * scale as f64)?;
59 let attn = if causal {
60 let seq_len = q_f32.dim(2)?;
61 if seq_len <= 1 {
62 // Generation step: single query attends to all KV — no masking needed
63 attn
64 } else {
65 // Build upper-triangular future mask (1 = future/masked position)
66 let mut mask_data = vec![0u8; seq_len * seq_len];
67 for i in 0..seq_len {
68 for j in (i + 1)..seq_len {
69 mask_data[i * seq_len + j] = 1;
70 }
71 }
72 let mask =
73 Tensor::from_vec(mask_data, (1, 1, seq_len, seq_len), q_f32.device())?;
74 // Use scalar broadcast instead of allocating full neg_inf tensor
75 let neg_inf = Tensor::new(f32::NEG_INFINITY, q_f32.device())?
76 .broadcast_as(attn.shape())?;
77 mask.broadcast_as(attn.shape())?
78 .where_cond(&neg_inf, &attn)?
79 }
80 } else {
81 attn
82 };
83 let attn = candle_nn::ops::softmax_last_dim(&attn)?;
84 attn.matmul(&v_f32)
85 }
86
87 fn silu_mul(&self, gate: &Tensor, up: &Tensor) -> Result<Tensor> {
88 candle_nn::ops::silu(&gate.contiguous()?)? * up.contiguous()?

Calls 6

dtypeMethod · 0.80
shapeMethod · 0.80
softmax_last_dimFunction · 0.50
cloneMethod · 0.45
matmulMethod · 0.45
deviceMethod · 0.45