// recurrent_sac.mbt — SAC with LSTM actor + LSTM critic (v0.38.2).
//
// Extension of v0.38.0 (vanilla SAC) and v0.38.1 (auto-tuned α) for
// partial observability. The POMDP corridor is the test bed: the
// agent only sees an indicator at t=0 (which side has the goal) and
// its current position; it must remember the indicator across the
// episode via the LSTM.
//
// Architecture:
//   - Actor  : LSTM cell → w_out_actor → softmax(logits) = π(a|h_t)
//   - Critic : LSTM cell → w_out_q → Q(h_t, a)  (one Q value per action)
//   - Twin Q + target Q with Polyak averaging (same as v0.38.0)
//   - Auto-tuned α from v0.38.1
//
// Forward path (per episode):
//   obs_seq = [obs_0, ..., obs_{T-1}]
//   h0, c0 = 0
//   for t in 0..T:
//     h_t = LSTM(obs_t, h_{t-1}, c_{t-1})
//     π_t = softmax(w_a · h_t + b_a)
//     Q1_t[i] = w_q1[i] · h_t + b_q1[i]   (for all i ∈ actions)
//     Q2_t[i] = w_q2[i] · h_t + b_q2[i]
//
// Soft target (per time step, for the critic update):
//   target_t = r_t + γ · (1 − done_t) · Σ_{a'} π_{t+1}(a'|s_{t+1}) · [Q̂_min(s_{t+1}, a') − α · log π_{t+1}(a'|s_{t+1})]
// where Q̂_min = min(Q̂1_target, Q̂2_target).
//
// Note: recurrent SAC uses BPTT through the whole episode for actor
// and critic updates — no replay buffer (per-agent-state). This is
// the standard on-policy recurrent RL formulation.

///|
/// LSTM Q-net: cell + linear head producing one Q value per action.
pub struct LstmQNet {
  cell : LstmCellParam
  w_out : Array[Array[Float]
]  // n_actions × d_h
  b_out : Array[Float]
}

///|
pub fn LstmQNet::new(
  d_x : Int,
  d_h : Int,
  n_actions : Int,
  seed : UInt64,
) -> LstmQNet {
  let cell = LstmCellParam::new(d_x, d_h, seed)
  let std = sqrtf(1.0F / Float::from_int(d_h))
  let rng = Xoshiro::from_state(seed + 11UL, seed + 12UL, seed + 13UL, seed + 14UL)
  let w_out = xavier_normal(n_actions, d_h, std, rng)
  let b_out : Array[Float] = Array::make(n_actions, 0.0F)
  { cell, w_out, b_out }
}

///|
/// Cache from one forward pass through an LstmQNet.
pub struct LstmQNetCache {
  cell_caches : Array[LstmCellCache]
  qs : Array[Array[Float]]
  hs : Array[Array[Float]]
}

