Skip to main content

make_causal_mask

Function make_causal_mask 

Source
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.