bunsen 0.24.1

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,
    },
};

/// Generates 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>()
}