// decision_transformer.mbt — Decision Transformer model: full forward
// (tokenize trajectory → GTrXL block → action prediction head)
// (v0.77.0).
//
// Architecture:
//   tokens_t  =  DTEmbeddings(rtg_t, state_t, action_t, timestep_t)    ∈ R^{d_model}
//   hidden_t  =  GTrXLBlock(tokens_t)                       ∈ R^{d_model}
//   a_pre_t   =  W_action · hidden_t + b_action             ∈ R^{action_dim}
//
// Loss: MSE between a_pre_t and the actual action_t. The MSE is
// computed by the trainer (v0.79.0); this file only ships forward.
//
// Reference: Chen et al. 2021 "Decision Transformer: Reinforcement
// Learning through Sequence Modeling". The original DT uses GPT-style
// causal attention; here we substitute a GTrXL block as the per-token
// recurrent memory primitive (the GTrXL block's gated residual update
// provides token-to-token carry).

///|
/// Decision Transformer model. Combines DTEmbeddings + GTrXLBlock +
/// action prediction head. The action head at the *last* timestep is
/// the policy output (the predicted next action for the desired return);
/// earlier timesteps are also predicted (used as MSE targets during
/// training).
pub struct DecisionTransformer {
  state_dim : Int
  action_dim : Int
  d_model : Int
  d_ff : Int
  max_timestep : Int
  embeddings : DTEmbeddings
  backbone : GTrXLBlock
  action_w : Array[Array[Float]]
  action_b : Array[Float]
}

///|
/// Build fresh DecisionTransformer.
///   - embeddings: DTEmbeddings::new with seed
///   - backbone: GTrXLBlock::new(d_model, d_ff, seed + 4UL)
///   - action_w: (action_dim × d_model), Xavier-normal scaled by sqrt(2/d_model)
///   - action_b: zeros
pub fn DecisionTransformer::new(
  state_dim : Int,
  action_dim : Int,
  d_model : Int,
  d_ff : Int,
  max_timestep : Int,
  seed : UInt64,
) -> DecisionTransformer {
  let embeddings = DTEmbeddings::new(
    state_dim, action_dim, d_model, max_timestep, seed,
  )
  let backbone = GTrXLBlock::new(d_model, d_ff, seed + 16UL)
  let rng = Xoshiro::from_state(seed + 20UL, seed + 21UL, seed + 22UL, seed + 23UL)
  let std_head = sqrtf(2.0F / Float::from_int(d_model))
  let action_w = xavier_normal(action_dim, d_model, std_head, rng)
  let action_b : Array[Float] = Array::make(action_dim, 0.0F)
  {
    state_dim,
    action_dim,
    d_model,
    d_ff,
    max_timestep,
    embeddings,
    backbone,
    action_w,
    action_b,
  }
}

///|
/// Full Decision Transformer forward. Returns
/// `(predicted_action_seq, block_cache)` where:
///   - `predicted_action_seq` is flat `[seq_len × action_dim]`
///   - `block_cache` is the GTrXLBlockCache for future BPTT
pub fn dt_forward(
  dt : DecisionTransformer,
  rtg_seq : Array[Float],
  state_seq : Array[Float],
  action_seq : Array[Float],
  timesteps : Array[Int],
  seq_len : Int,
) -> (Array[Float], GTrXLBlockCache) {
  // Tokenize the full trajectory.
  let tokens = dt_embed_trajectory(
    dt.embeddings, rtg_seq, state_seq, action_seq, timesteps, seq_len,
  )
  // Run through the GTrXL backbone.
  let (hidden_seq, cache) = gtrxl_block_seq_forward(
    dt.backbone, tokens, seq_len,
  )
  // Action prediction head.
  let predicted_action_seq : Array[Float] = Array::make(
    seq_len * dt.action_dim, 0.0F,
  )
  for t in 0.. Array[Float] {
  let (predicted_action_seq, _) = dt_forward(
    dt, rtg_seq, state_seq, action_seq, timesteps, seq_len,
  )
  let off = (seq_len - 1) * dt.action_dim
  let action : Array[Float] = Array::make(dt.action_dim, 0.0F)
  for i in 0..