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