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