// attention_mask.mbt — Attention mask utilities for transformer
// self-attention.
//
// Two common mask shapes are supported:
//
// Per-head (num_heads × seq_len × seq_len) → fed directly
// to MHA.
// Head-broadcast (seq_len × seq_len) → expanded to
// per-head via
// mask_broadcast.
//
// Mask values are ADDITIVE — they are added to the attention scores
// before softmax. Use a large negative value (e.g. -1e9) to suppress
// a position; use 0 to leave it unconstrained.
//
// Causal mask is the upper-triangular -1e9 matrix used in
// autoregressive (decoder-style) attention.
///|
/// Build a causal mask of shape (seq_len, seq_len). Positions
/// `(i, j)` with `j > i` are masked (set to `-1e9`); `j <= i` are
/// unconstrained (set to 0).
///
/// Returns row-major flat `Array[Float]` of length `seq_len * seq_len`.
pub fn causal_mask(seq_len : Int) -> Array[Float] {
let mask : Array[Float] = Array::make(seq_len * seq_len, 0.0F)
for i in 0.. i {
mask[i * seq_len + j] = -1000000000.0F
}
}
}
mask
}
///|
/// Broadcast a 2D mask (seq_len × seq_len) to per-head shape
/// (num_heads × seq_len × seq_len). All heads share the same mask.
pub fn mask_broadcast(
mask_2d : Array[Float],
seq_len : Int,
num_heads : Int,
) -> Array[Float] {
if mask_2d.length() != seq_len * seq_len {
abort(
"mask_broadcast: mask length \{mask_2d.length()} != seq_len^2 \{seq_len * seq_len}",
)
}
let out : Array[Float] = Array::make(num_heads * seq_len * seq_len, 0.0F)
for h in 0.. Array[Float] {
if allowed.length() != seq_len {
abort(
"mask_from_allowed: outer length \{allowed.length()} != seq_len \{seq_len}",
)
}
let mask : Array[Float] = Array::make(seq_len * seq_len, 0.0F)
for i in 0.. Array[Float] {
mask_broadcast(causal_mask(seq_len), seq_len, num_heads)
}