Files
rustytorch/crates/training/rtx-transformers/tests/integration_example_test.rs
T
2026-03-04 00:08:42 +00:00

159 lines
5.0 KiB
Rust

//! 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<dyn TrainableTransformerModel> = 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<dyn TrainableTransformerModel> = 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());
}
}