// gtrxl_qnetwork_continuous.mbt — Recurrent continuous-action Q-network
// with GTrXL block as recurrent memory (v0.74.0).
//
// Architecture: MLP w1 (concat(state, action) -> hidden) -> ReLU ->
// GTrXL block (maintains a d_model-dimensional hidden state across
// time via gated residual updates) -> MLP w2 (hidden -> scalar Q).
//
// [state_t ; action_t] ∈ R^{state_dim + action_dim}
// x_proj_t = ReLU(w1 · [state_t ; action_t] + b1) ∈ R^{d_model}
// hidden_t = gtrxl_block_token_step(x_proj_t).y ∈ R^{d_model}
// q_t = w2 · hidden_t + b2 ∈ R
//
// This is the GTrXL counterpart of `LSTMQNetworkContinuous` (v0.62.0)
// for the GTrXL_DDPG agent (v0.74.0). The recurrent memory is the
// GTrXL block's per-token gated residual update (carried token-to-token
// via the y → x residual).
//
// Reference: Parisotto et al. 2020; twin-critic TD pattern from
// Fujimoto et al. 2018 (TD3) adapted to recurrent actors.
///|
/// Recurrent continuous-action Q-network. MLP hidden dim equals the
/// GTrXL block's d_model so the ReLU-projected [state; action] and the
/// block's token dim match. `mlp_b2` is a scalar (the critic outputs a
/// single Q-value per timestep).
pub struct GTrXLQNetworkContinuous {
state_dim : Int
action_dim : Int
d_model : Int
d_ff : Int
mlp_w1 : Array[Array[Float]]
mlp_b1 : Array[Float]
block : GTrXLBlock
mlp_w2 : Array[Array[Float]]
mut mlp_b2 : Float
}
///|
/// Build a fresh GTrXLQNetworkContinuous. MLP weights use
/// Xavier-normal init scaled by sqrtf(2 / fan_in) (He-style for ReLU).
/// GTrXL block uses its own init. Zero biases.
pub fn GTrXLQNetworkContinuous::new(
state_dim : Int,
action_dim : Int,
d_model : Int,
d_ff : Int,
seed : UInt64,
) -> GTrXLQNetworkContinuous {
let in_dim = state_dim + action_dim
let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let std1 = sqrtf(2.0F / Float::from_int(in_dim))
let mlp_w1 = xavier_normal(d_model, in_dim, std1, rng1)
let mlp_b1 : Array[Float] = Array::make(d_model, 0.0F)
let block = GTrXLBlock::new(d_model, d_ff, seed + 4UL)
let rng2 = Xoshiro::from_state(seed + 8UL, seed + 9UL, seed + 10UL, seed + 11UL)
let std2 = sqrtf(2.0F / Float::from_int(d_model))
let mlp_w2 = xavier_normal(1, d_model, std2, rng2)
{
state_dim,
action_dim,
d_model,
d_ff,
mlp_w1,
mlp_b1,
block,
mlp_w2,
mlp_b2: 0.0F,
}
}
///|
/// Single-step forward. Returns the scalar Q-value for the given
/// (state, action) pair plus the per-step GTrXL intermediates for a
/// future BPTT backward.
pub fn gtrxl_qnetwork_continuous_step(
qnet : GTrXLQNetworkContinuous,
state : Array[Float],
action : Array[Float],
) -> (Float, GTrXLTokenCache) {
// sa = concat(state, action)
let sa = vec_concat(state, action)
// x_proj = ReLU(w1 · sa + b1)
let x_proj_pre = matvec(qnet.mlp_w1, qnet.mlp_b1, sa)
let x_proj = relu_forward(x_proj_pre)
// hidden_next = gtrxl_block_token_step(x_proj)
let (hidden_next, cache) = gtrxl_block_token_step(qnet.block, x_proj)
// q = w2 · hidden_next + b2
let q_pre = matvec(qnet.mlp_w2, [qnet.mlp_b2], hidden_next)
(q_pre[0], cache)
}
///|
/// Sequence forward. `state_seq` is flat row-major
/// `[seq_len × state_dim]`, `action_seq` is flat row-major
/// `[seq_len × action_dim]`. Returns `(q_seq, full_cache)` where
/// `q_seq` is flat `[seq_len]` and `full_cache` stitches per-step
/// GTrXL intermediates into a single record for future BPTT.
pub fn gtrxl_qnetwork_continuous_seq_forward(
qnet : GTrXLQNetworkContinuous,
state_seq : Array[Float],
action_seq : Array[Float],
seq_len : Int,
) -> (Array[Float], GTrXLBlockCache) {
let q_seq : Array[Float] = Array::make(seq_len, 0.0F)
let x_in_buf : Array[Float] = Array::make(seq_len * qnet.d_model, 0.0F)
let gate_pre_buf : Array[Float] = Array::make(
seq_len * qnet.d_model, 0.0F,
)
let hidden_pre_buf : Array[Float] = Array::make(seq_len * qnet.d_ff, 0.0F)
let hidden_post_buf : Array[Float] = Array::make(
seq_len * qnet.d_ff, 0.0F,
)
let gated_buf : Array[Float] = Array::make(seq_len * qnet.d_model, 0.0F)
for t in 0..