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