1use anyhow::Result;
2use candle_core::{Device, Tensor};
3
4pub fn make_causal_mask(seq: usize, num_masked: usize, device: &Device) -> Result<Tensor> {
8 let data: Vec<f32> = (0..seq)
9 .flat_map(|q| {
10 (0..seq).map(move |k| {
11 if k <= q && k >= num_masked {
12 0.0f32
13 } else {
14 f32::NEG_INFINITY
15 }
16 })
17 })
18 .collect();
19 Ok(Tensor::from_vec(data, (seq, seq), device)?
20 .unsqueeze(0)?
21 .unsqueeze(0)?)
22}
23
24#[cfg(test)]
25mod tests {
26 use super::*;
27
28 #[test]
29 fn plain_causal_mask_shape_and_values() {
30 let device = Device::Cpu;
31 let m = make_causal_mask(3, 0, &device).unwrap();
32 assert_eq!(m.dims(), &[1, 1, 3, 3]);
33 let v: Vec<f32> = m.flatten_all().unwrap().to_vec1().unwrap();
34 let inf = f32::NEG_INFINITY;
35 assert_eq!(v, vec![0.0, inf, inf, 0.0, 0.0, inf, 0.0, 0.0, 0.0]);
36 }
37
38 #[test]
39 fn num_masked_blocks_prefix_keys() {
40 let device = Device::Cpu;
41 let m = make_causal_mask(3, 1, &device).unwrap();
42 let v: Vec<f32> = m.flatten_all().unwrap().to_vec1().unwrap();
43 let inf = f32::NEG_INFINITY;
44 assert_eq!(v, vec![inf, inf, inf, inf, 0.0, inf, inf, 0.0, 0.0]);
46 }
47}