// gru_policy.mbt — REINFORCE with GRU policy on POMDP corridor (v0.34.0).
//
// Same setup as `lstm_policy.mbt` but with a GRU cell instead of
// LSTM. GRU has no cell state — only the hidden state evolves:
// z_t = σ(W_z · [h_{t-1}; x_t] + b_z)
// r_t = σ(W_r · [h_{t-1}; x_t] + b_r)
// n_t = tanh(W_n · [r_t ⊙ h_{t-1}; x_t] + b_n)
// h_t = (1 - z_t) ⊙ h_{t-1} + z_t ⊙ n_t
//
// The output projection W_out · h_t + b_out is identical to the
// LSTM policy. REINFORCE updates accumulate per-step d_logits
// gradients (using the TD-style ∇log π gradient), then backprop
// through the GRU via `gru_cell_backward`.
///|
/// GRU policy: cell + output projection.
pub struct GruPolicy {
cell : GruCellParam
w_out : Array[Array[Float]] // n_actions × d_h
b_out : Array[Float]
}
///|
/// Build a GRU policy. d_x = 2 + n_cells, d_h = hidden_dim,
/// n_actions = 2 (left/right).
pub fn GruPolicy::new(
n_cells : Int,
d_h : Int,
seed : UInt64,
) -> GruPolicy {
let d_x = 2 + n_cells
let cell = GruCellParam::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 for the GRU policy.
pub struct GruPolicyCache {
cell_caches : Array[GruCellCache]
probs : Array[Array[Float]]
hs : Array[Array[Float]]
}
///|
/// Run the GRU policy forward. Returns (action_seq, cache).
pub fn gru_policy_forward(
policy : GruPolicy,
obs_seq : Array[Array[Float]],
h0 : Array[Float],
rng : Xoshiro,
) -> (Array[Int], GruPolicyCache) {
let n = obs_seq.length()
let cell_caches : Array[GruCellCache] = []
let probs : Array[Array[Float]] = []
let hs : Array[Array[Float]] = []
let actions : Array[Int] = []
let mut cur_h = h0
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 {
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,
}
let obs_seq : Array[Array[Float]] = []
for _t in 0..