| 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()? |