// decision_transformer_trainer.mbt — Decision Transformer offline
// training step (v0.79.0).
//
// Performs one supervised learning step on a sampled trajectory window:
//   loss = (1 / N) · Σ_t || a_pre_t - a_t ||²
// where `a_pre_t` is the DT's action-head prediction and `a_t` is the
// ground-truth action from the trajectory buffer.
//
// Scope of v0.79.0:
//   - dt_mse_loss: scalar MSE between predicted and actual action seqs
//   - dt_train_step: one SGD step on the DT's action head (action_w,
//     action_b). Returns the updated DecisionTransformer + scalar loss.
//   - The DT embeddings + GTrXL backbone gradients are deferred (BPTT
//     through GTrXL block not yet implemented; see Batch G deferral).
//
// Reference: Chen et al. 2021, Section 4 "Training": pure supervised
// learning with MSE on action prediction, no critic, no bootstrapping.

///|
/// Compute the per-element MSE loss between predicted and target action
/// sequences. Shapes: both are flat `[seq_len × action_dim]`. Returns
/// (1 / N) · Σ (pred - target)² as a Float.
pub fn dt_mse_loss(
  predicted_action_seq : Array[Float],
  target_action_seq : Array[Float],
  n_elements : Int,
) -> Float {
  if n_elements <= 0 {
    return 0.0F
  }
  let mut sum_sq = 0.0F
  for i in 0.. Array[Float] {
  let grad : Array[Float] = Array::make(n_elements, 0.0F)
  if n_elements <= 0 {
    return grad
  }
  let scale = 2.0F / Float::from_int(n_elements)
  for i in 0.. (Array[Array[Float]], Array[Float]) {
  let d_w : Array[Array[Float]] = Array::make(
    action_dim, Array::make(d_model, 0.0F),
  )
  let d_b : Array[Float] = Array::make(action_dim, 0.0F)
  for t in 0.. (DecisionTransformer, Float) {
  // 1. Forward: get the predicted action sequence from dt_forward's
  //    intermediate (caller already computed hidden_seq).
  //    To avoid recomputing forward here, the caller passes the
  //    forward's predicted_action_seq via `prev_action_seq`.
  // 2. Loss.
  let n_elements = seq_len * dt.action_dim
  let loss = dt_mse_loss(prev_action_seq, target_action_seq, n_elements)
  // 3. Gradient w.r.t. action head.
  let grad_action = dt_mse_grad_action(prev_action_seq, target_action_seq, n_elements)
  let (d_w, d_b) = dt_action_head_grad(
    grad_action, hidden_seq, seq_len, dt.action_dim, dt.d_model,
  )
  // 4. SGD step on action_w, action_b.
  let (new_action_w, new_action_b) = sgd_update_arrays(
    flatten_2d(dt.action_w), dt.action_b, flatten_2d(d_w), d_b, lr,
  )
  let new_action_w_2d = unflatten_2d(new_action_w, dt.action_dim, dt.d_model)
  let updated : DecisionTransformer = { ..dt, action_w: new_action_w_2d, action_b: new_action_b }
  (updated, loss)
}

// Helper: flatten a 2D row-major Array[Array[Float]] into a 1D Array[Float].
fn flatten_2d(a : Array[Array[Float]]) -> Array[Float] {
  let mut total = 0
  for i in 0.. Array[Array[Float]] {
  let out : Array[Array[Float]] = Array::make(
    rows, Array::make(cols, 0.0F),
  )
  let mut off = 0
  for i in 0..