//! Tests for the integration example module //! //! NOTE: Disabled until transformer training API is fully implemented #![cfg(feature = "disabled_tests")] use rtx_tensor::{Device, Tensor}; use rtx_transformers::{ TransformerError, training::{MockTransformerModel, TrainableTransformerModel}, }; use std::collections::HashMap; #[cfg(test)] mod integration_tests { use super::*; fn create_test_device() -> Device { Device::cpu() } #[test] fn test_mock_model_creation() { let model = MockTransformerModel::new(); assert!(model.is_ok()); } #[test] fn test_mock_model_implements_trainable() { let model = MockTransformerModel::new().unwrap(); // Verify it can be used as TrainableTransformerModel let _trainable: Box = Box::new(model); } #[test] fn test_mock_model_forward_pass() { let mut model = MockTransformerModel::new().unwrap(); // Create test input let input = Tensor::randn(&[2, 10], &create_test_device()).unwrap(); // Forward pass should work let output = model.forward(&input); assert!(output.is_ok()); let output = output.unwrap(); assert_eq!(output.shape().dims()[0], 2); // batch size assert_eq!(output.shape().dims()[1], 10); // sequence length assert_eq!(output.shape().dims()[2], 512); // hidden size } #[test] fn test_mock_model_parameters() { let model = MockTransformerModel::new().unwrap(); let params = model.parameters(); // Should have the expected parameters assert!(params.contains_key("embedding.weight")); assert!(params.contains_key("transformer.layer.0.attention.query.weight")); assert!(params.contains_key("transformer.layer.0.attention.key.weight")); assert!(params.contains_key("transformer.layer.0.attention.value.weight")); assert!(params.contains_key("transformer.layer.0.ffn.dense1.weight")); assert!(params.contains_key("transformer.layer.0.ffn.dense2.weight")); assert_eq!(params.len(), 6); } #[test] fn test_mock_model_update_parameters() { let mut model = MockTransformerModel::new().unwrap(); // Get initial parameter value let initial_params = model.parameters(); let initial_embedding = initial_params.get("embedding.weight").unwrap().clone(); // Create update let mut updates = HashMap::new(); let update_tensor = Tensor::ones(initial_embedding.shape().dims(), &create_test_device()).unwrap(); updates.insert("embedding.weight".to_string(), update_tensor); // Update parameters let result = model.update_parameters(&updates); assert!(result.is_ok()); // Check parameter was updated let updated_params = model.parameters(); let updated_embedding = updated_params.get("embedding.weight").unwrap(); // The updated parameter should be different from initial // (initial + ones != initial) assert!(updated_embedding != &initial_embedding); } #[test] fn test_mock_model_training_mode() { let mut model = MockTransformerModel::new().unwrap(); // Initially should be in eval mode assert!(!model.is_training()); // Set to training mode model.set_training(true); assert!(model.is_training()); // Set back to eval mode model.set_training(false); assert!(!model.is_training()); } #[test] fn test_mock_model_debug_impl() { let model = MockTransformerModel::new().unwrap(); let debug_str = format!("{:?}", model); assert!(debug_str.contains("MockTransformerModel")); } #[test] fn test_mock_model_as_trait_object() { let model = MockTransformerModel::new().unwrap(); let mut trainable: Box = Box::new(model); // Test that all trait methods work let input = Tensor::randn(&[1, 5], &create_test_device()).unwrap(); let forward_result = trainable.forward(&input); assert!(forward_result.is_ok()); let params = trainable.parameters(); assert!(!params.is_empty()); trainable.set_training(true); let model_type = trainable.model_type(); assert_eq!(model_type, "MockTransformer"); } #[test] fn test_mock_model_info() { let model = MockTransformerModel::new().unwrap(); let info = model.model_info(); // Should return some basic info assert!(info.contains_key("architecture")); assert!(info.contains_key("num_parameters")); } #[test] fn test_mock_model_backward_pass() { let mut model = MockTransformerModel::new().unwrap(); // Create mock loss let loss = Tensor::from_data(vec![1.0], vec![1], &create_test_device()).unwrap(); // Backward pass should work (even if it's a no-op for mock) let result = model.backward(&loss); assert!(result.is_ok()); } }