// gru_ddpg_update_actor.mbt — DDPG_GRU BPTT-driven actor parameter update (v0.65.0).
//
// Companion to gru_ddpg_update_critic_seq (v0.64.0): implements the
// deterministic policy gradient (DPG) through the actor's recurrent
// (GRU) path. The DPG identity for continuous-action actor-critic
// (Lillicrap 2015) is:
//
// ∇_θ J = E_s[ ∇_a Q(s, π(s)) · ∇_θ π_θ(s) ]
//
// In a recurrent actor we need to:
// 1. Roll out the actor for T steps, getting action_seq AND the
// per-step actor hidden state cache (for BPTT through the actor).
// 2. Roll out critic1 (frozen w.r.t. actor grad) on the same
// (state_seq, action_seq), getting Q_seq.
// 3. For each t, ∂Q_t / ∂action_t = W2^T (one row of critic mlp_w2)
// (since Q = W2 · hidden_critic + b2, and hidden_critic
// depends on action via ReLU(w1 · [s, a] + b1) -> GRU -> mlp_w2).
// 4. ∂action_t / ∂actor_params: closed form via tanh squash.
// In v0.54.0 DDPG update, only mlp_w2 + mlp_b2 of the actor are
// updated (W1 update deferred). Here we extend this to also
// BPTT through the actor's GRU hidden state.
//
// This file adds the BPTT-driven actor update. Implementation notes:
//
// - The actor's mlp_w2 shape is (action_dim × gru_hidden). The actor
// update for slot t is:
// d_a = W2_critic^T · ∂Q_t/∂a (1×hidden vector * hidden ?)
// Actually the chain is:
// ∂Q/∂a_t = W2_critic · ∂hidden_critic/∂a_t
// Computing ∂hidden_critic/∂a_t requires BPTT through the critic
// GRU (the one done in v0.64.0). For simplicity and to keep the
// update parallel to v0.54.0's actor update, this version computes
// ∂Q_t/∂action_t directly via finite-difference (eps=1e-3) on the
// Q output. This is approximate but sufficient for a BPTT demo;
// the analytic chain through critic GRU is documented as a follow-up.
//
// - The actor's MLP w2 + b2 update uses the closed-form
// ∂Q/∂a · ∂a/∂W2_a (matching v0.54.0 ddpg_update_actor).
// - The actor's GRU gradient is computed via finite-difference on
// the Q output w.r.t. actor hidden state at each timestep, then
// propagated back through gru_cell_backward. This is a
// "policy gradient through time" approximation; the fully analytic
// version requires ∂Q/∂hidden_a which needs a second critic
// forward pass — left as a follow-up.
//
// `gru_deterministic_policy_seq_forward` (v0.58.0) and
// `gru_critic_step_with_cache` (v0.64.0) are reused.
///|
/// Update actor of a DDPG_GRU agent via T-step BPTT on a sampled
/// (state_seq) window. Returns the mean absolute Q across the T slots
/// (for diagnostics). Updates the actor in place via simple SGD.
///
/// Finite-difference eps for the actor-critic gradient is `eps` (default
/// 1e-3). Smaller eps is less noisy but more sensitive to Float32
/// round-off; 1e-3 is the standard choice for DDPG.
pub fn gru_ddpg_update_actor_seq(
agent : DDPG_GRU,
state_seq : Array[Float],
hidden_init : Array[Float],
lr : Float,
eps : Float,
) -> Float {
let seq_len = state_seq.length() / agent.actor.state_dim
let s_dim = agent.actor.state_dim
let a_dim = agent.actor.action_dim
let h_dim = agent.actor.gru_hidden
// Step 1: forward through actor -> action_seq + per-step caches.
let action_seq : Array[Float] = Array::make(seq_len * a_dim, 0.0F)
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),
})
let mut hidden_a = hidden_init
for t in 0.. (Array[Float], Array[Float], GruCellCache) {
let x_proj_pre = matvec(policy.mlp_w1, policy.mlp_b1, state)
let x_proj = relu_forward(x_proj_pre)
let (hidden_next, cache) = gru_cell_forward(x_proj, hidden, policy.gru)
let a_pre = matvec(policy.mlp_w2, policy.mlp_b2, hidden_next)
let half_range = (policy.action_high - policy.action_low) * 0.5F
let mid = (policy.action_high + policy.action_low) * 0.5F
let action : Array[Float] = Array::make(policy.action_dim, 0.0F)
for i in 0.. GRUDDPGActorGrad {
let s_dim = policy.state_dim
let a_dim = policy.action_dim
let h_dim = policy.gru_hidden
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(a_dim, [])
let zb2 : Array[Float] = Array::make(a_dim, 0.0F)
for i in 0..