// 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..= 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
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
}
let st_off = t * s_dim
let at_off = t * a_dim
for h_idx in 0.. (Float, Array[Float], LstmCellCache) {
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) = lstm_cell_forward(x_proj, hidden, cell, qnet.lstm)
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 lstm_critic_apply_sgd(
qnet : LSTMQNetworkContinuous,
grad : LSTMDDPGCriticGrad,
lr : Float,
) -> Unit {
let h_dim = qnet.hidden
for h_idx in 0..