///|
/// Forward pass through the Q-net over a sequence of observations.
/// Returns the Q-values per time step + a cache for BPTT.
pub fn lstm_qnet_forward(
  qnet : LstmQNet,
  obs_seq : Array[Array[Float]],
  h0 : Array[Float],
  c0 : Array[Float],
) -> (Array[Array[Float]], LstmQNetCache) {
  let n = obs_seq.length()
  let n_actions = qnet.w_out.length()
  let cell_caches : Array[LstmCellCache] = []
  let qs : Array[Array[Float]] = []
  let hs : Array[Array[Float]] = []
  let mut cur_h = h0
  let mut cur_c = c0
  for t in 0.. RecurrentSac {
  let n_actions = 2
  let policy = LstmPolicy::new(n_cells, d_h, seed)
  let d_x = 2 + n_cells
  let q1 = LstmQNet::new(d_x, d_h, n_actions, seed + 1UL)
  let q2 = LstmQNet::new(d_x, d_h, n_actions, seed + 2UL)
  let q1_target = LstmQNet::new(d_x, d_h, n_actions, seed + 3UL)
  let q2_target = LstmQNet::new(d_x, d_h, n_actions, seed + 4UL)
  // Sync target = online initially.
  copy_lstm_qnet(q1_target, q1)
  copy_lstm_qnet(q2_target, q2)
  { policy, q1, q2, q1_target, q2_target, log_alpha: log_alpha_init, target_entropy, alpha_lr }
}

///|
/// Current α = exp(log_alpha).
pub fn recurrent_sac_get_alpha(sac : RecurrentSac) -> Float {
  expf(sac.log_alpha)
}

///|
/// Episode trace for recurrent SAC. Stores obs_seq, next_obs_seq,
/// actions, rewards, dones for both actor and critic updates.
pub(all) struct RecurrentSacEpisode {
  obs_seq : Array[Array[Float]]
  next_obs_seq : Array[Array[Float]]
  actions : Array[Int]
  rewards : Array[Float]
  dones : Array[Bool]
  initial_pos : Int
  goal_side : Int
}

///|
/// Convenience constructor for synthetic episodes (tests).
pub fn RecurrentSacEpisode::new(
  obs_seq : Array[Array[Float]],
  next_obs_seq : Array[Array[Float]],
  actions : Array[Int],
  rewards : Array[Float],
  dones : Array[Bool],
  initial_pos : Int,
  goal_side : Int,
) -> RecurrentSacEpisode {
  { obs_seq, next_obs_seq, actions, rewards, dones, initial_pos, goal_side }
}

///|
/// Roll out one episode on the corridor POMDP, returning the trace.
pub fn recurrent_sac_rollout_episode(
  env_n_cells : Int,
  sac : RecurrentSac,
  max_steps : Int,
  rng : Xoshiro,
) -> RecurrentSacEpisode {
  let (u_raw, _) = box_muller(rng)
  let side = if u_raw > 0.0 { 1 } else { 0 }
  let env : CorridorEnv = CorridorEnv::{
    n_cells: env_n_cells,
    goal_side: side,
    step_penalty: -0.1F,
    goal_reward: 1.0F,
    max_steps,
  }
  let start = env_n_cells / 2
  let obs_seq : Array[Array[Float]] = []
  let next_obs_seq : Array[Array[Float]] = []
  let actions : Array[Int] = []
  let rewards : Array[Float] = []
  let dones : Array[Bool] = []
  let mut pos = start
  let mut done = false
  let mut t = 0
  // Push first obs.
  obs_seq.push(corridor_observation(env, pos))
  while !done && t < max_steps {
    // Sample action from current policy using running hidden state.
    let h_t : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
    let c_t : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
    // Re-run policy from start up to t to get current hidden state.
    // (Cheaper: maintain hidden state across the loop.)
    let _ = h_t
    let _ = c_t
    // Re-derive by re-running forward from t=0..t.
    let partial_obs : Array[Array[Float]] = []
    for k in 0.. Array[Float] {
  let n = ep.actions.length()
  let n_a = 2
  let targets : Array[Float] = Array::make(n, 0.0F)
  // Forward pass policy on next_obs_seq to get π_{t+1}.
  let h0 : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
  let c0 : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
  let (_, policy_cache) = lstm_policy_forward(sac.policy, ep.next_obs_seq, h0, c0, rng)
  // Forward pass target critics on next_obs_seq.
  let (_, q1_target_cache) = lstm_qnet_forward(sac.q1_target, ep.next_obs_seq, h0, c0)
  let (_, q2_target_cache) = lstm_qnet_forward(sac.q2_target, ep.next_obs_seq, h0, c0)
  let alpha = recurrent_sac_get_alpha(sac)
  for t in 0.. 0.0F { logf(pi_t[i]) } else { -20.0F }
        expected = expected + pi_t[i] * (q_min - alpha * lp)
      }
      targets[t] = ep.rewards[t] + gamma * expected
    }
  }
  targets
}

///|
/// Critic MSE update via BPTT. Applies SGD on q1 and q2 (online
/// critics only; target critics updated separately via Polyak).
/// Returns mean squared TD error.
pub fn recurrent_sac_critic_update(
  sac : RecurrentSac,
  ep : RecurrentSacEpisode,
  gamma : Float,
  lr : Float,
  rng : Xoshiro,
) -> Float {
  let targets = recurrent_sac_soft_target_seq(sac, ep, gamma, rng)
  let h0 : Array[Float] = Array::make(sac.q1.cell.d_h, 0.0F)
  let c0 : Array[Float] = Array::make(sac.q1.cell.d_h, 0.0F)
  let (q1_pred, q1_cache) = lstm_qnet_forward(sac.q1, ep.obs_seq, h0, c0)
  let (q2_pred, q2_cache) = lstm_qnet_forward(sac.q2, ep.obs_seq, h0, c0)
  let n = ep.actions.length()
  let mut total_loss = 0.0F
  for t in 0.. 0 {
    total_loss / Float::from_int(n)
  } else {
    0.0F
  }
}

