// lstm_ddpg_update.mbt — LSTM_DDPG BPTT-driven critic parameter update (v0.66.0).
//
// Parallel to gru_ddpg_update_critic_seq (v0.64.0) but for the LSTM
// variant. Closes the BPTT-driven parameter update half of v0.63.0
// LSTM_DDPG.
//
// Pipeline for one critic update (T-step BPTT):
//   1. T-step forward through the LSTM critic, storing per-step
//      (x_proj, hidden, cell, cache_lstm) for BPTT.
//   2. Compute per-step TD targets via `lstm_ddpg_compute_td_target_seq`.
//   3. T-step BPTT backward through critic1:
//        d_q_t = 2 * (q_t - q_target_t)
//        d_hidden_t = d_hidden_{t+1} + W2^T · d_q_t
//        d_cell_t = d_cell_{t+1} + (cell-state gradient from output gate)
//        d_x_proj_t = ReLU'(x_proj_t) ⊙ d_hidden_t
//        accumulate d_mlp_w1 / d_mlp_b1 / d_mlp_w2 / d_mlp_b2
//        lstm_cell_backward propagates through 4 gates + cell state,
//        accumulating d_w_f / d_w_i / d_w_c / d_w_o + biases.
//   4. SGD step on critic params + accumulated LSTM grads.
//
// Reuses `lstm_cell_backward` (lstm_cell.mbt v0.29.0) and
// `LstmCellGrad::zero` for the LSTM gradient buffer.

///|
/// Per-critic gradient buffer matching the shape of an
/// `LSTMQNetworkContinuous`.
pub struct LSTMDDPGCriticGrad {
  mlp_w1 : Array[Array[Float]]
  mlp_b1 : Array[Float]
  mlp_w2 : Array[Array[Float]]
  mut mlp_b2 : Float
  lstm_grad : LstmCellGrad
}

///|
/// Build a fresh `LSTMDDPGCriticGrad` matching the shape of the
/// supplied `LSTMQNetworkContinuous`. Caller is responsible for
/// zero-initialising once per training step.
pub fn LSTMDDPGCriticGrad::zero(
  qnet : LSTMQNetworkContinuous,
) -> LSTMDDPGCriticGrad {
  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.lstm_hidden
  // Allocate per-step cache storage.
  let x_proj_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  let hidden_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  let cell_seq : Array[Float] = Array::make(seq_len * h_dim, 0.0F)
  let lstm_cache_seq : Array[LstmCellCache] = Array::make(seq_len, {
    x: Array::make(h_dim, 0.0F), h_prev: Array::make(h_dim, 0.0F),
    c_prev: Array::make(h_dim, 0.0F), f: Array::make(h_dim, 0.0F),
    ig: Array::make(h_dim, 0.0F), c_tilde: Array::make(h_dim, 0.0F),
    o: Array::make(h_dim, 0.0F), c_t: Array::make(h_dim, 0.0F),
    tanh_c_t: 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
  let mut cell = cell_init
  for t in 0..