// beam_search_planner.mbt — Beam search planner for Trajectory
// Transformer (v0.81.0).
//
// The TT itself does NOT directly predict actions (it predicts
// next_state, reward, done, value at each step). So to plan a
// trajectory we maintain K candidate "beams" — each is a partial
// trajectory of (state, action, reward, done) — and at each planning
// step expand them by:
//   1. Forward the TT to predict (next_state, reward, done, value) for
//      the current beam tail.
//   2. For each beam, sample K action candidates (uniform random in
//      [action_low, action_high] for now — could be replaced by a
//      policy proposal later).
//   3. Append (predicted_next_state, sampled_action, predicted_reward,
//      predicted_done) to form K new candidate trajectories.
//   4. Score each new candidate by the predicted value + accumulated
//      discounted reward.
//   5. Keep top-K.
//
// After `horizon` planning steps, return the first action of the
// highest-scoring beam (the action to execute in the environment now).
//
// Reference: Janner et al. 2021 Section 4 "Planning as Conditional
// Sequence Modeling" — beam search with stochastic action proposals.
//
// Scope of v0.81.0:
//   - BeamSearchPlanner struct (wraps a TrajectoryTransformer + K +
//     horizon + gamma)
//   - beam_search_plan: full pipeline from history → best first action

///|
/// One beam in the search. Holds a flat row-major Float buffer plus
/// the running discounted-return score. We store (state, action,
/// reward, done) per step up to the current beam length.
pub struct Beam {
  state_seq : Array[Float]
  action_seq : Array[Float]
  reward_seq : Array[Float]
  done_seq : Array[Float]
  length : Int
  score : Float
}

///|
/// BeamSearchPlanner config.
pub struct BeamSearchPlanner {
  tt : TrajectoryTransformer
  beam_width : Int
  horizon : Int
  gamma : Float
  action_low : Float
  action_high : Float
}

///|
/// Build a fresh BeamSearchPlanner. `beam_width = K`, `horizon = H`,
/// `gamma` is the discount used in accumulated score.
pub fn BeamSearchPlanner::new(
  tt : TrajectoryTransformer,
  beam_width : Int,
  horizon : Int,
  gamma : Float,
  action_low : Float,
  action_high : Float,
) -> BeamSearchPlanner {
  { tt, beam_width, horizon, gamma, action_low, action_high }
}

///|
/// Initialize K beams from a single (state, action, reward, done)
/// history. Each beam starts with the same history but a separate
/// (independent) score of 0.
fn beam_search_init_beams(
  state_seq : Array[Float],
  action_seq : Array[Float],
  reward_seq : Array[Float],
  done_seq : Array[Float],
  history_len : Int,
  state_dim : Int,
  action_dim : Int,
  beam_width : Int,
) -> Array[Beam] {
  let beams : Array[Beam] = Array::make(beam_width, {
    state_seq: Array::make(history_len * state_dim, 0.0F),
    action_seq: Array::make(history_len * action_dim, 0.0F),
    reward_seq: Array::make(history_len, 0.0F),
    done_seq: Array::make(history_len, 0.0F),
    length: 0,
    score: 0.0F,
  })
  for k in 0.. Float {
  let mut s = 0.0F
  let mut discount = 1.0F
  let mut terminated = false
  for t in 0.. 0.5F {
      terminated = true
    }
  }
  s
}

///|
/// Pick the top-K beams by score from a list of K*K candidates.
/// Returns a fresh Array[Beam] of length K containing the highest
/// scorers (ties broken by index).
fn beam_search_topk(
  candidates : Array[Beam],
  beam_width : Int,
) -> Array[Beam] {
  // Sort indices by score descending (insertion sort — fine for K*K small).
  let n = candidates.length()
  let mut indices : Array[Int] = Array::make(n, 0)
  for i in 0.. best_score {
        best_score = s
        best_idx = j
      }
    }
    chosen.push(candidates[indices[best_idx]])
    // Remove the chosen index.
    let new_indices : Array[Int] = []
    for j in 0.. Beam {
  let new_state : Array[Float] = Array::make(
    (beam.length + 1) * state_t.length(), 0.0F,
  )
  for i in 0.. Array[Float] {
  let tt = planner.tt
  let state_dim = tt.state_dim
  let action_dim = tt.action_dim
  // Initial beams.
  let mut beams = beam_search_init_beams(
    state_seq, action_seq, reward_seq, done_seq, history_len,
    state_dim, action_dim, planner.beam_width,
  )
  // Timesteps: simple 0..length mapping (re-built each step).
  let timesteps_buf : Array[Int] = Array::make(planner.horizon + history_len, 0)
  for t in 0..<(planner.horizon + history_len) {
    timesteps_buf[t] = t
  }
  // Expand H steps.
  for _h in 0.. best.score {
      best = beams[k]
    }
  }
  let first_action : Array[Float] = Array::make(action_dim, 0.0F)
  for i in 0.. Float {
  let mut p = 1.0F
  for _i in 0..