130 lines
3.7 KiB
Rust
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]);
|
|
}
|