// gru_ddpg_update.mbt — DDPG_GRU BPTT-driven critic parameter update (v0.64.0).
//
// Closes the open loop on the v0.60.0 DDPG_GRU agent: the forward path
// (select_action / compute_td_target_seq / soft_update) was already
// shipped; this version adds the per-step critic parameter update via
// T-step back-propagation through time (BPTT).
//
// Pipeline for one critic update:
// 1. T-step forward through the critic (source) on (state_seq,
// action_seq), storing per-step (x_proj, hidden, cache_gru) for BPTT.
// 2. Compute the per-step TD target via
// `ddpg_gru_compute_td_target_seq` (in gru_ddpg.mbt) which reuses
// critic1_target / critic2_target for the next-step Q.
// 3. T-step BPTT backward through the critic:
// d_q_seq[t] = 2 * (q_t - q_target_t)
// d_hidden_t = d_hidden_t+1 + W2^T · d_q_seq[t]
// d_x_proj_t = ReLU'(x_proj_t) ⊙ d_hidden_t
// SGD step on mlp params at slot t (w1 / b1 / w2 / b2)
// d_h propagates back through gru_cell_backward (accumulates
// d_W_z / d_W_r / d_W_n / d_b_z / d_b_r / d_b_n into `grad`).
// 4. SGD step on the accumulated GRU gradients.
//
// Twin-critic symmetric: the same loop runs on critic2, sharing the
// TD targets.
//
// `gru_cell_backward` (gru_cell.mbt v0.30.0), `gru_cell_forward`,
// and `GruCellGrad::zero` are reused from prior versions.
///|
/// Per-critic gradient buffer matching the shape of a `GRUQNetworkContinuous`.
/// Reset to zero before each BPTT step.
pub struct GRUDDPGCriticGrad {
mlp_w1 : Array[Array[Float]]
mlp_b1 : Array[Float]
mlp_w2 : Array[Array[Float]]
mut mlp_b2 : Float
gru_grad : GruCellGrad
}
///|
/// Build a fresh `GRUDDPGCriticGrad` matching the shape of the
/// supplied `GRUQNetworkContinuous`. Caller is responsible for
/// zero-initialising once per training step.
pub fn GRUDDPGCriticGrad::zero(qnet : GRUQNetworkContinuous) -> GRUDDPGCriticGrad {
let s_dim = qnet.state_dim
let a_dim = qnet.action_dim
let h_dim = qnet.hidden
let in_dim = s_dim + a_dim
let zw1 : Array[Array[Float]] = Array::make(h_dim, [])
let zb1 : Array[Float] = Array::make(h_dim, 0.0F)
let zw2 : Array[Array[Float]] = Array::make(1, [])
zw2[0] = Array::make(h_dim, 0.0F)
for i in 0.. Float {
let seq_len = action_seq.length()
let s_dim = agent.actor.state_dim
let a_dim = agent.actor.action_dim
let h_dim = agent.actor.gru_hidden
// Allocate per-step cache storage.
let x_proj_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
let gru_cache_seq : Array[GruCellCache] = Array::make(seq_len, {
x: Array::make(h_dim, 0.0F), h_prev: Array::make(h_dim, 0.0F),
z: Array::make(h_dim, 0.0F), r: Array::make(h_dim, 0.0F),
s: Array::make(h_dim, 0.0F), n: Array::make(h_dim, 0.0F),
h_t: Array::make(h_dim, 0.0F),
})
// Step 1: forward through critic1 (source).
let q_seq : Array[Float] = Array::make(seq_len, 0.0F)
let mut hidden = hidden_init
for t in 0..= 0; t = t - 1 {
let td_err = q_seq[t] - td_seq[t]
total_abs_td = total_abs_td + (if td_err < 0.0F { -td_err } else { td_err })
let d_q = 2.0F * td_err // gradient of (q - q_target)^2 / 2
// d_hidden = d_h + W2^T · d_q
let d_hidden_from_q : Array[Float] = Array::make(h_dim, 0.0F)
for k in 0.. 0.0F { 1.0F } else { 0.0F }
d_x_proj[k] = d_hidden_from_q[k] * gate
}
// Accumulate MLP gradients (slot-t contribution).
let st_off = t * s_dim
let at_off = t * a_dim
for h_idx in 0.. (Float, Array[Float], GruCellCache) {
let sa = vec_concat(state, action)
let x_proj_pre = matvec(qnet.mlp_w1, qnet.mlp_b1, sa)
let x_proj = relu_forward(x_proj_pre)
let (hidden_next, cache) = gru_cell_forward(x_proj, hidden, qnet.gru)
let q_pre = matvec(qnet.mlp_w2, [qnet.mlp_b2], hidden_next)
(q_pre[0], x_proj, cache)
}
///|
/// SGD step on a critic's parameters using accumulated BPTT gradients.
fn critic_apply_sgd(
qnet : GRUQNetworkContinuous,
grad : GRUDDPGCriticGrad,
lr : Float,
) -> Unit {
let h_dim = qnet.hidden
for h_idx in 0.. Unit {
for i in 0..