///|
/// Apply accumulated gradients to an LstmQNet (output projection
/// + LSTM cell weights/biases) via SGD step. Mirrors the helper
/// used by `lstm_policy_gradient_update` for the actor.
fn apply_lstm_qnet_grad(
  qnet : LstmQNet,
  cache : LstmQNetCache,
  d_logits_per_t : Array[Array[Float]],
  lr : Float,
) -> Unit {
  let n = cache.hs.length()
  let n_actions = qnet.w_out.length()
  let d_h = qnet.cell.d_h
  // 1) Output projection gradients.
  for t in 0.. Unit {
  let n = ep.actions.length()
  let n_actions = sac.policy.w_out.length()
  let d_h = sac.policy.cell.d_h
  let h0 : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
  let c0 : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
  // Forward pass policy and q1 (critic) over obs_seq.
  let (_, policy_cache) = lstm_policy_forward(sac.policy, ep.obs_seq, h0, c0, rng)
  let (_, q1_cache) = lstm_qnet_forward(sac.q1, ep.obs_seq, h0, c0)
  let alpha = recurrent_sac_get_alpha(sac)
  // Per-step d_logits for actor (gradient of -Q + α·H wrt logits).
  // For each time step, for each action a:
  //   d_logits[i] = Σ_a π(a) · (indicator_{i=a} − π(i)) · (−Q[i] + α · (log π(i) + 1))
  // = −(Q(i) − Σ_a Q(a)π(a)) · (1 − π(i)) + Σ_{a≠i} ... actually simpler form:
  //   d_logits[i] = Σ_a π(a) · (1{i=a} − π(i)) · (−Q(a) + α·(log π(i) + 1))
  // Let v_i = Q(i) − α·(log π(i) + 1). Then d_logits[i] = Σ_a π(a) · (1{i=a} − π(i)) · (−v_a)
  //                                  = −Σ_a π(a) · 1{i=a} · v_a + Σ_a π(a)² · v_a
  //                                  = −π(i) · v_i + π(i) · Σ_a π(a) · v_a
  //                                  = π(i) · (Σ_a π(a) · v_a − v_i)
  let d_logits_per_t : Array[Array[Float]] = Array::make(n, [])
  for t in 0.. 0.0F { logf(pi[i]) } else { -20.0F }
      v_arr[i] = q[i] - alpha * (lp + 1.0F)
      v_mean = v_mean + pi[i] * v_arr[i]
    }
    let d_logits : Array[Float] = Array::make(n_actions, 0.0F)
    for i in 0.. Float {
  let h0 : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
  let c0 : Array[Float] = Array::make(sac.policy.cell.d_h, 0.0F)
  let (_, policy_cache) = lstm_policy_forward(sac.policy, ep.obs_seq, h0, c0, rng)
  let n_a = 2
  let mut sum_neg_ent = 0.0F
  let n = ep.actions.length()
  for t in 0.. 0.0F { logf(pi[i]) } else { -20.0F }
      neg_ent = neg_ent + pi[i] * lp
    }
    sum_neg_ent = sum_neg_ent + neg_ent
  }
  let mean_neg_ent = if n > 0 { sum_neg_ent / Float::from_int(n) } else { 0.0F }
  let delta = mean_neg_ent - sac.target_entropy
  sac.log_alpha = sac.log_alpha + sac.alpha_lr * delta
  delta
}

///|
/// Polyak soft target update for both critics.
pub fn recurrent_sac_soft_update(sac : RecurrentSac, tau : Float) -> Unit {
  copy_lstm_qnet_polyak(sac.q1_target, sac.q1, tau)
  copy_lstm_qnet_polyak(sac.q2_target, sac.q2, tau)
}

///|
/// Copy online Q-net into target Q-net (initial sync).
fn copy_lstm_qnet(target : LstmQNet, source : LstmQNet) -> Unit {
  // Copy w_out / b_out.
  for i in 0.. Unit {
  for i in 0.. Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let mut total_return = 0.0F
  for _ep in 0.. 0 {
    total_return / Float::from_int(n_episodes)
  } else {
    0.0F
  }
}