Skip to main content

zsfm_nn/
mask.rs

1use anyhow::Result;
2use candle_core::{Device, Tensor};
3
4/// Additive causal mask of shape `[1, 1, seq, seq]`: `0.0` where `num_masked <= k <= q`,
5/// `-inf` elsewhere. `num_masked` lets a prefix of keys be masked out regardless of query
6/// position (used for left-padded contexts); pass `0` for a plain causal mask.
7pub 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        // key 0 is masked for every query now.
45        assert_eq!(v, vec![inf, inf, inf, inf, 0.0, inf, inf, 0.0, 0.0]);
46    }
47}