// lstm_policy.mbt — REINFORCE with LSTM policy on POMDP corridor (v0.33.2).
//
// POMDP corridor:
//   - N cells in a line. At t=0 the agent sees an indicator
//     (`[1, 0]` for "goal on left" or `[0, 1]` for "goal on right").
//     At later steps only the position is observable.
//   - 2 actions: 0=left, 1=right.
//   - Reward: +1 at goal, -0.1 per step.
//
// LSTM policy:
//   - Observation = concat(indicator_2, position_one_hot_N) =
//     length 2 + N.
//   - LSTM cell processes the sequence.
//   - Output projection: logits = W_out · h_t + b_out (n_actions × d_h).
//   - Sample action from softmax(logits).
//
// REINFORCE with LSTM:
//   - Per-step gradient: d_logits[i] = (1{i == a_t} - π(a|h_t)) · G_t
//   - Output layer: d_W_out += d_logits ⊗ h_t, d_b_out += d_logits
//   - LSTM backward: feed d_h_t into lstm_cell_backward at each
//     timestep, accumulating parameter gradients.

///|
/// Corridor POMDP: indicator at t=0 says which side has the goal.
pub(all) struct CorridorEnv {
  n_cells : Int
  goal_side : Int  // 0 = left (cell 0), 1 = right (cell n_cells-1)
  step_penalty : Float
  goal_reward : Float
  max_steps : Int
}

///|
/// Build a corridor with the goal on the right by default.
pub fn CorridorEnv::new(n_cells : Int, max_steps : Int) -> CorridorEnv {
  { n_cells, goal_side: 1, step_penalty: -0.1F, goal_reward: 1.0F, max_steps }
}

///|
/// Goal cell index (0 or n_cells-1).
pub fn CorridorEnv::goal_cell(self : CorridorEnv) -> Int {
  if self.goal_side == 0 {
    0
  } else {
    self.n_cells - 1
  }
}

///|
/// Build the observation vector at time t. Concat of
/// (indicator, position_one_hot).
pub fn corridor_observation(
  env : CorridorEnv,
  position : Int,
) -> Array[Float] {
  let obs : Array[Float] = Array::make(2 + env.n_cells, 0.0F)
  if env.goal_side == 0 {
    obs[0] = 1.0F
  } else {
    obs[1] = 1.0F
  }
  if position >= 0 && position < env.n_cells {
    obs[2 + position] = 1.0F
  }
  obs
}

///|
/// Step the corridor. pos = current position, action = 0 (left)
/// or 1 (right). Going off the grid keeps the agent in place.
pub fn CorridorEnv::step(self : CorridorEnv, pos : Int, action : Int) -> (Int, Float, Bool) {
  let mut new_pos = pos
  if action == 0 {
    if pos > 0 {
      new_pos = pos - 1
    }
  } else if action == 1 {
    if pos < self.n_cells - 1 {
      new_pos = pos + 1
    }
  }
  let goal = self.goal_cell()
  if new_pos == goal {
    (new_pos, self.goal_reward, true)
  } else {
    (new_pos, self.step_penalty, false)
  }
}

///|
/// LSTM policy: cell + output projection.
pub struct LstmPolicy {
  cell : LstmCellParam
  w_out : Array[Array[Float]]  // n_actions × d_h
  b_out : Array[Float]
}

///|
/// Build an LSTM policy. d_x = 2 + n_cells, d_h = hidden_dim,
/// n_actions = 2 (left/right).
pub fn LstmPolicy::new(
  n_cells : Int,
  d_h : Int,
  seed : UInt64,
) -> LstmPolicy {
  let d_x = 2 + n_cells
  let cell = LstmCellParam::new(d_x, d_h, seed)
  let n_actions = 2
  let std = sqrtf(1.0F / Float::from_int(d_h))
  let rng = Xoshiro::from_state(seed + 1UL, seed + 2UL, seed + 3UL, seed + 4UL)
  let w_out = xavier_normal(n_actions, d_h, std, rng)
  let b_out : Array[Float] = Array::make(n_actions, 0.0F)
  { cell, w_out, b_out }
}

///|
/// Forward cache: per-step cell caches + per-step projected info.
pub struct LstmPolicyCache {
  cell_caches : Array[LstmCellCache]
  probs : Array[Array[Float]]
  hs : Array[Array[Float]]
}

///|
/// Run the LSTM policy forward. Returns (action_seq, cache).
pub fn lstm_policy_forward(
  policy : LstmPolicy,
  obs_seq : Array[Array[Float]],
  h0 : Array[Float],
  c0 : Array[Float],
  rng : Xoshiro,
) -> (Array[Int], LstmPolicyCache) {
  let n = obs_seq.length()
  let cell_caches : Array[LstmCellCache] = []
  let probs : Array[Array[Float]] = []
  let hs : Array[Array[Float]] = []
  let actions : Array[Int] = []
  let mut cur_h = h0
  let mut cur_c = c0
  for t in 0.. max_l {
        max_l = logits[i]
      }
    }
    let ps : Array[Float] = Array::make(n_actions, 0.0F)
    let mut sum_e = 0.0F
    for i in 0.. Unit {
  // Use the min length so we don't read past the end of `returns`
  // when the env terminates before `max_steps`.
  let actions_n = actions.length()
  let returns_n = returns.length()
  let n = if actions_n < returns_n { actions_n } else { returns_n }
  let n_actions = policy.w_out.length()
  let d_h = policy.cell.d_h
  let d_w_out : Array[Array[Float]] = Array::make(n_actions, [])
  for i in 0.. Unit {
  for i in 0.. Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let mut total_return = 0.0F
  let start = env_n_cells / 2
  for _ep in 0.. 0.0 { 1 } else { 0 }
    let env : CorridorEnv = CorridorEnv::{
      n_cells: env_n_cells,
      goal_side: side,
      step_penalty: -0.1F,
      goal_reward: 1.0F,
      max_steps,
    }
    // Build a fixed-length obs_seq: indicator is always present,
    // position is the start (so the LSTM must rely on the
    // indicator alone to decide direction).
    let obs_seq : Array[Array[Float]] = []
    for _t in 0.. Float {
  let mut s = 0.0F
  for r in rewards {
    s = s + r
  }
  s
}