147 lines
4.4 KiB
Rust
147 lines
4.4 KiB
Rust
/// Tests for the public API of rtx-multimodal
|
|
/// Following strict TDD - testing only public interfaces
|
|
use rtx_tensor::{Device, Tensor};
|
|
|
|
#[test]
|
|
fn test_multimodal_config() {
|
|
// Test the main config structure
|
|
use rtx_multimodal::MultimodalConfig;
|
|
|
|
let config = MultimodalConfig::new(768, 12);
|
|
assert_eq!(config.hidden_dim, 768);
|
|
assert_eq!(config.num_heads, 12);
|
|
|
|
// Test with Flash Attention
|
|
let config_with_flash = config.with_flash_attention(true);
|
|
assert!(config_with_flash.vision_config.use_flash_attention);
|
|
assert!(config_with_flash.audio_config.use_flash_attention);
|
|
}
|
|
|
|
#[test]
|
|
fn test_fusion_config() {
|
|
use rtx_multimodal::fusion::{FusionConfig, ModalityFusionStrategy};
|
|
|
|
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 default strategy
|
|
let strategy = config.fusion_strategy.clone();
|
|
// Just check it exists - don't check specific value since it could change
|
|
match strategy {
|
|
ModalityFusionStrategy::Concatenation
|
|
| ModalityFusionStrategy::BilinearFusion
|
|
| ModalityFusionStrategy::TensorFusion
|
|
| ModalityFusionStrategy::AttentionFusion
|
|
| ModalityFusionStrategy::HierarchicalAttention
|
|
| ModalityFusionStrategy::GatedFusion
|
|
| ModalityFusionStrategy::ContrastiveFusion => assert!(true),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_cross_modal_config() {
|
|
use rtx_multimodal::cross_modal_attention::{CrossModalConfig, FusionStrategy};
|
|
|
|
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 fusion strategy
|
|
match config.fusion_strategy {
|
|
FusionStrategy::EarlyFusion
|
|
| FusionStrategy::AttentionFusion
|
|
| FusionStrategy::LateFusion
|
|
| FusionStrategy::HierarchicalFusion => assert!(true),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_vision_config() {
|
|
use rtx_multimodal::vision::VisionConfig;
|
|
|
|
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() {
|
|
use rtx_multimodal::audio::AudioConfig;
|
|
|
|
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() {
|
|
use rtx_multimodal::multimodal_preprocessing::MultimodalPreprocessingConfig;
|
|
|
|
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_error_types() {
|
|
use rtx_multimodal::{MultimodalError, Result};
|
|
|
|
// 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_modality_fusion_strategy() {
|
|
use rtx_multimodal::fusion::ModalityFusionStrategy;
|
|
|
|
// 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 Debug trait is implemented
|
|
let debug_str = format!("{:?}", strategy);
|
|
assert!(!debug_str.is_empty());
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_cross_modal_fusion_strategy() {
|
|
use rtx_multimodal::cross_modal_attention::FusionStrategy;
|
|
|
|
let strategies = vec![
|
|
FusionStrategy::EarlyFusion,
|
|
FusionStrategy::AttentionFusion,
|
|
FusionStrategy::LateFusion,
|
|
FusionStrategy::HierarchicalFusion,
|
|
];
|
|
|
|
for strategy in strategies {
|
|
let debug_str = format!("{:?}", strategy);
|
|
assert!(!debug_str.is_empty());
|
|
}
|
|
}
|