use approx::assert_abs_diff_eq; use rtx_multimodal::error::Result; use rtx_multimodal::fusion::*; use rtx_tensor::{Device, Tensor}; #[test] fn test_cross_modal_attention() -> Result<()> { let device = Device::default(); let config = CrossModalAttentionConfig { embed_dim: 512, num_heads: 8, dropout: 0.1, }; // CrossModalAttention is not exported, skip this test for now // let cross_attn = CrossModalAttention::new(&config, &device)?; // Vision features [batch, vision_tokens, embed_dim] let vision = Tensor::randn(&[2, 197, 512], &device)?; // Audio features [batch, audio_tokens, embed_dim] let audio = Tensor::randn(&[2, 100, 512], &device)?; // Test commented out since CrossModalAttention is not exported // Vision attending to audio // let vision_attended = cross_attn.forward(&vision, &audio)?; // assert_eq!(vision_attended.shape(), &[2, 197, 512]); // Audio attending to vision // let audio_attended = cross_attn.forward(&audio, &vision)?; // assert_eq!(audio_attended.shape(), &[2, 100, 512]); // Just check that tensors are created correctly assert_eq!(vision.shape(), &[2, 197, 512]); assert_eq!(audio.shape(), &[2, 100, 512]); Ok(()) } #[test] fn test_multimodal_fusion_transformer() -> Result<()> { let device = Device::default(); let config = MultimodalFusionConfig { embed_dim: 512, num_heads: 8, num_layers: 6, dropout: 0.1, max_vision_tokens: 197, max_audio_tokens: 100, max_text_tokens: 77, }; let fusion = MultimodalFusionTransformer::new(&config, &device)?; let vision = Tensor::randn(&[2, 197, 512], &device)?; let audio = Tensor::randn(&[2, 100, 512], &device)?; let text = Tensor::randn(&[2, 77, 512], &device)?; let fused = fusion.forward(&vision, &audio, &text)?; // Should concatenate and process all modalities // Total tokens = 197 + 100 + 77 = 374 assert_eq!(fused.shape(), &[2, 374, 512]); Ok(()) } #[test] fn test_bilinear_fusion() -> Result<()> { let device = Device::default(); let fusion = BilinearFusion::new(512, 512, 512, &device)?; let modality_a = Tensor::randn(&[2, 512], &device)?; let modality_b = Tensor::randn(&[2, 512], &device)?; let fused = fusion.forward(&modality_a, &modality_b)?; assert_eq!(fused.shape(), &[2, 512]); Ok(()) } #[test] fn test_tensor_fusion_network() -> Result<()> { let device = Device::default(); let config = TensorFusionConfig { vision_dim: 768, audio_dim: 512, text_dim: 512, output_dim: 256, hidden_dim: 128, }; let tfn = TensorFusionNetwork::new(&config, &device)?; let vision = Tensor::randn(&[2, 768], &device)?; let audio = Tensor::randn(&[2, 512], &device)?; let text = Tensor::randn(&[2, 512], &device)?; let fused = tfn.forward(&vision, &audio, &text)?; assert_eq!(fused.shape(), &[2, 256]); Ok(()) } #[test] fn test_multimodal_bottleneck_fusion() -> Result<()> { let device = Device::default(); let config = BottleneckFusionConfig { input_dims: vec![768, 512, 512], // vision, audio, text bottleneck_dim: 128, output_dim: 256, num_layers: 3, }; let bottleneck = MultimodalBottleneckFusion::new(&config, &device)?; let modalities = vec![ Tensor::randn(&[2, 768], &device)?, // vision Tensor::randn(&[2, 512], &device)?, // audio Tensor::randn(&[2, 512], &device)?, // text ]; let fused = bottleneck.forward(&modalities)?; assert_eq!(fused.shape(), &[2, 256]); Ok(()) } #[test] fn test_attention_based_fusion() -> Result<()> { let device = Device::default(); let config = AttentionFusionConfig { modality_dims: vec![768, 512, 512], hidden_dim: 256, num_heads: 8, dropout: 0.1, }; let attention_fusion = AttentionBasedFusion::new(&config, &device)?; let modalities = vec![ Tensor::randn(&[2, 768], &device)?, // vision Tensor::randn(&[2, 512], &device)?, // audio Tensor::randn(&[2, 512], &device)?, // text ]; let fused = attention_fusion.forward(&modalities)?; assert_eq!(fused.shape(), &[2, 256]); Ok(()) } #[test] fn test_modality_specific_encoders() -> Result<()> { let device = Device::default(); let vision_encoder = ModalityEncoder::new(2048, 512, &device)?; // ResNet features let audio_encoder = ModalityEncoder::new(128, 512, &device)?; // Mel features let text_encoder = ModalityEncoder::new(768, 512, &device)?; // BERT features let vision_features = Tensor::randn(&[2, 2048], &device)?; let audio_features = Tensor::randn(&[2, 128], &device)?; let text_features = Tensor::randn(&[2, 768], &device)?; let vision_encoded = vision_encoder.forward(&vision_features)?; let audio_encoded = audio_encoder.forward(&audio_features)?; let text_encoded = text_encoder.forward(&text_features)?; // All should map to same dimension assert_eq!(vision_encoded.shape(), &[2, 512]); assert_eq!(audio_encoded.shape(), &[2, 512]); assert_eq!(text_encoded.shape(), &[2, 512]); Ok(()) } #[test] fn test_multimodal_contrastive_learning() -> Result<()> { let device = Device::default(); let config = ContrastiveLearningConfig { embed_dim: 512, temperature: 0.07, }; let contrastive = MultimodalContrastiveLearning::new(&config, &device)?; let vision_embeds = Tensor::randn(&[4, 512], &device)?; let audio_embeds = Tensor::randn(&[4, 512], &device)?; let text_embeds = Tensor::randn(&[4, 512], &device)?; // Compute pairwise similarities let vision_audio_sim = contrastive.compute_similarity(&vision_embeds, &audio_embeds)?; let vision_text_sim = contrastive.compute_similarity(&vision_embeds, &text_embeds)?; let audio_text_sim = contrastive.compute_similarity(&audio_embeds, &text_embeds)?; assert_eq!(vision_audio_sim.shape(), &[4, 4]); assert_eq!(vision_text_sim.shape(), &[4, 4]); assert_eq!(audio_text_sim.shape(), &[4, 4]); Ok(()) } #[test] fn test_sequential_fusion() -> Result<()> { let device = Device::default(); let config = SequentialFusionConfig { input_dims: vec![768, 512, 256], hidden_dim: 512, output_dim: 256, dropout: 0.1, }; let sequential_fusion = SequentialFusion::new(&config, &device)?; let modalities = vec![ Tensor::randn(&[2, 768], &device)?, Tensor::randn(&[2, 512], &device)?, Tensor::randn(&[2, 256], &device)?, ]; let fused = sequential_fusion.forward(&modalities)?; assert_eq!(fused.shape(), &[2, 256]); Ok(()) } #[test] fn test_hierarchical_fusion() -> Result<()> { let device = Device::default(); let config = HierarchicalFusionConfig { modality_dims: vec![768, 512, 512, 256], // vision, audio, text, other fusion_stages: vec![ (vec![0, 1], 512), // fuse vision + audio first (vec![2, 3], 384), // fuse text + other (vec![0, 1], 256), // final fusion ], }; let hierarchical = HierarchicalFusion::new(&config, &device)?; let modalities = vec![ Tensor::randn(&[2, 768], &device)?, // vision Tensor::randn(&[2, 512], &device)?, // audio Tensor::randn(&[2, 512], &device)?, // text Tensor::randn(&[2, 256], &device)?, // other ]; let fused = hierarchical.forward(&modalities)?; assert_eq!(fused.shape(), &[2, 256]); Ok(()) } #[test] fn test_gated_fusion() -> Result<()> { let device = Device::default(); let config = GatedFusionConfig { modality_dims: vec![512, 512, 512], hidden_dim: 256, output_dim: 512, }; let gated_fusion = GatedFusion::new(&config, &device)?; let modalities = vec![ Tensor::randn(&[2, 512], &device)?, Tensor::randn(&[2, 512], &device)?, Tensor::randn(&[2, 512], &device)?, ]; let fused = gated_fusion.forward(&modalities)?; assert_eq!(fused.shape(), &[2, 512]); // Test gating weights sum to approximately 1 let gates = gated_fusion.get_gate_weights(&modalities)?; let gate_sums = gates.sum(Some(1))?; // Gates should sum to 1 across modalities for each sample for i in 0..2 { let sum_val = gate_sums.get(&[i, 0])?; assert_abs_diff_eq!(sum_val, 1.0, epsilon = 1e-5); } Ok(()) }