Files
foxhunt/ml/tests/target_update_tests.rs
jgrusewski ac0a83e4f7 refactor(ml): reorganize tests — move integration tests to ml/tests/
Move test files from ml/src/*/tests/ to ml/tests/. Convert
crate-internal imports to public API imports. Rename files to
follow naming conventions (no wave/priority prefixes).

Files moved:
- ml/src/dqn/tests/ -> ml/tests/ (4 files)
- ml/src/trainers/dqn/tests/ -> ml/tests/ (5 files)

Renames:
- target_update_comprehensive_tests.rs -> target_update_tests.rs
- p0_integration_tests.rs -> dqn_trainer_integration_tests.rs
- p1_integration_tests.rs -> dqn_trainer_p1_tests.rs
- ensemble_uncertainty_hyperopt_tests.rs -> ensemble_hyperopt_tests.rs

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-20 14:07:22 +01:00

385 lines
12 KiB
Rust

//! Comprehensive TDD tests for Polyak soft updates (WAVE 26 P1.12)
//!
//! Verifies:
//! 1. Tau is configurable and defaults to 0.001
//! 2. Soft update formula: theta_target = tau * theta_online + (1 - tau) * theta_target
//! 3. Target network divergence computation
//! 4. Convergence behavior over multiple updates
//! 5. Boundary conditions (tau = 0, tau = 1)
use candle_core::{DType, Device, Tensor};
use candle_nn::VarMap;
use ml::dqn::target_update::{polyak_update, hard_update, convergence_half_life};
use ml::trainers::dqn::DQNHyperparameters;
/// Helper: Create VarMap with uniform values
fn create_varmap(value: f32) -> VarMap {
let varmap = VarMap::new();
let device = Device::Cpu;
// Create test tensors
let weight = (Tensor::ones(&[10, 10], DType::F32, &device).unwrap() * (value as f64)).unwrap();
let bias = (Tensor::ones(&[10], DType::F32, &device).unwrap() * (value as f64)).unwrap();
let mut data = varmap.data().lock().unwrap();
data.insert("layer1.weight".to_string(), candle_core::Var::from_tensor(&weight).unwrap());
data.insert("layer1.bias".to_string(), candle_core::Var::from_tensor(&bias).unwrap());
drop(data);
varmap
}
/// Helper: Extract mean value from VarMap
fn get_mean_value(varmap: &VarMap) -> f32 {
let data = varmap.data().lock().unwrap();
let mut sum = 0.0;
let mut count = 0;
for (_, tensor) in data.iter() {
let t: &Tensor = tensor.as_ref();
sum += t.mean_all().unwrap().to_scalar::<f32>().unwrap();
count += 1;
}
sum / count as f32
}
/// Helper: Compute L2 divergence between two VarMaps
fn compute_network_divergence(online: &VarMap, target: &VarMap) -> f64 {
let online_data = online.data().lock().unwrap();
let target_data = target.data().lock().unwrap();
let mut total_divergence = 0.0;
let mut param_count = 0;
for (name, online_tensor) in online_data.iter() {
if let Some(target_tensor) = target_data.get(name) {
let online_t: &Tensor = online_tensor.as_ref();
let target_t: &Tensor = target_tensor.as_ref();
// L2 norm: sqrt(sum((online - target)^2))
let diff = (online_t - target_t).unwrap();
let squared = (&diff * &diff).unwrap();
let sum_squared = squared.sum_all().unwrap().to_scalar::<f32>().unwrap() as f64;
total_divergence += sum_squared.sqrt();
param_count += 1;
}
}
// Average divergence across all parameters
if param_count > 0 {
total_divergence / param_count as f64
} else {
0.0
}
}
#[test]
fn test_tau_default_value() {
// GIVEN: DQN config should default to tau=0.001
let config = DQNHyperparameters::default();
// THEN: Tau should be 0.001 (Rainbow DQN standard)
assert_eq!(config.tau, 0.001, "Default tau should be 0.001");
println!("Default tau = {} (Rainbow DQN standard)", config.tau);
}
#[test]
fn test_soft_update_formula_correctness() {
// GIVEN: Online network at 1.0, target at 0.0
let online_vars = create_varmap(1.0);
let target_vars = create_varmap(0.0);
let tau = 0.3; // Use larger tau for easier verification
// WHEN: Apply single Polyak update
polyak_update(&online_vars, &target_vars, tau).unwrap();
// THEN: theta_target = tau * theta_online + (1 - tau) * theta_target
// = 0.3 * 1.0 + 0.7 * 0.0 = 0.3
let result = get_mean_value(&target_vars);
assert!(
(result - 0.3).abs() < 1e-6,
"Expected 0.3, got {}. Formula: tau*1.0 + (1-tau)*0.0",
result
);
println!("Soft update formula correct: {:.6} (expected 0.3)", result);
}
#[test]
fn test_network_divergence_computation() {
// GIVEN: Two networks with known values
let online = create_varmap(1.0);
let target = create_varmap(0.5);
// WHEN: Compute divergence
let divergence = compute_network_divergence(&online, &target);
// THEN: Should be non-zero and finite
assert!(divergence > 0.0, "Divergence should be positive");
assert!(divergence.is_finite(), "Divergence should be finite");
// compute_network_divergence returns average L2 norm across parameters
// Layer 1 (10x10 weights): L2 = sqrt(0.5^2 * 100) = sqrt(25) = 5.0
// Layer 2 (10 bias): L2 = sqrt(0.5^2 * 10) = sqrt(2.5) ~= 1.58
// Average: (5.0 + 1.58) / 2 ~= 3.29
assert!(
divergence > 2.0 && divergence < 5.0,
"Expected divergence ~=3.29, got {}",
divergence
);
println!("Network divergence: {:.4} (L2 norm)", divergence);
}
#[test]
fn test_divergence_decreases_with_updates() {
// GIVEN: Networks starting far apart
let online = create_varmap(1.0);
let target = create_varmap(0.0);
let initial_divergence = compute_network_divergence(&online, &target);
// WHEN: Apply multiple soft updates
for _ in 0..10 {
polyak_update(&online, &target, 0.1).unwrap();
}
let final_divergence = compute_network_divergence(&online, &target);
// THEN: Divergence should decrease significantly
assert!(
final_divergence < initial_divergence * 0.5,
"Expected divergence to decrease by >50%, initial={:.4}, final={:.4}",
initial_divergence,
final_divergence
);
println!(
"Divergence decreased: {:.4} -> {:.4} ({:.1}% reduction)",
initial_divergence,
final_divergence,
100.0 * (1.0 - final_divergence / initial_divergence)
);
}
#[test]
fn test_tau_boundary_condition_zero() {
// GIVEN: Online at 1.0, target at 0.0
let online = create_varmap(1.0);
let target = create_varmap(0.0);
// WHEN: tau = 0 (no update)
polyak_update(&online, &target, 0.0).unwrap();
// THEN: Target should remain 0.0
let result = get_mean_value(&target);
assert!(
result.abs() < 1e-6,
"tau=0 should not update target, got {}",
result
);
println!("tau=0 boundary condition: target unchanged ({:.6})", result);
}
#[test]
fn test_tau_boundary_condition_one() {
// GIVEN: Online at 1.0, target at 0.0
let online = create_varmap(1.0);
let target = create_varmap(0.0);
// WHEN: tau = 1.0 (full copy, equivalent to hard update)
polyak_update(&online, &target, 1.0).unwrap();
// THEN: Target should equal online (1.0)
let result = get_mean_value(&target);
assert!(
(result - 1.0).abs() < 1e-6,
"tau=1.0 should copy online to target, got {}",
result
);
println!("tau=1.0 boundary condition: target = online ({:.6})", result);
}
#[test]
fn test_rainbow_tau_convergence_rate() {
// GIVEN: Rainbow's tau = 0.001
let tau = 0.001;
let half_life = convergence_half_life(tau);
// THEN: Should be approximately 693 steps
assert!(
(half_life - 693.0).abs() < 5.0,
"Rainbow tau=0.001 should give ~693 step half-life, got {:.0}",
half_life
);
// Verify empirically
let online = create_varmap(1.0);
let target = create_varmap(0.0);
// Apply 693 updates
for _ in 0..693 {
polyak_update(&online, &target, tau).unwrap();
}
let result = get_mean_value(&target);
// After 693 steps, should reach 50% of online value
assert!(
(result - 0.5).abs() < 0.05,
"After 693 steps with tau=0.001, target should ~=0.5, got {}",
result
);
println!(
"Rainbow tau=0.001 convergence: half-life={:.0} steps, empirical={:.4} (expected 0.5)",
half_life, result
);
}
#[test]
fn test_soft_vs_hard_update_stability() {
// GIVEN: Initial networks
let online = create_varmap(1.0);
let target_soft = create_varmap(0.0);
let target_hard = create_varmap(0.0);
// WHEN: Apply 10 soft updates vs 1 hard update
for _ in 0..10 {
polyak_update(&online, &target_soft, 0.1).unwrap();
}
hard_update(&online, &target_hard).unwrap();
let soft_result = get_mean_value(&target_soft);
let hard_result = get_mean_value(&target_hard);
// THEN: Hard update should jump directly to 1.0
// Soft updates should be gradual (10 steps at tau=0.1 ~= 0.65)
assert!(
(hard_result - 1.0).abs() < 1e-6,
"Hard update should copy fully, got {}",
hard_result
);
// After 10 steps with tau=0.1: (1-0.1)^10 ~= 0.349 remains -> 0.651 updated
assert!(
soft_result > 0.6 && soft_result < 0.7,
"Soft updates should be gradual ~=0.65, got {}",
soft_result
);
println!(
"Update comparison: soft={:.4} (gradual), hard={:.4} (instant)",
soft_result, hard_result
);
}
#[test]
fn test_divergence_with_changing_online_network() {
// GIVEN: Online network that changes over time
let mut divergences = Vec::new();
for step in 0..5 {
let online = create_varmap(step as f32);
let target = create_varmap(0.0);
// Apply tau=0.2 update
polyak_update(&online, &target, 0.2).unwrap();
let div = compute_network_divergence(&online, &target);
divergences.push(div);
}
// THEN: Divergence should increase as online network moves further
for i in 1..divergences.len() {
assert!(
divergences[i] > divergences[i - 1],
"Divergence should increase with larger online values: step {}: {:.4} vs {:.4}",
i,
divergences[i - 1],
divergences[i]
);
}
println!("Divergence tracking with changing online: {:?}", divergences);
}
#[test]
fn test_multiple_parameter_layers() {
// GIVEN: VarMaps with multiple layers
let online = VarMap::new();
let target = VarMap::new();
let device = Device::Cpu;
// Add 3 layers
let mut online_data = online.data().lock().unwrap();
let mut target_data = target.data().lock().unwrap();
for i in 1..=3 {
let w = (Tensor::ones(&[8, 8], DType::F32, &device).unwrap() * 1.0).unwrap();
let b = (Tensor::ones(&[8], DType::F32, &device).unwrap() * 1.0).unwrap();
let t_w = (Tensor::ones(&[8, 8], DType::F32, &device).unwrap() * 0.0).unwrap();
let t_b = (Tensor::ones(&[8], DType::F32, &device).unwrap() * 0.0).unwrap();
online_data.insert(format!("layer{}.weight", i), candle_core::Var::from_tensor(&w).unwrap());
online_data.insert(format!("layer{}.bias", i), candle_core::Var::from_tensor(&b).unwrap());
target_data.insert(format!("layer{}.weight", i), candle_core::Var::from_tensor(&t_w).unwrap());
target_data.insert(format!("layer{}.bias", i), candle_core::Var::from_tensor(&t_b).unwrap());
}
drop(online_data);
drop(target_data);
// WHEN: Apply soft update
polyak_update(&online, &target, 0.2).unwrap();
// THEN: All layers should be updated uniformly
let target_data = target.data().lock().unwrap();
for i in 1..=3 {
let weight = target_data.get(&format!("layer{}.weight", i)).unwrap();
let w_mean: f32 = weight.as_ref().mean_all().unwrap().to_scalar().unwrap();
assert!(
(w_mean - 0.2).abs() < 1e-6,
"Layer {} weight should be 0.2, got {}",
i,
w_mean
);
}
println!("Multiple layers updated uniformly");
}
#[test]
#[should_panic(expected = "Tau must be in [0.0, 1.0]")]
fn test_invalid_tau_panics() {
let online = create_varmap(1.0);
let target = create_varmap(0.0);
// Should panic with tau > 1.0
let _ = polyak_update(&online, &target, 1.5);
}
#[test]
fn test_convergence_half_life_different_tau_values() {
let test_cases = vec![
(0.001, 693.0), // Rainbow DQN
(0.005, 138.0), // 5x faster
(0.01, 69.0), // 10x faster
(0.05, 13.5), // 50x faster
(0.1, 6.6), // 100x faster
];
for (tau, expected_half_life) in test_cases {
let half_life = convergence_half_life(tau);
assert!(
(half_life - expected_half_life).abs() < expected_half_life * 0.1,
"tau={} should give half-life~={:.1}, got {:.1}",
tau,
expected_half_life,
half_life
);
}
println!("Convergence half-life verified for multiple tau values");
}