// td3.mbt — Twin Delayed DDPG (Fujimoto et al. 2018).
//
// TD3 improves on DDPG (v0.54.0) with three modifications:
//   1. Twin Q-networks (already in our DDPG — same dual critic structure).
//   2. Target policy smoothing: when computing the Q-target, add
//      clipped Gaussian noise to the target action to reduce
//      variance from Q-function approximation errors.
//        a' = clip(μ̂(s') + clip(ε, -clip_range, clip_range), a_low, a_high)
//        where ε ~ N(0, target_noise_std)
//   3. Delayed policy update: update the actor and target networks
//      every `policy_delay` critic updates (instead of every step).
//
// The TD3 agent reuses the same networks as DDPG (deterministic actor,
// twin continuous-action critics, twin targets) — only the *update
// logic* differs.
//
// Reference:
//   Fujimoto, van Hoof, Meger, "Addressing Function Approximation
//   Error in Actor-Critic Methods" (ICML 2018).

///|
/// TD3 agent. Same networks as DDPG; the only new fields are
/// `target_noise_std` and `policy_delay`.
pub struct TD3 {
  actor : DeterministicPolicy
  critic1 : QNetworkContinuous
  critic2 : QNetworkContinuous
  actor_target : DeterministicPolicy
  critic1_target : QNetworkContinuous
  critic2_target : QNetworkContinuous
  gamma : Float
  tau : Float
  exploration_noise : Float
  // TD3-specific.
  target_noise_std : Float
  target_noise_clip : Float
  policy_delay : Int
}

///|
/// Construct a TD3 agent. Mirrors DDPG::new with two extra parameters:
/// `target_noise_std` (Gaussian noise stddev for target smoothing) and
/// `policy_delay` (number of critic updates between actor + target updates).
pub fn TD3::new(
  state_dim : Int,
  action_dim : Int,
  hidden : Int,
  action_low : Float,
  action_high : Float,
  gamma : Float,
  tau : Float,
  exploration_noise : Float,
  target_noise_std : Float,
  target_noise_clip : Float,
  policy_delay : Int,
  seed : UInt64,
) -> TD3 {
  let actor = DeterministicPolicy::new(
    state_dim, action_dim, hidden, action_low, action_high, seed,
  )
  let critic1 = QNetworkContinuous::new(state_dim, action_dim, hidden, seed + 1UL)
  let critic2 = QNetworkContinuous::new(state_dim, action_dim, hidden, seed + 2UL)
  let actor_target = DeterministicPolicy::new(
    state_dim, action_dim, hidden, action_low, action_high, seed + 3UL,
  )
  let critic1_target = QNetworkContinuous::new(
    state_dim, action_dim, hidden, seed + 4UL,
  )
  let critic2_target = QNetworkContinuous::new(
    state_dim, action_dim, hidden, seed + 5UL,
  )
  policy_copy(actor_target, actor)
  qnet_continuous_copy(critic1_target, critic1)
  qnet_continuous_copy(critic2_target, critic2)
  {
    actor,
    critic1,
    critic2,
    actor_target,
    critic1_target,
    critic2_target,
    gamma,
    tau,
    exploration_noise,
    target_noise_std,
    target_noise_clip,
    policy_delay,
  }
}

///|
/// Deterministic policy action (no noise). For evaluation.
pub fn td3_select_action_eval(
  td3 : TD3,
  state : Array[Float],
) -> Array[Float] {
  let hidden : Array[Float] = Array::make(td3.actor.hidden, 0.0F)
  deterministic_policy_forward(td3.actor, state, hidden)
}

///|
/// Action with Gaussian exploration noise (clipped to action range).
pub fn td3_select_action(
  td3 : TD3,
  state : Array[Float],
  rng : Xoshiro,
) -> Array[Float] {
  let action = td3_select_action_eval(td3, state)
  let noisy : Array[Float] = Array::make(td3.actor.action_dim, 0.0F)
  for a in 0.. td3.actor.action_high {
      v = td3.actor.action_high
    }
    ignore(noisy.set(a, v))
  }
  noisy
}

///|
/// Compute the target action for TD3: target actor + clipped Gaussian
/// noise, clipped to action range. This is "target policy smoothing".
fn td3_smoothed_target_action(
  td3 : TD3,
  next_state : Array[Float],
  rng : Xoshiro,
  hidden_buf : Array[Float],
) -> Array[Float] {
  let target_act = deterministic_policy_forward(
    td3.actor_target, next_state, hidden_buf,
  )
  let smoothed : Array[Float] = Array::make(td3.actor.action_dim, 0.0F)
  let clip = td3.target_noise_clip
  let action_low = td3.actor.action_low
  let action_high = td3.actor.action_high
  for a in 0.. clip {
      clip
    } else if noise < -clip {
      -clip
    } else {
      noise
    }
    let mut v = target_act[a] + clipped_noise
    if v < action_low {
      v = action_low
    } else if v > action_high {
      v = action_high
    }
    ignore(smoothed.set(a, v))
  }
  smoothed
}

