// trajectory_transformer_agent.mbt — Trajectory Transformer agent:
// full planning pipeline (history → beam search → best action) (v0.83.0).
//
// Wires:
// - TrajectoryTransformer (v0.80.0) for the model
// - BeamSearchPlanner (v0.81.0) for the planning algorithm
//
// The agent maintains a sliding-window history of (state, action,
// reward, done). On each environment step, it appends the new
// transition, runs beam search over the history, and returns the
// first action of the highest-scoring planned trajectory.
//
// Scope of v0.83.0:
// - TrajectoryTransformerAgent struct + constructor
// - tt_agent_step: append transition + run beam search + return action
// - tt_agent_reset: clear history (called at episode boundary)
//
// Reference: Janner et al. 2021, Section 4 (the 'TT pipeline' for
// model-based control).
///|
/// Trajectory Transformer agent: model + planner + sliding history.
pub(all) struct TrajectoryTransformerAgent {
tt : TrajectoryTransformer
planner : BeamSearchPlanner
context_length : Int
action_dim : Int
state_dim : Int
// Sliding window history (state, action, reward, done). When the
// history exceeds `context_length`, the oldest entry is dropped.
mut history_states : Array[Float]
mut history_actions : Array[Float]
mut history_rewards : Array[Float]
mut history_dones : Array[Float]
mut history_len : Int
}
///|
/// Build a fresh TrajectoryTransformerAgent. `context_length` is the
/// number of past steps fed into the planner at inference time
/// (the planner sees this many transitions before planning forward).
pub fn TrajectoryTransformerAgent::new(
tt : TrajectoryTransformer,
planner : BeamSearchPlanner,
context_length : Int,
) -> TrajectoryTransformerAgent {
let state_dim = tt.state_dim
let action_dim = tt.action_dim
{
tt,
planner,
context_length,
action_dim,
state_dim,
history_states: Array::make(context_length * state_dim, 0.0F),
history_actions: Array::make(context_length * action_dim, 0.0F),
history_rewards: Array::make(context_length, 0.0F),
history_dones: Array::make(context_length, 0.0F),
history_len: 0,
}
}
///|
/// Append one transition (state, action, reward, done) to the sliding
/// history. If history is full, the oldest entry is dropped. The
/// transition added here is the (s_t, a_{t-1}, r_t, d_t) tuple — i.e.
/// the action taken in the *previous* state, observed reward, and done
/// flag. The caller passes `prev_action` (the action executed before
/// observing `state`).
fn tt_agent_append(
agent : TrajectoryTransformerAgent,
state : Array[Float],
prev_action : Array[Float],
prev_reward : Float,
prev_done : Float,
) -> Unit {
let cl = agent.context_length
if agent.history_len < cl {
// Append at the tail.
let t = agent.history_len
let st_off = t * agent.state_dim
let at_off = t * agent.action_dim
for k in 0.. Array[Float] {
tt_agent_append(agent, state, prev_action, prev_reward, prev_done)
// Use only the first history_len entries for the planner.
let cl = agent.context_length
let state_seq : Array[Float] = Array::make(
agent.history_len * agent.state_dim, 0.0F,
)
let action_seq : Array[Float] = Array::make(
agent.history_len * agent.action_dim, 0.0F,
)
let reward_seq : Array[Float] = Array::make(agent.history_len, 0.0F)
let done_seq : Array[Float] = Array::make(agent.history_len, 0.0F)
let n_states_copy = agent.history_len * agent.state_dim
for i in 0.. Unit {
agent.history_len = 0
// Zero out the buffers (optional but keeps state clean).
for i in 0..