diff --git a/crates/ml/src/trainers/dqn/trainer/metrics.rs b/crates/ml/src/trainers/dqn/trainer/metrics.rs index e33d1ed3f..9872eb37c 100644 --- a/crates/ml/src/trainers/dqn/trainer/metrics.rs +++ b/crates/ml/src/trainers/dqn/trainer/metrics.rs @@ -275,32 +275,24 @@ impl DQNTrainer { crate::cuda_pipeline::dtoh_bf16_to_f32(stream, q_out, &mut host_q).ok()?; host_q.truncate(sample_size * total_actions); - // Compute gap (best - second_best) and per-action averages on host + // Compute gap (best - second_best) WITHIN DIRECTION BRANCH ONLY + // and per-action averages on host. + // Q-values layout: [dir_0..dir_2, mag_0..mag_2, ord_0..ord_2, urg_0..urg_2] + // The gap measures direction conviction (best_dir - second_best_dir). let cols = total_actions; let rows = sample_size; + let (b0, b1, _, _) = { + let ag = self.agent.read().await; + ag.branch_sizes() + }; let mut gaps = Vec::with_capacity(rows); let mut col_sums = [0.0_f64; 9]; for r in 0..rows { let row_offset = r * cols; + // Direction gap: best vs second_best within first b0 columns only let mut best_val = f32::NEG_INFINITY; let mut second_best = f32::NEG_INFINITY; - for c in 0..cols.min(9) { - let v = host_q.get(row_offset + c).copied().unwrap_or(0.0); - if let Some(s) = col_sums.get_mut(c) { - *s += v as f64; - } - if v > best_val { - second_best = best_val; - best_val = v; - } else if v > second_best { - second_best = v; - } else { - // v <= second_best: no update needed - } - } - // Also check remaining columns (9..total_actions) for best/second_best - // (order_type + urgency branches, if total_actions > 9) - for c in 9..cols { + for c in 0..b0.min(cols) { let v = host_q.get(row_offset + c).copied().unwrap_or(0.0); if v > best_val { second_best = best_val; @@ -310,6 +302,19 @@ impl DQNTrainer { } } gaps.push(best_val - second_best); + // Per-action sums: map branch outputs to 9 dir×mag combos + // dir_d × mag_m → index d*b1 + m (for display only) + for d in 0..b0.min(3) { + let q_dir = host_q.get(row_offset + d).copied().unwrap_or(0.0) as f64; + for m in 0..b1.min(3) { + let q_mag = host_q.get(row_offset + b0 + m).copied().unwrap_or(0.0) as f64; + let combo_idx = d * b1.min(3) + m; + if let Some(s) = col_sums.get_mut(combo_idx) { + // Combined Q ≈ Q_dir + Q_mag (additive advantage) + *s += q_dir + q_mag; + } + } + } } let rows_f64 = rows as f64;