// lstm_deterministic_policy.mbt — Recurrent deterministic policy for
// partially-observable continuous-control RL (v0.61.0).
//
// Architecture: MLP w1 (state → hidden) → ReLU → LSTM cell (over time,
// tracking both hidden state h_t and cell state c_t) → MLP w2
// (hidden → action) → tanh squash to [action_low, action_high].
//
// state_t ∈ R^{state_dim}
// x_proj_t = ReLU(w1 · state_t + b1) ∈ R^{lstm_hidden}
// (h_t, c_t) = LSTM_cell(x_proj_t, h_{t-1}, c_{t-1})
// a_pre_t = w2 · h_t + b2 ∈ R^{action_dim}
// action_t = tanh(a_pre_t) * (high - low) / 2
// + (high + low) / 2 ∈ R^{action_dim}
//
// This is the LSTM counterpart of `GRUDeterministicPolicy` (v0.58.0)
// for environments where the observation alone is insufficient and a
// recurrent summary (with explicit memory cell) is needed. The LSTM
// maintains both a hidden state h_t and a cell state c_t (vanilla
// LSTM with 4 gates: forget / input / candidate / output).
//
// Reference: Hausknecht & Stone 2015 "Deep Recurrent Q-Learning for
// Partially Observable MDPs" (the recurrent actor pattern); the
// LSTM variant follows Gers 2000 "Learning to Forget".
///|
/// Recurrent deterministic policy. The LSTM hidden dimension equals
/// the MLP hidden dimension so the ReLU-projected state and the LSTM
/// cell input/output match.
pub struct LSTMDeterministicPolicy {
state_dim : Int
action_dim : Int
lstm_hidden : Int
mlp_w1 : Array[Array[Float]]
mlp_b1 : Array[Float]
lstm : LstmCellParam
mlp_w2 : Array[Array[Float]]
mlp_b2 : Array[Float]
action_low : Float
action_high : Float
}
///|
/// Build a fresh LSTMDeterministicPolicy. The MLP weights get
/// Xavier-normal init scaled by `sqrtf(2.0 / fan_in)` (He-style for
/// ReLU), and the LSTM cell reuses its own init (Xavier-normal across
/// the 4 gate matrices). Zero biases. `action_low` / `action_high`
/// define the squash range.
pub fn LSTMDeterministicPolicy::new(
state_dim : Int,
action_dim : Int,
lstm_hidden : Int,
action_low : Float,
action_high : Float,
seed : UInt64,
) -> LSTMDeterministicPolicy {
let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let mlp_w1_std = sqrtf(2.0F / Float::from_int(state_dim))
let mlp_w1 = xavier_normal(lstm_hidden, state_dim, mlp_w1_std, rng1)
let mlp_b1 : Array[Float] = Array::make(lstm_hidden, 0.0F)
let lstm = LstmCellParam::new(lstm_hidden, lstm_hidden, seed + 4UL)
let rng2 = Xoshiro::from_state(seed + 5UL, seed + 6UL, seed + 7UL, seed + 8UL)
let mlp_w2_std = sqrtf(2.0F / Float::from_int(lstm_hidden))
let mlp_w2 = xavier_normal(action_dim, lstm_hidden, mlp_w2_std, rng2)
let mlp_b2 : Array[Float] = Array::make(action_dim, 0.0F)
{
state_dim,
action_dim,
lstm_hidden,
mlp_w1,
mlp_b1,
lstm,
mlp_w2,
mlp_b2,
action_low,
action_high,
}
}
///|
/// Single-step forward. `obs` is the current observation (length
/// state_dim), `hidden` is the previous LSTM hidden state (length
/// lstm_hidden), `cell` is the previous LSTM cell state (length
/// lstm_hidden). Returns `(action, new_hidden, new_cell)` where
/// action is the squashed continuous action (length action_dim).
pub fn lstm_deterministic_policy_step(
policy : LSTMDeterministicPolicy,
obs : Array[Float],
hidden : Array[Float],
cell : Array[Float],
) -> (Array[Float], Array[Float], Array[Float]) {
// x_proj = ReLU(w1 · obs + b1)
let x_proj_pre = matvec(policy.mlp_w1, policy.mlp_b1, obs)
let x_proj = relu_forward(x_proj_pre)
// (hidden_next, cell_next) = LSTM_cell(x_proj, hidden, cell)
let (hidden_next, cache) = lstm_cell_forward(
x_proj, hidden, cell, policy.lstm,
)
let cell_next = cache.c_t
// a_pre = w2 · hidden_next + b2
let a_pre = matvec(policy.mlp_w2, policy.mlp_b2, hidden_next)
// squash: action = tanh(a_pre) * (high - low) / 2 + (high + low) / 2
let half_range = (policy.action_high - policy.action_low) * 0.5F
let mid = (policy.action_high + policy.action_low) * 0.5F
let action : Array[Float] = Array::make(policy.action_dim, 0.0F)
for i in 0.. (Array[Float], Array[Float], Array[Float]) {
let action_seq : Array[Float] = Array::make(seq_len * policy.action_dim, 0.0F)
let mut hidden = hidden_init
let mut cell = cell_init
for t in 0..