// dueling_dqn.mbt — Dueling DQN (v0.37.1).
//
// Dueling DQN (Wang et al. 2016) splits Q(s, a) into a value
// stream V(s) and an advantage stream A(s, a):
//
// Q(s, a; θ) = V(s; θ_v) + A(s, a; θ_a) - (1/|A|) · Σ_{a'} A(s, a'; θ_a)
//
// The mean-subtraction makes the advantage stream centred: at any
// state the average advantage across actions is 0, so Q reflects
// the relative value of each action. This helps when many actions
// have similar values at a state (no need to learn the magnitude
// of V separately for each action).
//
// Reuses:
// - `LinearQNet` from dqn.mbt (we use one net for V, one for A)
// - `ReplayBuffer`, `qnet_copy` from dqn.mbt
///|
/// Dueling Q-network: two LinearQNets (value + advantage).
pub struct DuelingQNet {
n_states : Int
n_actions : Int
value : LinearQNet
advantage : LinearQNet
}
///|
pub fn DuelingQNet::new(
n_states : Int,
n_actions : Int,
seed : UInt64,
) -> DuelingQNet {
let value = LinearQNet::new(n_states, n_actions, seed)
let advantage = LinearQNet::new(n_states, n_actions, seed + 17UL)
{ n_states, n_actions, value, advantage }
}
///|
/// Forward: Q(s, a) = V(s) + A(s, a) - mean_a A(s, a).
pub fn dueling_q_forward(net : DuelingQNet, x : Array[Float]) -> Array[Float] {
let n_a = net.n_actions
let v = q_forward(net.value, x)
let a = q_forward(net.advantage, x)
let v_scalar = v[0] // value stream outputs length n_actions; we use index 0 as V(s).
// Compute mean advantage.
let mut mean_a = 0.0F
for i in 0.. Int {
let mut best = 0
let mut best_v = q[0]
for i in 1.. best_v {
best_v = q[i]
best = i
}
}
best
}
///|
/// Mean of an array (helper).
fn array_mean(a : Array[Float]) -> Float {
if a.length() == 0 {
return 0.0F
}
let mut s = 0.0F
for i in 0.. Float {
let n = states.length()
let n_a = online.n_actions
let n_s = online.n_states
let mut total_abs_delta = 0.0F
for k in 0.. Float {
let mut m = a[0]
for i in 1.. m {
m = a[i]
}
}
m
}
///|
/// Copy weights from `src` into `dst`.
pub fn dueling_qnet_copy(dst : DuelingQNet, src : DuelingQNet) -> Unit {
qnet_copy(dst.value, src.value)
qnet_copy(dst.advantage, src.advantage)
}
///|
/// ε-greedy action using dueling Q-net.
pub fn dueling_eps_greedy_action(
q_net : DuelingQNet,
state : Int,
epsilon : Float,
rng : Xoshiro,
) -> Int {
if epsilon <= 0.0F {
let x = rl_one_hot(state, q_net.n_states)
let q = dueling_q_forward(q_net, x)
return dueling_q_argmax(q)
}
let (u, _) = box_muller(rng)
let u_pos = if u < 0.0 { -u } else { u }
if Float::from_double(u_pos) < epsilon {
let (u2, _) = box_muller(rng)
let n_a = q_net.n_actions
let idx_raw = if u2 < 0.0 { -u2 } else { u2 }
let idx = Float::from_double(idx_raw * Double::from_int(n_a)).to_int()
if idx < 0 {
0
} else if idx >= n_a {
n_a - 1
} else {
idx
}
} else {
let x = rl_one_hot(state, q_net.n_states)
let q = dueling_q_forward(q_net, x)
dueling_q_argmax(q)
}
}
///|
/// Train Dueling DQN. Returns mean episode return.
pub fn train_dueling_dqn(
env : GridWorld,
q_net : DuelingQNet,
target : DuelingQNet,
n_episodes : Int,
gamma : Float,
lr : Float,
epsilon_start : Float,
epsilon_end : Float,
max_steps : Int,
buffer_capacity : Int,
warmup_episodes : Int,
sync_every : Int,
batch_size : Int,
seed : UInt64,
) -> Float {
let rng = Xoshiro::from_state(seed, seed + 7UL, seed + 13UL, seed + 17UL)
let buffer = ReplayBuffer::new(buffer_capacity)
let mut total_return = 0.0F
let mut return_count = 0
for ep in 0..= warmup_episodes && buffer.len() >= batch_size {
let (ss, aa, rr, nss, dd) = buffer.sample(batch_size, rng)
let _ = dueling_dqn_update_step(q_net, target, ss, aa, rr, nss, dd, gamma, lr)
}
if (ep + 1) % sync_every == 0 {
dueling_qnet_copy(target, q_net)
}
}
if return_count > 0 {
total_return / Float::from_int(return_count)
} else {
0.0F
}
}