pub fn make_causal_mask(
seq: usize,
num_masked: usize,
device: &Device,
) -> Result<Tensor>Expand description
Additive causal mask of shape [1, 1, seq, seq]: 0.0 where num_masked <= k <= q,
-inf elsewhere. num_masked lets a prefix of keys be masked out regardless of query
position (used for left-padded contexts); pass 0 for a plain causal mask.