// trajectory_transformer.mbt — Trajectory Transformer primitive:
// 4-modality (state, action, reward, done) tokenization + 4 prediction
// heads (next_state, reward, done, value) on a GTrXL backbone (v0.80.0).
//
// Reference: Janner et al. 2021 "Trajectory Transformer: One Model to
// Plan Them All". The original TT uses a single transformer to model
// the joint distribution over (states, actions, rewards, done-flags) in
// a continuous-control trajectory. Here we substitute a GTrXL block as
// the per-token recurrent memory primitive (same substitution as
// Decision Transformer in Batch H).
//
// Per-timestep token layout (similar to DT but with reward + done):
//   token_t = Linear_state(s_t) + Linear_action(a_t) + Linear_reward(r_t) + Linear_done(d_t) + Embed_timestep(t)
//
// Four prediction heads (applied to the post-GTrXL hidden state):
//   next_state_t = W_state · hidden_t + b_state     ∈ R^{state_dim}
//   reward_t     = W_reward · hidden_t + b_reward   ∈ R
//   done_t       = W_done · hidden_t + b_done       ∈ R
//   value_t      = W_value · hidden_t + b_value     ∈ R
//
// Scope of v0.80.0:
//   - TTEmbeddings struct + constructor (4 modality projections + timestep)
//   - TrajectoryTransformer struct + constructor (embeddings + GTrXLBlock + 4 heads)
//   - tt_forward: full forward pass — tokenize → GTrXL → 4 heads
//   - tt_predict_next (helper): one-step prediction for inference
//   - tt_predict_value (helper): value-only prediction at a given position

///|
/// Trajectory Transformer token embeddings. Four modality projections
/// (state, action, reward, done) + a timestep embedding lookup. All
/// five add into the same d_model vec.
pub struct TTEmbeddings {
  state_dim : Int
  action_dim : Int
  d_model : Int
  max_timestep : Int
  state_w : Array[Array[Float]]
  action_w : Array[Array[Float]]
  reward_w : Array[Array[Float]]
  done_w : Array[Array[Float]]
  timestep_emb : Array[Array[Float]]
}

///|
/// Build fresh TTEmbeddings.
///   - state_w: (d_model × state_dim), Xavier-normal scaled by sqrt(2/state_dim)
///   - action_w: (d_model × action_dim), Xavier-normal scaled by sqrt(2/action_dim)
///   - reward_w: (d_model × 1), Xavier-normal scaled by sqrt(2/1)
///   - done_w: (d_model × 1), Xavier-normal scaled by sqrt(2/1)
///   - timestep_emb: (max_timestep × d_model), std = 0.02
pub fn TTEmbeddings::new(
  state_dim : Int,
  action_dim : Int,
  d_model : Int,
  max_timestep : Int,
  seed : UInt64,
) -> TTEmbeddings {
  let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
  let std_state = sqrtf(2.0F / Float::from_int(state_dim))
  let state_w = xavier_normal(d_model, state_dim, std_state, rng1)
  let rng2 = Xoshiro::from_state(seed + 4UL, seed + 5UL, seed + 6UL, seed + 7UL)
  let std_action = sqrtf(2.0F / Float::from_int(action_dim))
  let action_w = xavier_normal(d_model, action_dim, std_action, rng2)
  let rng3 = Xoshiro::from_state(seed + 8UL, seed + 9UL, seed + 10UL, seed + 11UL)
  let std_scalar = sqrtf(2.0F / 1.0F)
  let reward_w = xavier_normal(d_model, 1, std_scalar, rng3)
  let rng4 = Xoshiro::from_state(seed + 12UL, seed + 13UL, seed + 14UL, seed + 15UL)
  let done_w = xavier_normal(d_model, 1, std_scalar, rng4)
  let rng5 = Xoshiro::from_state(seed + 16UL, seed + 17UL, seed + 18UL, seed + 19UL)
  let std_ts = 0.02F
  let timestep_emb = xavier_normal(max_timestep, d_model, std_ts, rng5)
  {
    state_dim,
    action_dim,
    d_model,
    max_timestep,
    state_w,
    action_w,
    reward_w,
    done_w,
    timestep_emb,
  }
}

