diff --git a/crates/ml/src/trainers/dqn/trainer/metrics.rs b/crates/ml/src/trainers/dqn/trainer/metrics.rs index 9b1c4a4f6..b13f32e39 100644 --- a/crates/ml/src/trainers/dqn/trainer/metrics.rs +++ b/crates/ml/src/trainers/dqn/trainer/metrics.rs @@ -111,12 +111,12 @@ impl DQNTrainer { let top5_coverage_pct = (top5_count as f64 / total_factored as f64) * 100.0; metrics.add_metric("top5_coverage_pct", top5_coverage_pct); } else { - // Fallback: 5 exposure levels + // Fallback: exposure-only (9 levels) let unique_actions = total_action_counts.iter() .filter(|&&count| count > 0).count(); - let action_diversity = (unique_actions as f64 / 5.0) * 100.0; + let action_diversity = (unique_actions as f64 / 9.0) * 100.0; metrics.add_metric("action_diversity", action_diversity); - metrics.add_metric("action_space_size", 5.0); + metrics.add_metric("action_space_size", 9.0); let active_threshold = (total_exposure as f64 * 0.005).max(1.0); let active_count = total_action_counts.iter() diff --git a/crates/ml/tests/smoke_test_real_data.rs b/crates/ml/tests/smoke_test_real_data.rs index 93a800d74..02561eb43 100644 --- a/crates/ml/tests/smoke_test_real_data.rs +++ b/crates/ml/tests/smoke_test_real_data.rs @@ -883,9 +883,11 @@ async fn smoke_e2e_dqn_training_loop() { .copied() .unwrap_or(0.0); info!(action_space, action_diversity, total_actions, "Action metrics"); + // Branching DQN: action_space=81 (9×3×3) when factored counts populated, + // or 9 (exposure-only fallback) for very short runs. assert!( - action_space > 5.0 || total_actions == 0.0, - "Branching DQN should report 81-action space, got {action_space}" + action_space >= 9.0 || total_actions == 0.0, + "Branching DQN should report >=9 action space, got {action_space}" ); // 6. Gradient norms should be non-zero (clipping is functional)