// position_encoding.mbt — sinusoidal and learnable position embeddings.
//
// Two complementary encodings for transformer input:
//
// 1. Sinusoidal (Vaswani 2017 §3.5) — fixed, non-trainable:
// PE(pos, 2i) = sin(pos / 10000^(2i / d_model))
// PE(pos, 2i+1) = cos(pos / 10000^(2i / d_model))
// Returns the full (max_len, d_model) table.
//
// 2. Learnable (BERT / ViT / GPT-2 style) — trainable parameter
// matrix initialised with N(0, 1) (or Xavier), shape
// (max_len, d_model). Backward only accumulates gradients at the
// active positions actually used.
//
// Combine via element-wise add: x + positional_embedding(x).
///|
/// Sinusoidal positional encoding table (non-trainable).
///
/// For position `pos` ∈ [0, max_len) and embedding index `j` ∈
/// [0, d_model):
///
/// PE[pos, j] = sin(pos / 10000^(⌊j/2⌋·2 / d_model)) if j is even
/// PE[pos, j] = cos(pos / 10000^(⌊j/2⌋·2 / d_model)) if j is odd
///
/// Returns row-major flat array of length `max_len * d_model`.
pub fn sinusoidal_position_encoding(
max_len : Int,
d_model : Int,
) -> Array[Float] {
let out : Array[Float] = Array::make(max_len * d_model, 0.0F)
for pos in 0.. Float {
expf(x * 4.605170185988091F) // 10000 = e^(ln 10000)
}
///|
/// Learnable positional embedding parameter container.
pub struct PositionalEmbedding {
max_len : Int
d_model : Int
// Row-major flat (max_len, d_model).
weight : Array[Float]
}
///|
/// Construct a learnable positional embedding with Xavier-normal
/// init (N(0, sqrt(2 / (1 + d_model)))).
pub fn PositionalEmbedding::new(
max_len : Int,
d_model : Int,
seed : UInt64,
) -> PositionalEmbedding {
let rng = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let std = sqrtf(2.0F / Float::from_int(1 + d_model))
let weight : Array[Float] = Array::make(max_len * d_model, 0.0F)
let mut i = 0
while i < max_len * d_model {
let (z1, _) = box_muller(rng)
weight[i] = Float::from_double(z1) * std
i = i + 1
}
{ max_len, d_model, weight }
}
///|
/// Slice the embedding table to the first `seq_len` rows. Returns a
/// fresh row-major array of length `seq_len * d_model`. Caller must
/// ensure `seq_len <= max_len`.
pub fn positional_embedding_forward(
pe : PositionalEmbedding,
seq_len : Int,
) -> Array[Float] {
if seq_len > pe.max_len {
abort("positional_embedding_forward: seq_len \{seq_len} > max_len \{pe.max_len}")
}
let d = pe.d_model
let out : Array[Float] = Array::make(seq_len * d, 0.0F)
for s in 0.. Array[Float] {
let d = pe.d_model
let d_weight : Array[Float] = Array::make(pe.max_len * d, 0.0F)
for s in 0..