Files
rustytorch/crates/models/rtx-multimodal/src/fusion/tests.rs
T
2026-03-04 00:08:42 +00:00

66 lines
2.2 KiB
Rust

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