/// 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()); } }