// 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..