159 lines
5.0 KiB
Rust
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());
|
|
}
|
|
}
|