// gtrxl_actor.mbt — Recurrent deterministic actor with GTrXL block as
// recurrent memory (v0.73.0).
//
// Architecture: state_t → MLP w1 → ReLU → embed_t (length d_model) →
// GTrXL block (maintains a d_model-dimensional hidden state across time
// via gated residual updates) → MLP w2 → tanh squash to [action_low,
// action_high].
//
//   state_t ∈ R^{state_dim}
//   embed_t = ReLU(w1 · state_t + b1)                  ∈ R^{d_model}
//   hidden_t = gtrxl_block_token_step(embed_t).y       ∈ R^{d_model}
//   a_pre_t = w2 · hidden_t + b2                       ∈ R^{action_dim}
//   action_t = tanh(a_pre_t) * (high - low) / 2
//            + (high + low) / 2                        ∈ R^{action_dim}
//
// This is the GTrXL counterpart of `GRUDeterministicPolicy` (v0.58.0)
// and `LSTMDeterministicPolicy` (v0.61.0). The recurrent memory here
// is the GTrXL block's per-token gated residual update (the GTrXL
// block has its own internal state in the form of `x_t`-to-`x_t`
// identity residual; we feed it the embedding from the MLP).
//
// Reference: Parisotto et al. 2020 "Stabilizing Transformers for
// Reinforcement Learning"; Hausknecht & Stone 2015 for the recurrent
// actor pattern.

///|
/// Recurrent deterministic actor. The GTrXL block's d_model equals the
/// MLP hidden dimension so the ReLU-projected state and the GTrXL
/// block's token dim match.
pub struct GTrXLDeterministicPolicy {
  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]]
  mlp_b2 : Array[Float]
  action_low : Float
  action_high : Float
}

///|
/// Build a fresh GTrXLDeterministicPolicy. The MLP weights get
/// Xavier-normal init scaled by `sqrtf(2.0 / fan_in)` (He-style for
/// ReLU), and the GTrXL block uses its own init (Xavier-normal across
/// ffn_w1/ffn_w2 + 0.1×std for the gate matrix). Zero biases.
/// `action_low` / `action_high` define the squash range.
pub fn GTrXLDeterministicPolicy::new(
  state_dim : Int,
  action_dim : Int,
  d_model : Int,
  d_ff : Int,
  action_low : Float,
  action_high : Float,
  seed : UInt64,
) -> GTrXLDeterministicPolicy {
  let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let mlp_w1_std = sqrtf(2.0F / Float::from_int(state_dim))
  let mlp_w1 = xavier_normal(d_model, state_dim, mlp_w1_std, 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 mlp_w2_std = sqrtf(2.0F / Float::from_int(d_model))
  let mlp_w2 = xavier_normal(action_dim, d_model, mlp_w2_std, rng2)
  let mlp_b2 : Array[Float] = Array::make(action_dim, 0.0F)
  {
    state_dim,
    action_dim,
    d_model,
    d_ff,
    mlp_w1,
    mlp_b1,
    block,
    mlp_w2,
    mlp_b2,
    action_low,
    action_high,
  }
}

///|
/// Single-step forward. `obs` is the current observation (length
/// state_dim). Returns `(action, block_cache)` where `block_cache`
/// holds the GTrXL per-step intermediates for a future BPTT
/// backward, and `action` is the squashed continuous action (length
/// action_dim).
pub fn gtrxl_actor_step(
  policy : GTrXLDeterministicPolicy,
  obs : Array[Float],
) -> (Array[Float], GTrXLTokenCache) {
  // embed = ReLU(w1 · obs + b1)
  let x_proj_pre = matvec(policy.mlp_w1, policy.mlp_b1, obs)
  let x_proj = relu_forward(x_proj_pre)
  // hidden_next = gtrxl_block_token_step(embed)
  let (hidden_next, cache) = gtrxl_block_token_step(policy.block, x_proj)
  // a_pre = w2 · hidden_next + b2
  let a_pre = matvec(policy.mlp_w2, policy.mlp_b2, hidden_next)
  // squash: action = tanh(a_pre) * (high - low) / 2 + (high + low) / 2
  let half_range = (policy.action_high - policy.action_low) * 0.5F
  let mid = (policy.action_high + policy.action_low) * 0.5F
  let action : Array[Float] = Array::make(policy.action_dim, 0.0F)
  for i in 0.. (Array[Float], GTrXLBlockCache) {
  let action_seq : Array[Float] = Array::make(
    seq_len * policy.action_dim, 0.0F,
  )
  // Stitch per-step caches into a single record. The block receives
  // the freshly-embedded token each step (no carry-over of an external
  // hidden state — the gated-residual inside the GTrXL block is the
  // recurrent memory itself, propagated token-to-token via `y`).
  let x_in_buf : Array[Float] = Array::make(seq_len * policy.d_model, 0.0F)
  let gate_pre_buf : Array[Float] = Array::make(
    seq_len * policy.d_model, 0.0F,
  )
  let hidden_pre_buf : Array[Float] = Array::make(seq_len * policy.d_ff, 0.0F)
  let hidden_post_buf : Array[Float] = Array::make(
    seq_len * policy.d_ff, 0.0F,
  )
  let gated_buf : Array[Float] = Array::make(seq_len * policy.d_model, 0.0F)
  for t in 0..