//! Debug test: Investigate target network initialization use anyhow::Result; use ml::dqn::dqn::{WorkingDQN, WorkingDQNConfig}; #[test] fn debug_initial_weight_equality() -> Result<()> { let config = WorkingDQNConfig { state_dim: 10, num_actions: 3, hidden_dims: vec![16], learning_rate: 0.001, gamma: 0.99, epsilon_start: 0.0, epsilon_end: 0.0, epsilon_decay: 1.0, replay_buffer_capacity: 1000, batch_size: 4, min_replay_size: 4, target_update_freq: 10, use_double_dqn: false, use_huber_loss: true, huber_delta: 1.0, leaky_relu_alpha: 0.01, gradient_clip_norm: 10.0, tau: 1.0, use_soft_updates: false, warmup_steps: 0, temperature_start: 1.0, temperature_decay: 0.99, }; let dqn = WorkingDQN::new(config)?; // Extract weights let online_vars = dqn.get_q_network_vars(); let target_vars = dqn.get_target_network_vars(); let online_data = online_vars.data().lock().unwrap(); let target_data = target_vars.data().lock().unwrap(); println!("Online network has {} variables", online_data.len()); println!("Target network has {} variables", target_data.len()); for (name, online_var) in online_data.iter() { if let Some(target_var) = target_data.get(name) { let online_tensor = online_var.as_tensor(); let target_tensor = target_var.as_tensor(); let online_flat = online_tensor.flatten_all().unwrap().to_vec1::().unwrap(); let target_flat = target_tensor.flatten_all().unwrap().to_vec1::().unwrap(); let differences: Vec = online_flat .iter() .zip(target_flat.iter()) .map(|(&a, &b)| (a - b).abs()) .collect(); let max_diff = differences.iter().cloned().fold(0.0f32, f32::max); let mean_diff = differences.iter().sum::() / differences.len() as f32; println!( "Variable '{}': {} weights, max_diff={:.8}, mean_diff={:.8}", name, online_flat.len(), max_diff, mean_diff ); // Sample some values if online_flat.len() > 0 { println!( " Sample: online[0]={:.8}, target[0]={:.8}, diff={:.8}", online_flat[0], target_flat[0], differences[0] ); } } else { println!("WARNING: Variable '{}' not found in target network!", name); } } Ok(()) }