// 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
}
}