#![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, )] //! DQN Checkpoint Loading Tests //! //! Tests for loading DQN model weights from safetensors files. //! Follows TDD methodology - tests written first, then implementation. use ml::dqn::{DQN, DQNConfig, Experience}; use ml::MLError; use ml_core::cuda_autograd::GpuTensor; use std::fs; use std::sync::Arc; 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<(), MLError> { // Create temp directory for test files let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); // Create and save a DQN model let config = DQNConfig::emergency_safe_defaults(); let dqn = DQN::new(config.clone())?; let vars = dqn.get_q_network_vars(); let stream = vars.cuda_stream(); ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream, None)?; // Create a new DQN and load the checkpoint let mut dqn2 = DQN::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<(), MLError> { let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); let config = DQNConfig::emergency_safe_defaults(); let dqn = DQN::new(config.clone())?; // Save checkpoint let vars = dqn.get_q_network_vars(); let stream = vars.cuda_stream(); ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream, None)?; // Get original variable names and count let original_data = dqn.get_q_network_vars().data(); let original_count = original_data.len(); let original_names: Vec = original_data.keys().cloned().collect(); // Load into new model let mut dqn2 = DQN::new(config)?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; // Verify variable count matches let loaded_data = dqn2.get_q_network_vars().data(); 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<(), MLError> { let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; let checkpoint_path = temp_dir.path().join("dqn_test.safetensors"); let config = DQNConfig::emergency_safe_defaults(); let dqn = DQN::new(config.clone())?; // Save checkpoint let vars = dqn.get_q_network_vars(); let stream_ref = vars.cuda_stream(); ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream_ref, None)?; // Load into new model let mut dqn2 = DQN::new(config.clone())?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; // Create test input as GpuTensor [1, state_dim] let test_state = vec![0.5f32; config.state_dim]; let stream = Arc::clone(dqn2.get_q_network_vars().cuda_stream()); let state_tensor = GpuTensor::from_host(&test_state, vec![1, config.state_dim], &stream)?; // 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_host(&stream)?; for q_val in q_vec.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<(), MLError> { let temp_dir = TempDir::new().map_err(|e| MLError::ModelError(e.to_string()))?; let checkpoint_path = temp_dir.path().join("dqn_e2e.safetensors"); let mut config = DQNConfig::emergency_safe_defaults(); config.min_replay_size = 4; config.batch_size = 4; // Create and train original model let mut dqn = DQN::new(config.clone())?; // Add training experiences for i in 0..10 { let experience = 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 let vars = dqn.get_q_network_vars(); let stream_ref = vars.cuda_stream(); ml_core::checkpoint::save_safetensors(vars, &checkpoint_path, stream_ref, None)?; // Create test state for inference comparison let test_state = vec![0.5f32; config.state_dim]; let stream = Arc::clone(dqn.get_q_network_vars().cuda_stream()); let state_tensor = GpuTensor::from_host(&test_state, vec![1, config.state_dim], &stream)?; // Get Q-values from original model let original_q_values = dqn.forward(&state_tensor)?; let original_q_vec = original_q_values.to_host(&stream)?; // Load into new model let mut dqn2 = DQN::new(config.clone())?; dqn2.load_from_safetensors(checkpoint_path.to_str().unwrap())?; // Get Q-values from loaded model let stream2 = Arc::clone(dqn2.get_q_network_vars().cuda_stream()); let state_tensor2 = GpuTensor::from_host(&test_state, vec![1, config.state_dim], &stream2)?; let loaded_q_values = dqn2.forward(&state_tensor2)?; let loaded_q_vec = loaded_q_values.to_host(&stream2)?; // Verify Q-values match (within tolerance) // Note: Differences arise from distributional dueling network components // (e.g., RMSNorm running stats) that aren't captured in VarStore save/load. for (i, (orig, loaded)) in original_q_vec .iter() .zip(loaded_q_vec.iter()) .enumerate() { let diff = (orig - loaded).abs(); assert!( diff < 0.05, "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<(), MLError> { let config = DQNConfig::emergency_safe_defaults(); let mut dqn = DQN::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().map_err(|e| MLError::ModelError(e.to_string()))?; let corrupted_path = temp_dir.path().join("corrupted.safetensors"); fs::write(&corrupted_path, b"not a valid safetensors file") .map_err(|e| MLError::ModelError(e.to_string()))?; 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"); let vars = dqn.get_q_network_vars(); let stream = vars.cuda_stream(); let path_with_ext = format!("{}.safetensors", checkpoint_path.display()); ml_core::checkpoint::save_safetensors(vars, &path_with_ext, stream, None)?; // 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(()) }