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

285 lines
8.3 KiB
Rust

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