//! Tests for multimodal fusion. #[cfg(test)] mod fusion_tests { use super::super::*; use rtx_tensor::Device; #[tokio::test] async fn test_modality_fusion_creation() { let config = FusionConfig::new(768, 768, 768, 768, 12); let device = Device::cuda(0).unwrap_or(Device::default()); let fusion = ModalityFusion::new(config, &device); assert!(fusion.is_ok()); let fusion = fusion.unwrap(); assert_eq!(fusion.config().vision_dim, 768); assert_eq!(fusion.config().num_heads, 12); } #[tokio::test] async fn test_trimodal_fusion() { let config = FusionConfig::new(512, 512, 512, 512, 8); let device = Device::cuda(0).unwrap_or(Device::default()); let mut fusion = ModalityFusion::new(config, &device).unwrap(); let vision = Tensor::randn(&[2, 197, 512], &device).unwrap(); let audio = Tensor::randn(&[2, 500, 512], &device).unwrap(); let text = Tensor::randn(&[2, 128, 512], &device).unwrap(); let result = fusion.forward_trimodal(&vision, &audio, &text); assert!(result.is_ok()); let output = result.unwrap(); assert_eq!(output.shape()[0], 2); // Batch size preserved assert_eq!(output.shape()[1], 512); // Output dimension } #[test] fn test_fusion_strategies() { let device = Device::cuda(0).unwrap_or(Device::default()); let strategies = vec![ ModalityFusionStrategy::Concatenation, ModalityFusionStrategy::BilinearFusion, ModalityFusionStrategy::TensorFusion, ModalityFusionStrategy::AttentionFusion, ModalityFusionStrategy::HierarchicalAttention, ModalityFusionStrategy::GatedFusion, ModalityFusionStrategy::ContrastiveFusion, ]; for strategy in strategies { let mut config = FusionConfig::new(256, 256, 256, 256, 4); config.fusion_strategy = strategy.clone(); let fusion = ModalityFusion::new(config, &device); assert!( fusion.is_ok(), "Failed to create fusion with strategy {:?}", strategy ); } } }