#![allow( clippy::assertions_on_constants, clippy::assertions_on_result_states, clippy::clone_on_copy, clippy::decimal_literal_representation, clippy::doc_markdown, clippy::empty_line_after_doc_comments, clippy::field_reassign_with_default, clippy::get_unwrap, clippy::identity_op, clippy::inconsistent_digit_grouping, clippy::indexing_slicing, clippy::integer_division, clippy::len_zero, clippy::let_underscore_must_use, clippy::manual_div_ceil, clippy::manual_let_else, clippy::manual_range_contains, clippy::modulo_arithmetic, clippy::needless_range_loop, clippy::non_ascii_literal, clippy::redundant_clone, clippy::shadow_reuse, clippy::shadow_same, clippy::shadow_unrelated, clippy::single_match_else, clippy::str_to_string, clippy::string_slice, clippy::tests_outside_test_module, clippy::too_many_lines, clippy::unnecessary_wraps, clippy::unseparated_literal_suffix, clippy::use_debug, clippy::useless_vec, clippy::wildcard_enum_match_arm, clippy::else_if_without_else, clippy::expect_used, clippy::missing_const_for_fn, clippy::similar_names, clippy::type_complexity, clippy::collapsible_else_if, clippy::doc_lazy_continuation, clippy::items_after_test_module, clippy::map_clone, clippy::multiple_unsafe_ops_per_block, clippy::unwrap_or_default, clippy::assign_op_pattern, clippy::needless_borrow, clippy::println_empty_string, clippy::unnecessary_cast, clippy::used_underscore_binding, clippy::create_dir, clippy::implicit_saturating_sub, clippy::exit, clippy::expect_fun_call, clippy::too_many_arguments, clippy::unnecessary_map_or, clippy::unwrap_used, dead_code, unused_imports, unused_variables, clippy::cloned_ref_to_slice_refs, clippy::neg_multiply, clippy::while_let_loop, clippy::bool_assert_comparison, clippy::excessive_precision, clippy::trivially_copy_pass_by_ref, clippy::op_ref, clippy::redundant_closure, clippy::unnecessary_lazy_evaluations, clippy::if_then_some_else_none, clippy::unnecessary_to_owned, clippy::single_component_path_imports, )] //! Comprehensive TDD tests for Polyak soft updates (WAVE 26 P1.12) //! //! Verifies: //! 1. Tau is configurable and defaults to 0.005 //! 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 std::sync::Arc; use std::sync::OnceLock; use cudarc::driver::CudaStream; use tracing::info; use ml_core::cuda_autograd::GpuVarStore; use ml_core::device::MlDevice; use ml::dqn::target_update::{polyak_update, hard_update, convergence_half_life, compute_network_divergence}; use ml::trainers::dqn::DQNHyperparameters; static SHARED_CUDA: OnceLock = OnceLock::new(); fn cuda_device() -> MlDevice { SHARED_CUDA.get_or_init(|| MlDevice::cuda(0).expect("CUDA device required")).clone() } fn make_stream() -> Arc { let device = cuda_device(); device.cuda_stream().expect("stream").clone() } /// Helper: Create GpuVarStore with uniform values (bf16 storage) fn create_varstore(value: f32, stream: &Arc) -> GpuVarStore { let mut store = GpuVarStore::new(stream.clone()); // Create test tensors with uniform value (bf16 to match register signature) let weight_host: Vec = vec![half::bf16::from_f32(value); 10 * 10]; let bias_host: Vec = vec![half::bf16::from_f32(value); 10]; let mut w_data = stream.alloc_zeros::(100).unwrap(); stream.memcpy_htod(&weight_host, &mut w_data).unwrap(); store.register("layer1.weight", w_data, vec![10, 10]).unwrap(); let mut b_data = stream.alloc_zeros::(10).unwrap(); stream.memcpy_htod(&bias_host, &mut b_data).unwrap(); store.register("layer1.bias", b_data, vec![10]).unwrap(); store } /// Helper: Extract mean value from GpuVarStore (test-only readback, bf16 -> f32) fn get_mean_value(store: &GpuVarStore, stream: &Arc) -> f32 { let mut sum = 0.0_f32; let mut count = 0; for name in store.param_names() { let param = store.get(name).expect("param must exist"); let n = param.data.len(); let mut buf = vec![half::bf16::ZERO; n]; stream.memcpy_dtoh(¶m.data, &mut buf).unwrap(); // test-only readback let param_sum: f32 = buf.iter().map(|x| x.to_f32()).sum(); sum += param_sum / n as f32; count += 1; } sum / count as f32 } #[test] fn test_tau_default_value() { // GIVEN: DQN config should default to tau=0.005 let config = DQNHyperparameters::default(); // THEN: Tau should be 0.005 (Rainbow DQN standard) assert!( (config.tau - 0.005).abs() < 1e-6, "Default tau should be 0.005, got {}", config.tau, ); info!(tau = config.tau, "Default tau (Rainbow DQN standard)"); } #[test] fn test_soft_update_formula_correctness() { let stream = make_stream(); // GIVEN: Online network at 1.0, target at 0.0 let online_vars = create_varstore(1.0, &stream); let mut target_vars = create_varstore(0.0, &stream); let tau = 0.3; // Use larger tau for easier verification // WHEN: Apply single Polyak update polyak_update(&online_vars, &mut target_vars, tau, &stream).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, &stream); assert!( (result - 0.3).abs() < 1e-5, "Expected 0.3, got {}. Formula: tau*1.0 + (1-tau)*0.0", result ); info!(result, "Soft update formula correct (expected 0.3)"); } #[test] fn test_network_divergence_computation() { let stream = make_stream(); // GIVEN: Two networks with known values let online = create_varstore(1.0, &stream); let target = create_varstore(0.5, &stream); // WHEN: Compute divergence let divergence = compute_network_divergence(&online, &target, &stream).unwrap(); // 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 ); info!(divergence, "Network divergence (L2 norm)"); } #[test] fn test_divergence_decreases_with_updates() { let stream = make_stream(); // GIVEN: Networks starting far apart let online = create_varstore(1.0, &stream); let mut target = create_varstore(0.0, &stream); let initial_divergence = compute_network_divergence(&online, &target, &stream).unwrap(); // WHEN: Apply multiple soft updates for _ in 0..10 { polyak_update(&online, &mut target, 0.1, &stream).unwrap(); } let final_divergence = compute_network_divergence(&online, &target, &stream).unwrap(); // 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 ); info!( initial_divergence, final_divergence, reduction_pct = 100.0 * (1.0 - final_divergence / initial_divergence), "Divergence decreased" ); } #[test] fn test_tau_boundary_condition_zero() { let stream = make_stream(); // GIVEN: Online at 1.0, target at 0.0 let online = create_varstore(1.0, &stream); let mut target = create_varstore(0.0, &stream); // WHEN: tau = 0 (no update) polyak_update(&online, &mut target, 0.0, &stream).unwrap(); // THEN: Target should remain 0.0 let result = get_mean_value(&target, &stream); assert!( result.abs() < 1e-5, "tau=0 should not update target, got {}", result ); info!(result, "tau=0 boundary condition: target unchanged"); } #[test] fn test_tau_boundary_condition_one() { let stream = make_stream(); // GIVEN: Online at 1.0, target at 0.0 let online = create_varstore(1.0, &stream); let mut target = create_varstore(0.0, &stream); // WHEN: tau = 1.0 (full copy, equivalent to hard update) polyak_update(&online, &mut target, 1.0, &stream).unwrap(); // THEN: Target should equal online (1.0) let result = get_mean_value(&target, &stream); assert!( (result - 1.0).abs() < 1e-5, "tau=1.0 should copy online to target, got {}", result ); info!(result, "tau=1.0 boundary condition: target = online"); } #[test] fn test_rainbow_tau_convergence_rate() { let stream = make_stream(); // 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_varstore(1.0, &stream); let mut target = create_varstore(0.0, &stream); // Apply 693 updates for _ in 0..693 { polyak_update(&online, &mut target, tau, &stream).unwrap(); } let result = get_mean_value(&target, &stream); // 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 ); info!( half_life, empirical = result, "Rainbow tau=0.001 convergence (expected empirical ~0.5)" ); } #[test] fn test_soft_vs_hard_update_stability() { let stream = make_stream(); // GIVEN: Initial networks let online = create_varstore(1.0, &stream); let mut target_soft = create_varstore(0.0, &stream); let mut target_hard = create_varstore(0.0, &stream); // WHEN: Apply 10 soft updates vs 1 hard update for _ in 0..10 { polyak_update(&online, &mut target_soft, 0.1, &stream).unwrap(); } hard_update(&online, &mut target_hard, &stream).unwrap(); let soft_result = get_mean_value(&target_soft, &stream); let hard_result = get_mean_value(&target_hard, &stream); // 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-5, "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 ); info!( soft_result, hard_result, "Update comparison: soft (gradual) vs hard (instant)" ); } #[test] fn test_divergence_with_changing_online_network() { let stream = make_stream(); // GIVEN: Online network that changes over time let mut divergences = Vec::new(); for step in 0..5 { let online = create_varstore(step as f32, &stream); let mut target = create_varstore(0.0, &stream); // Apply tau=0.2 update polyak_update(&online, &mut target, 0.2, &stream).unwrap(); let div = compute_network_divergence(&online, &target, &stream).unwrap(); 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] ); } info!(?divergences, "Divergence tracking with changing online"); } #[test] fn test_multiple_parameter_layers() { let stream = make_stream(); // GIVEN: GpuVarStores with multiple layers let mut online = GpuVarStore::new(stream.clone()); let mut target = GpuVarStore::new(stream.clone()); // Add 3 layers (bf16 storage) for i in 1..=3 { let w_host: Vec = vec![half::bf16::from_f32(1.0); 8 * 8]; let b_host: Vec = vec![half::bf16::from_f32(1.0); 8]; let tw_host: Vec = vec![half::bf16::from_f32(0.0); 8 * 8]; let tb_host: Vec = vec![half::bf16::from_f32(0.0); 8]; let mut w_data = stream.alloc_zeros::(64).unwrap(); stream.memcpy_htod(&w_host, &mut w_data).unwrap(); online.register(&format!("layer{i}.weight"), w_data, vec![8, 8]).unwrap(); let mut b_data = stream.alloc_zeros::(8).unwrap(); stream.memcpy_htod(&b_host, &mut b_data).unwrap(); online.register(&format!("layer{i}.bias"), b_data, vec![8]).unwrap(); let mut tw_data = stream.alloc_zeros::(64).unwrap(); stream.memcpy_htod(&tw_host, &mut tw_data).unwrap(); target.register(&format!("layer{i}.weight"), tw_data, vec![8, 8]).unwrap(); let mut tb_data = stream.alloc_zeros::(8).unwrap(); stream.memcpy_htod(&tb_host, &mut tb_data).unwrap(); target.register(&format!("layer{i}.bias"), tb_data, vec![8]).unwrap(); } // WHEN: Apply soft update polyak_update(&online, &mut target, 0.2, &stream).unwrap(); // THEN: All layers should be updated uniformly for i in 1..=3 { let param = target.get(&format!("layer{i}.weight")).expect("param must exist"); let n = param.data.len(); let mut buf = vec![half::bf16::ZERO; n]; stream.memcpy_dtoh(¶m.data, &mut buf).unwrap(); // test-only readback let w_mean: f32 = buf.iter().map(|x| x.to_f32()).sum::() / n as f32; assert!( (w_mean - 0.2).abs() < 1e-5, "Layer {} weight should be 0.2, got {}", i, w_mean ); } info!("Multiple layers updated uniformly"); } #[test] #[should_panic(expected = "Tau must be in [0.0, 1.0]")] fn test_invalid_tau_panics() { let stream = make_stream(); let online = create_varstore(1.0, &stream); let mut target = create_varstore(0.0, &stream); // Should panic with tau > 1.0 let _ = polyak_update(&online, &mut target, 1.5, &stream); } #[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 ); } info!("Convergence half-life verified for multiple tau values"); }