///|
/// Embed a single TT trajectory token. All four modality inputs are
/// vector projections; `timestep_t` is an Int in `[0, max_timestep)`.
/// Returns Array[Float] of length d_model.
pub fn tt_embed_single_token(
  emb : TTEmbeddings,
  state_t : Array[Float],
  action_t : Array[Float],
  reward_t : Float,
  done_t : Float,
  timestep_t : Int,
) -> Array[Float] {
  let state_proj : Array[Float] = Array::make(emb.d_model, 0.0F)
  for i in 0.. Array[Float] {
  let tokens : Array[Float] = Array::make(seq_len * emb.d_model, 0.0F)
  for t in 0.. TrajectoryTransformer {
  let embeddings = TTEmbeddings::new(
    state_dim, action_dim, d_model, max_timestep, seed,
  )
  let backbone = GTrXLBlock::new(d_model, d_ff, seed + 20UL)
  let rng1 = Xoshiro::from_state(seed + 24UL, seed + 25UL, seed + 26UL, seed + 27UL)
  let std_h = sqrtf(2.0F / Float::from_int(d_model))
  let state_w = xavier_normal(state_dim, d_model, std_h, rng1)
  let state_b : Array[Float] = Array::make(state_dim, 0.0F)
  let rng2 = Xoshiro::from_state(seed + 28UL, seed + 29UL, seed + 30UL, seed + 31UL)
  let reward_w = xavier_normal(1, d_model, std_h, rng2)
  let rng3 = Xoshiro::from_state(seed + 32UL, seed + 33UL, seed + 34UL, seed + 35UL)
  let done_w = xavier_normal(1, d_model, std_h, rng3)
  let rng4 = Xoshiro::from_state(seed + 36UL, seed + 37UL, seed + 38UL, seed + 39UL)
  let value_w = xavier_normal(1, d_model, std_h, rng4)
  {
    state_dim,
    action_dim,
    d_model,
    d_ff,
    max_timestep,
    embeddings,
    backbone,
    state_w,
    state_b,
    reward_w,
    reward_b: 0.0F,
    done_w,
    done_b: 0.0F,
    value_w,
    value_b: 0.0F,
  }
}

///|
/// Full Trajectory Transformer forward pass. Returns:
///   - pred_next_state_seq: [seq_len × state_dim]
///   - pred_reward_seq:     [seq_len]
///   - pred_done_seq:       [seq_len]
///   - pred_value_seq:      [seq_len]
///   - block_cache:         GTrXLBlockCache for future BPTT
pub fn tt_forward(
  tt : TrajectoryTransformer,
  state_seq : Array[Float],
  action_seq : Array[Float],
  reward_seq : Array[Float],
  done_seq : Array[Float],
  timesteps : Array[Int],
  seq_len : Int,
) -> (
  Array[Float],
  Array[Float],
  Array[Float],
  Array[Float],
  GTrXLBlockCache,
) {
  let tokens = tt_embed_trajectory(
    tt.embeddings, state_seq, action_seq, reward_seq, done_seq,
    timesteps, seq_len,
  )
  let (hidden_seq, cache) = gtrxl_block_seq_forward(
    tt.backbone, tokens, seq_len,
  )
  // Head projections.
  let pred_next_state_seq : Array[Float] = Array::make(
    seq_len * tt.state_dim, 0.0F,
  )
  let pred_reward_seq : Array[Float] = Array::make(seq_len, 0.0F)
  let pred_done_seq : Array[Float] = Array::make(seq_len, 0.0F)
  let pred_value_seq : Array[Float] = Array::make(seq_len, 0.0F)
  for t in 0.. (Array[Float], Float, Float, Float) {
  let (ns_seq, r_seq, d_seq, v_seq, _) = tt_forward(
    tt, state_seq, action_seq, reward_seq, done_seq, timesteps, seq_len,
  )
  let off = (seq_len - 1) * tt.state_dim
  let next_state : Array[Float] = Array::make(tt.state_dim, 0.0F)
  for i in 0.. Float {
  let (_, _, _, v_seq, _) = tt_forward(
    tt, state_seq, action_seq, reward_seq, done_seq, timesteps, seq_len,
  )
  v_seq[seq_len - 1]
}