bunsen 0.22.0

bunsen is a batteries included common library for burn
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
use burn::{
    Tensor,
    prelude::{
        Backend,
        Bool,
    },
};

/// Generate a Bool causal mask [1, `seq_len`, `n_past` + `seq_len`].
/// `true` = masked (future positions blocked), `false` = attend.
pub fn causal_mask<B: Backend>(
    seq_len: usize,
    n_past: usize,
    device: &B::Device,
) -> Tensor<B, 3, Bool> {
    let total = n_past + seq_len;
    Tensor::<B, 2, Bool>::tril_mask([seq_len, total], n_past as i64, device).unsqueeze::<3>()
}