// gtrxl_block.mbt — Gated Transformer-XL (GTrXL) block (v0.72.0).
//
// Implementation of the GTrXL gating scheme from Parisotto et al.
// 2020 "Stabilizing Transformers for Reinforcement Learning".
//
// Key innovation: instead of additive residual `out = main_path(x) + x`,
// GTrXL uses a gating branch:
// out = sigmoid(W_gate · LN(x)) ⊙ main_path(LN(x)) + x
// This dampens irrelevant residual contributions and stabilizes training
// over long horizons in RL.
//
// Scope of v0.72.0:
// - GTrXLBlock struct + constructor (single per-token gated FFN)
// - gtrxl_block_token_step (single-token forward; returns intermediates)
// - gtrxl_block_seq_forward (T-step forward; returns per-step intermediates)
// - GTrXLBlockCache (forward-pass cache for future BPTT)
//
// BPTT-driven update deferred. MHAttn integration deferred (per-token
// FFN path is the minimum useful GTrXL primitive; the full attention
// forward will be added in a follow-up).
///|
/// Gated Transformer-XL block. Per-token gated FFN + residual.
/// Differs from `TransformerBlock` (v0.23.3) by gating the FFN path
/// with a learned W_gate projection + element-wise sigmoid.
pub struct GTrXLBlock {
d_model : Int
d_ff : Int
ffn_w1 : Array[Array[Float]]
ffn_b1 : Array[Float]
ffn_w2 : Array[Array[Float]]
ffn_b2 : Array[Float]
ffn_gate_w : Array[Array[Float]]
ffn_gate_b : Array[Float]
}
///|
/// Per-token forward intermediates for one GTrXL block step. Returned
/// from `gtrxl_block_token_step` and stitched together by
/// `gtrxl_block_seq_forward`. Future BPTT backward will read these.
pub struct GTrXLTokenCache {
x_input : Array[Float]
ffn_gate_pre : Array[Float]
ffn_hidden_pre : Array[Float]
ffn_hidden_post : Array[Float]
gated_ffn : Array[Float]
}
///|
/// Forward-pass cache for a full sequence through one GTrXL block.
/// Each field is `[seq_len × ...]` flattened. Reserved for the future
/// BPTT-driven update (v0.72.0 only ships forward).
pub struct GTrXLBlockCache {
seq_len : Int
x_input : Array[Float]
ffn_gate_pre : Array[Float]
ffn_hidden_pre : Array[Float]
ffn_hidden_post : Array[Float]
gated_ffn : Array[Float]
y_seq : Array[Float]
}
///|
/// Build a fresh GTrXLBlock. Gate matrix init has small std (gain
/// ~0.1) so initial gate signal is roughly sigmoid(0) ≈ 0.5 (not too
/// aggressive at startup).
pub fn GTrXLBlock::new(
d_model : Int,
d_ff : Int,
seed : UInt64,
) -> GTrXLBlock {
let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let std1 = sqrtf(2.0F / Float::from_int(d_model))
let ffn_w1 = xavier_normal(d_ff, d_model, std1, rng1)
let ffn_b1 : Array[Float] = Array::make(d_ff, 0.0F)
let rng2 = Xoshiro::from_state(seed + 4UL, seed + 5UL, seed + 6UL, seed + 7UL)
let ffn_w2 = xavier_normal(d_model, d_ff, std1, rng2)
let ffn_b2 : Array[Float] = Array::make(d_model, 0.0F)
let rng3 = Xoshiro::from_state(seed + 8UL, seed + 9UL, seed + 10UL, seed + 11UL)
let ffn_gate_w = xavier_normal(d_model, d_model, 0.1F * std1, rng3)
let ffn_gate_b : Array[Float] = Array::make(d_model, 0.0F)
{ d_model, d_ff, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ffn_gate_w, ffn_gate_b }
}
///|
/// Single-token GTrXL gated FFN + residual forward.
///
/// y = sigmoid(W_gate · x) ⊙ FFN(x) + x
///
/// `x` is a single token (length d_model). Returns
/// `(y, intermediates)` where intermediates are the forward pass
/// quantities needed for a future BPTT backward.
pub fn gtrxl_block_token_step(
block : GTrXLBlock,
x : Array[Float],
) -> (Array[Float], GTrXLTokenCache) {
// Gate path: sigmoid(W_gate · x + b_gate)
let gate_pre : Array[Float] = matvec(block.ffn_gate_w, block.ffn_gate_b, x)
let gate : Array[Float] = Array::make(block.d_model, 0.0F)
for k in 0.. (Array[Float], GTrXLBlockCache) {
let y : Array[Float] = Array::make(seq_len * block.d_model, 0.0F)
let x_in_buf : Array[Float] = Array::make(seq_len * block.d_model, 0.0F)
let gate_pre_buf : Array[Float] = Array::make(seq_len * block.d_model, 0.0F)
let hidden_pre_buf : Array[Float] = Array::make(seq_len * block.d_ff, 0.0F)
let hidden_post_buf : Array[Float] = Array::make(seq_len * block.d_ff, 0.0F)
let gated_buf : Array[Float] = Array::make(seq_len * block.d_model, 0.0F)
for t in 0..