// gru_qnetwork_continuous.mbt — Recurrent continuous-action Q-network
// for partially-observable environments (v0.59.0).
//
// Architecture: MLP w1 (concat(state, action) -> hidden) -> ReLU ->
// GRU cell (over time) -> MLP w2 (hidden -> scalar Q).
//
// [state_t ; action_t] ∈ R^{state_dim + action_dim}
// x_proj_t = ReLU(w1 · [state_t ; action_t] + b1) ∈ R^{hidden}
// hidden_t = GRU_cell(x_proj_t, hidden_{t-1}) ∈ R^{hidden}
// q_t = w2 · hidden_t + b2 ∈ R
//
// This is the recurrent counterpart of `QNetworkContinuous` (v0.54.0)
// for the DDPG_GRU agent. The GRU hidden state mixes past state-action
// pairs so the critic can break the Markov assumption.
///|
/// Recurrent continuous-action Q-network. MLP hidden dim equals GRU
/// hidden dim so the ReLU-projected [state; action] and the GRU cell
/// input/output match. `mlp_b2` is a scalar (the critic outputs a
/// single Q-value per timestep).
pub struct GRUQNetworkContinuous {
state_dim : Int
action_dim : Int
hidden : Int
mlp_w1 : Array[Array[Float]]
mlp_b1 : Array[Float]
gru : GruCellParam
mlp_w2 : Array[Array[Float]]
mut mlp_b2 : Float
}
///|
/// Build a fresh GRUQNetworkContinuous. MLP weights use Xavier-normal
/// init scaled by sqrtf(2 / fan_in) (He-style for ReLU). GRU cell
/// uses its own init. Zero biases.
pub fn GRUQNetworkContinuous::new(
state_dim : Int,
action_dim : Int,
hidden : Int,
seed : UInt64,
) -> GRUQNetworkContinuous {
let in_dim = state_dim + action_dim
let rng1 = Xoshiro::from_state(seed, seed + 1UL, seed + 2UL, seed + 3UL)
let std1 = sqrtf(2.0F / Float::from_int(in_dim))
let mlp_w1 = xavier_normal(hidden, in_dim, std1, rng1)
let mlp_b1 : Array[Float] = Array::make(hidden, 0.0F)
let gru = GruCellParam::new(hidden, hidden, seed + 4UL)
let rng2 = Xoshiro::from_state(seed + 5UL, seed + 6UL, seed + 7UL, seed + 8UL)
let std2 = sqrtf(2.0F / Float::from_int(hidden))
let mlp_w2 = xavier_normal(1, hidden, std2, rng2)
{ state_dim, action_dim, hidden, mlp_w1, mlp_b1, gru, mlp_w2, mlp_b2: 0.0F }
}
///|
/// Single-step forward. Returns the scalar Q-value for the given
/// (state, action) pair mixed with the previous hidden state, plus
/// the new hidden state for the next step.
pub fn gru_qnetwork_continuous_step(
qnet : GRUQNetworkContinuous,
state : Array[Float],
action : Array[Float],
hidden : Array[Float],
) -> (Float, Array[Float]) {
// sa = concat(state, action)
let sa = vec_concat(state, action)
// x_proj = ReLU(w1 · sa + b1)
let x_proj_pre = matvec(qnet.mlp_w1, qnet.mlp_b1, sa)
let x_proj = relu_forward(x_proj_pre)
// hidden_next = GRU_cell(x_proj, hidden)
let (hidden_next, _cache) = gru_cell_forward(x_proj, hidden, qnet.gru)
// q = w2 · hidden_next + b2
let q_pre = matvec(qnet.mlp_w2, [qnet.mlp_b2], hidden_next)
(q_pre[0], hidden_next)
}
///|
/// Sequence forward. `state_seq` is flat row-major `[seq_len × state_dim]`,
/// `action_seq` is flat row-major `[seq_len × action_dim]`. Returns
/// `(q_seq, final_hidden)` where `q_seq` is flat `[seq_len]`.
pub fn gru_qnetwork_continuous_seq_forward(
qnet : GRUQNetworkContinuous,
state_seq : Array[Float],
action_seq : Array[Float],
seq_len : Int,
hidden_init : Array[Float],
) -> (Array[Float], Array[Float]) {
let q_seq : Array[Float] = Array::make(seq_len, 0.0F)
let mut hidden = hidden_init
for t in 0..