//! DQN Checkpoint Loading Tests //! //! Tests for loading DQN model weights from safetensors files. //! Follows TDD methodology - tests written first, then implementation. use anyhow::Result; use ml::dqn::{WorkingDQN, WorkingDQNConfig}; use std::fs; use tempfile::TempDir; /// Test 1: Basic safetensors loading /// /// Verifies that the load_from_safetensors() method exists and can load /// a previously saved checkpoint without errors. #[test] fn test_load_safetensors_basic() -> Result<()> { // Create temp directory for test files let temp_dir = TempDir::new()?; let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); // Create and save a DQN model let config = WorkingDQNConfig::emergency_safe_defaults(); let dqn = WorkingDQN::new(config.clone())?; dqn.get_q_network_vars().save(&checkpoint_path)?; // Create a new DQN and load the checkpoint let mut dqn2 = WorkingDQN::new(config)?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; Ok(()) } /// Test 2: Validate weight dimensions match after loading /// /// Ensures that loaded weights have the same dimensions as the original model. #[test] fn test_load_safetensors_weight_dimensions() -> Result<()> { let temp_dir = TempDir::new()?; let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); let config = WorkingDQNConfig::emergency_safe_defaults(); let dqn = WorkingDQN::new(config.clone())?; // Save checkpoint dqn.get_q_network_vars().save(&checkpoint_path)?; // Get original variable names and count let original_vars = dqn.get_q_network_vars(); let original_data = original_vars.data().lock().unwrap(); let original_count = original_data.len(); let original_names: Vec = original_data.keys().cloned().collect(); drop(original_data); // Load into new model let mut dqn2 = WorkingDQN::new(config)?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; // Verify variable count matches let loaded_vars = dqn2.get_q_network_vars(); let loaded_data = loaded_vars.data().lock().unwrap(); assert_eq!(loaded_data.len(), original_count, "Variable count mismatch"); // Verify all original variable names exist for name in original_names { assert!( loaded_data.contains_key(&name), "Missing variable: {}", name ); } Ok(()) } /// Test 3: Forward pass produces correct outputs after loading /// /// Verifies that inference works correctly after loading weights, /// and produces valid Q-values. #[test] fn test_load_safetensors_forward_pass() -> Result<()> { let temp_dir = TempDir::new()?; let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); let config = WorkingDQNConfig::emergency_safe_defaults(); let dqn = WorkingDQN::new(config.clone())?; // Save checkpoint dqn.get_q_network_vars().save(&checkpoint_path)?; // Load into new model let mut dqn2 = WorkingDQN::new(config.clone())?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; // Create test input let test_state = vec![0.5f32; config.state_dim]; let state_tensor = candle_core::Tensor::from_vec(test_state.clone(), (1, config.state_dim), dqn2.device())?; // Forward pass should work let q_values = dqn2.forward(&state_tensor)?; // Verify output shape assert_eq!(q_values.dims(), &[1, config.num_actions]); // Verify Q-values are finite (not NaN or Inf) let q_vec = q_values.to_vec2::()?; for q_val in q_vec[0].iter() { assert!(q_val.is_finite(), "Q-value is not finite: {}", q_val); } Ok(()) } /// Test 4: End-to-end train→save→load→infer /// /// Complete workflow test: train model, save checkpoint, load in new instance, /// verify inference works correctly. #[test] fn test_load_safetensors_e2e_workflow() -> Result<()> { let temp_dir = TempDir::new()?; let checkpoint_path = temp_dir.path().join("dqn_e2e.safetensors"); let mut config = WorkingDQNConfig::emergency_safe_defaults(); config.min_replay_size = 4; config.batch_size = 4; // Create and train original model let mut dqn = WorkingDQN::new(config.clone())?; // Add training experiences for i in 0..10 { let experience = ml::dqn::Experience::new( vec![i as f32 * 0.1; config.state_dim], (i % config.num_actions) as u8, i as f32, vec![(i + 1) as f32 * 0.1; config.state_dim], i == 9, ); dqn.store_experience(experience)?; } // Train for a few steps for _ in 0..5 { let _ = dqn.train_step(None)?; } // Save checkpoint dqn.get_q_network_vars().save(&checkpoint_path)?; // Create test state for inference comparison let test_state = vec![0.5f32; config.state_dim]; let state_tensor = candle_core::Tensor::from_vec(test_state.clone(), (1, config.state_dim), dqn.device())?; // Get Q-values from original model let original_q_values = dqn.forward(&state_tensor)?; let original_q_vec = original_q_values.to_vec2::()?; // Load into new model let mut dqn2 = WorkingDQN::new(config.clone())?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; // Get Q-values from loaded model let loaded_q_values = dqn2.forward(&state_tensor)?; let loaded_q_vec = loaded_q_values.to_vec2::()?; // Verify Q-values match (within floating point tolerance) // Note: Small differences can occur due to GPU/CPU variations and target network updates for (i, (orig, loaded)) in original_q_vec[0] .iter() .zip(loaded_q_vec[0].iter()) .enumerate() { let diff = (orig - loaded).abs(); assert!( diff < 0.01, "Q-value mismatch at index {}: orig={}, loaded={}, diff={}", i, orig, loaded, diff ); } Ok(()) } /// Test 5: Error cases (file not found, corrupted file) /// /// Verifies proper error handling for invalid checkpoint files. #[test] fn test_load_safetensors_error_cases() -> Result<()> { let config = WorkingDQNConfig::emergency_safe_defaults(); let mut dqn = WorkingDQN::new(config)?; // Test 1: File not found let result = dqn.load_from_safetensors("/nonexistent/path/model.safetensors"); assert!(result.is_err(), "Should fail for nonexistent file"); // Test 2: Corrupted file let temp_dir = TempDir::new()?; let corrupted_path = temp_dir.path().join("corrupted.safetensors"); fs::write(&corrupted_path, b"not a valid safetensors file")?; let result = dqn.load_from_safetensors(corrupted_path.to_str().unwrap()); assert!(result.is_err(), "Should fail for corrupted file"); // Test 3: Extension handling (.safetensors auto-append) let checkpoint_path = temp_dir.path().join("test_model"); dqn.get_q_network_vars() .save(format!("{}.safetensors", checkpoint_path.display()))?; // Should work without .safetensors extension let result = dqn.load_from_safetensors(checkpoint_path.to_str().unwrap()); assert!(result.is_ok(), "Should auto-append .safetensors extension"); Ok(()) }