Files
rustytorch/crates/models/rtx-multimodal/tests/basic_functionality_test.rs
T
2026-03-04 00:08:42 +00:00

130 lines
3.7 KiB
Rust

/// TDD Test: Verify basic rtx-multimodal functionality
/// RED -> GREEN -> REFACTOR
use rtx_multimodal::{
MultimodalConfig, MultimodalError, Result,
audio::AudioConfig,
cross_modal_attention::{CrossModalConfig, FusionStrategy},
fusion::{FusionConfig, ModalityFusionStrategy},
multimodal_preprocessing::MultimodalPreprocessingConfig,
vision::VisionConfig,
};
use rtx_tensor::{Device, Tensor};
#[test]
fn test_multimodal_config_creation() {
// Test basic configuration creation
let config = MultimodalConfig::new(768, 12);
assert_eq!(config.hidden_dim, 768);
assert_eq!(config.num_heads, 12);
}
#[test]
fn test_fusion_config_creation() {
let config = FusionConfig::new(512, 512, 512, 512, 8);
assert_eq!(config.vision_dim, 512);
assert_eq!(config.audio_dim, 512);
assert_eq!(config.text_dim, 512);
assert_eq!(config.output_dim, 512);
assert_eq!(config.num_heads, 8);
}
#[test]
fn test_cross_modal_config() {
let config = CrossModalConfig::new(384, 6);
assert_eq!(config.hidden_dim, 384);
assert_eq!(config.num_heads, 6);
assert_eq!(config.head_dim, 64); // 384 / 6
}
#[test]
fn test_vision_config() {
let config = VisionConfig::new(224, 16, 768, 12);
assert_eq!(config.image_size, 224);
assert_eq!(config.patch_size, 16);
assert_eq!(config.hidden_dim, 768);
assert_eq!(config.num_heads, 12);
}
#[test]
fn test_audio_config() {
let config = AudioConfig::new(80, 1000, 512, 8);
assert_eq!(config.mel_bins, 80);
assert_eq!(config.max_seq_len, 1000);
assert_eq!(config.hidden_dim, 512);
assert_eq!(config.num_heads, 8);
}
#[test]
fn test_preprocessing_config() {
let config = MultimodalPreprocessingConfig::default();
assert_eq!(config.vision_config.image_size, 224);
assert_eq!(config.audio_config.sample_rate, 16000);
assert_eq!(config.text_config.max_sequence_length, 512);
}
#[test]
fn test_modality_fusion_strategies() {
// Test all strategy variants exist and can be created
let strategies = vec![
ModalityFusionStrategy::Concatenation,
ModalityFusionStrategy::BilinearFusion,
ModalityFusionStrategy::TensorFusion,
ModalityFusionStrategy::AttentionFusion,
ModalityFusionStrategy::HierarchicalAttention,
ModalityFusionStrategy::GatedFusion,
ModalityFusionStrategy::ContrastiveFusion,
];
for strategy in strategies {
// Test that we can create a config with each strategy
let mut _config = FusionConfig::new(256, 256, 256, 256, 4);
_config.fusion_strategy = strategy.clone();
// Test Debug trait is implemented
let debug_str = format!("{:?}", strategy);
assert!(!debug_str.is_empty());
}
}
#[test]
fn test_cross_modal_fusion_strategies() {
let strategies = vec![
FusionStrategy::EarlyFusion,
FusionStrategy::AttentionFusion,
FusionStrategy::LateFusion,
FusionStrategy::HierarchicalFusion,
];
for strategy in strategies {
let mut _config = CrossModalConfig::new(256, 4);
_config.fusion_strategy = strategy.clone();
let debug_str = format!("{:?}", strategy);
assert!(!debug_str.is_empty());
}
}
#[test]
fn test_error_handling() {
// Test error creation
let error = MultimodalError::tensor("Test error".to_string());
let error_string = format!("{:?}", error);
assert!(error_string.contains("Test"));
// Test Result type
let result: Result<()> = Ok(());
assert!(result.is_ok());
}
#[test]
fn test_basic_tensor_operations() {
let device = Device::default();
// Test tensor creation
let tensor = Tensor::ones(&[2, 3], &device);
assert!(tensor.is_ok());
let t = tensor.unwrap();
assert_eq!(t.shape().dims(), &[2, 3]);
}