///|
/// Run one TD3 update step on a mini-batch. Returns the mean absolute
/// TD error across both critics. Updates critics every call; updates
/// actor + targets only when `step_count % policy_delay == 0`.
///
/// Implementation:
///   1. Compute smoothed target action a'.
///   2. Compute twin Q-targets on (next_state, a').
///   3. Twin-min: q_target = r + γ(1-d)·min(Q̂1, Q̂2).
///   4. Update both online critics via analytic gradient.
///   5. If step_count % policy_delay == 0, update actor + soft-sync targets.
pub fn td3_update_step(
  td3 : TD3,
  batch : (Array[Float], Array[Float], Array[Float], Array[Float], Array[Float]),
  lr_critic : Float,
  lr_actor : Float,
  step_count : Int,
  rng : Xoshiro,
) -> Float {
  let (states_b, actions_b, rewards_b, next_states_b, dones_b) = batch
  let batch_size = rewards_b.length()
  let s_dim = td3.actor.state_dim
  let a_dim = td3.actor.action_dim
  let hidden_dim = td3.actor.hidden
  let mut total_abs_td = 0.0F
  let hidden_q1 : Array[Float] = Array::make(hidden_dim, 0.0F)
  let hidden_q2 : Array[Float] = Array::make(hidden_dim, 0.0F)
  let hidden_q1_t : Array[Float] = Array::make(hidden_dim, 0.0F)
  let hidden_q2_t : Array[Float] = Array::make(hidden_dim, 0.0F)
  let hidden_actor_t : Array[Float] = Array::make(hidden_dim, 0.0F)
  for k in 0.. 0 {
    td3_actor_update(td3, batch, lr_actor)
    td3_soft_update(td3)
  }
  total_abs_td / Float::from_int(batch_size)
}

///|
/// Apply critic gradient (same as DDPG's critic_apply_grad).
fn td3_apply_critic_grad(
  net : QNetworkContinuous,
  state : Array[Float],
  action : Array[Float],
  hidden : Array[Float],
  delta : Float,
) -> Unit {
  for h in 0.. 0.0F { 1.0F } else { 0.0F }
    let grad_h_eff = grad_h * gate
    net.b1[h] = net.b1[h] - grad_h_eff
    for j in 0.. Unit {
  let (states_b, _actions_b, _rewards_b, _next_states_b, _dones_b) = batch
  let batch_size = states_b.length() / td3.actor.state_dim
  let s_dim = td3.actor.state_dim
  let a_dim = td3.actor.action_dim
  let hidden_dim = td3.actor.hidden
  let hidden_act : Array[Float] = Array::make(hidden_dim, 0.0F)
  let hidden_q : Array[Float] = Array::make(hidden_dim, 0.0F)
  let range = (td3.actor.action_high - td3.actor.action_low) * 0.5F
  for k in 0.. 0.0F { 1.0F } else { 0.0F }
        g_a_j = g_a_j + td3.critic1.w1[h][s_dim + a] *
          td3.critic1.w2[h] * gate
      }
      let mut s_pre = td3.actor.b2[a]
      for h in 0.. Unit {
  let tau = td3.tau
  let one_minus_tau = 1.0F - tau
  for h in 0.. (Array[Float], Float, Bool),
  td3 : TD3,
  n_episodes : Int,
  buffer_capacity : Int,
  warmup_episodes : Int,
  batch_size : Int,
  max_steps : Int,
  seed : UInt64,
) -> Float {
  let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
  let buffer = ContinuousReplayBuffer::new(buffer_capacity, env_state_dim, env_action_dim)
  let mut total_return = 0.0F
  let mut return_count = 0
  let mut global_step = 0
  for ep in 0..= warmup_episodes && buffer.len() >= batch_size {
      let batch = buffer.sample(batch_size, rng)
      let _ = td3_update_step(td3, batch, 0.001F, 0.0001F, global_step, rng)
    }
    let _ = env_action_dim
    let _ = env_state_dim
  }
  if return_count > 0 {
    total_return / Float::from_int(return_count)
  } else {
    0.0F
  }
}