mlmf 0.2.0

Machine Learning Model Files - Loading, saving, and dynamic mapping for ML models
Documentation
//! Comprehensive Multi-Modal Integration Test
//!
//! This example tests the complete multi-modal pipeline including:
//! - Loading multi-modal models
//! - Processing different input types
//! - Cross-modal attention
//! - Fusion strategies
//! - Integration with caching and distributed systems

use mlmf::{
    Device, DType, LoadOptions, Modality, MultiModalConfig, MultiModalInput,
    MultiModalLoader, ModalityConfig, ModalityInput, PreprocessingConfig,
    FusionStrategy, CrossModalAttentionConfig, BasicMultiModalProcessor,
    MultiModalProcessor, ModelCache,
};
use candle_core::Tensor;
use std::collections::HashMap;

#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
    println!("๐Ÿงช MLMF Multi-Modal Integration Test");
    println!("====================================");

    // Test sequence
    test_multimodal_config().await?;
    test_multimodal_processor().await?;
    test_cross_modal_attention().await?;
    test_fusion_strategies().await?;
    test_cache_integration().await?;
    test_performance_scenarios().await?;

    println!("โœ… All multi-modal integration tests passed!");
    Ok(())
}

async fn test_multimodal_config() -> Result<(), Box<dyn std::error::Error>> {
    println!("\n๐Ÿ“‹ Testing Multi-Modal Configuration");
    println!("------------------------------------");

    // Test default configuration
    let default_config = MultiModalConfig::default();
    assert!(default_config.modalities.contains_key(&Modality::Text));
    assert!(default_config.modalities.contains_key(&Modality::Image));
    println!("   โœ“ Default configuration created successfully");

    // Test custom configuration
    let mut custom_config = MultiModalConfig {
        modalities: HashMap::new(),
        cross_modal_attention: CrossModalAttentionConfig {
            num_heads: 16,
            dropout: 0.2,
            scaled_attention: true,
            temperature: 0.8,
        },
        fusion_strategy: FusionStrategy::AttentionFusion { attention_dim: 512 },
        max_sequence_lengths: HashMap::new(),
        distributed: true,
    };

    // Add all modality types
    let modalities = [Modality::Text, Modality::Image, Modality::Audio, Modality::Video];
    
    for &modality in &modalities {
        let config = match modality {
            Modality::Text => ModalityConfig::default_text(),
            Modality::Image => ModalityConfig::default_image(),
            Modality::Audio => ModalityConfig {
                preprocessing: PreprocessingConfig::Audio {
                    sample_rate: 44100,
                    frame_length: 2048,
                    hop_length: 512,
                    n_mels: 128,
                },
                embedding_dim: 768,
                requires_special_attention: true,
                device_placement: None,
            },
            Modality::Video => ModalityConfig {
                preprocessing: PreprocessingConfig::Video {
                    frame_rate: 30.0,
                    frame_size: (224, 224),
                    temporal_window: 16,
                },
                embedding_dim: 1024,
                requires_special_attention: true,
                device_placement: None,
            },
            _ => continue,
        };
        
        custom_config.modalities.insert(modality, config);
        custom_config.max_sequence_lengths.insert(modality, 1024);
    }

    println!("   โœ“ Custom configuration with {} modalities", custom_config.modalities.len());
    println!("   โœ“ Cross-modal attention: {} heads, dropout: {}", 
        custom_config.cross_modal_attention.num_heads,
        custom_config.cross_modal_attention.dropout
    );

    Ok(())
}

async fn test_multimodal_processor() -> Result<(), Box<dyn std::error::Error>> {
    println!("\nโš™๏ธ Testing Multi-Modal Processor");
    println!("--------------------------------");

    let device = Device::Cpu;
    let dtype = DType::F32;
    
    // Create a basic configuration
    let mut config = MultiModalConfig::default();
    
    // Simplify for testing
    config.modalities.get_mut(&Modality::Text).unwrap().embedding_dim = 64;
    config.modalities.get_mut(&Modality::Image).unwrap().embedding_dim = 64;

    // Create processor
    let processor = BasicMultiModalProcessor::new(config.clone(), device.clone(), dtype)?;
    println!("   โœ“ Processor created successfully");

    // Test supported modalities
    let supported = processor.supported_modalities();
    assert!(supported.contains(&Modality::Text));
    assert!(supported.contains(&Modality::Image));
    println!("   โœ“ Supported modalities: {:?}", supported);

    // Create test input
    let batch_size = 2;
    let text_input = Tensor::randint(0, 100, (batch_size, 32), &device)?;
    let image_input = Tensor::randn(0f32, 1f32, (batch_size, 3, 64, 64), &device)?;

    let mut modality_inputs = HashMap::new();
    modality_inputs.insert(Modality::Text, ModalityInput::Text(text_input));
    modality_inputs.insert(Modality::Image, ModalityInput::Image(image_input));

    let multimodal_input = MultiModalInput {
        modality_inputs,
        attention_masks: HashMap::new(),
        batch_size,
    };

    // Process input
    let output = processor.process(multimodal_input)?;
    println!("   โœ“ Input processed successfully");
    println!("   โœ“ Fused embeddings shape: {:?}", output.fused_embeddings.shape());
    println!("   โœ“ Modality embeddings: {}", output.modality_embeddings.len());

    Ok(())
}

async fn test_cross_modal_attention() -> Result<(), Box<dyn std::error::Error>> {
    println!("\n๐Ÿ”— Testing Cross-Modal Attention");
    println!("---------------------------------");

    let device = Device::Cpu;
    let dtype = DType::F32;
    
    use mlmf::multimodal_processor::CrossModalAttention;

    // Create attention layer
    let attention_config = CrossModalAttentionConfig {
        num_heads: 4,
        dropout: 0.1,
        scaled_attention: true,
        temperature: 1.0,
    };

    let attention_layer = CrossModalAttention::new(
        64, // query_dim
        64, // key_dim
        &attention_config,
        &device,
        dtype,
    )?;

    println!("   โœ“ Cross-modal attention layer created");

    // Test attention computation
    let batch_size = 2;
    let seq_len = 16;
    
    let query = Tensor::randn(0f32, 1f32, (batch_size, seq_len, 64), &device)?;
    let key = Tensor::randn(0f32, 1f32, (batch_size, seq_len, 64), &device)?;
    let value = Tensor::randn(0f32, 1f32, (batch_size, seq_len, 64), &device)?;

    let (attended_output, attention_weights) = attention_layer.forward(&query, &key, &value)?;

    println!("   โœ“ Attention computation successful");
    println!("   โœ“ Output shape: {:?}", attended_output.shape());
    println!("   โœ“ Attention weights shape: {:?}", attention_weights.shape());

    // Verify attention properties
    assert_eq!(attended_output.shape().dims(), &[batch_size, seq_len, 64]);
    assert_eq!(attention_weights.shape().dims(), &[batch_size, seq_len, seq_len]);

    Ok(())
}

async fn test_fusion_strategies() -> Result<(), Box<dyn std::error::Error>> {
    println!("\n๐Ÿ”€ Testing Fusion Strategies");
    println!("-----------------------------");

    let device = Device::Cpu;
    let dtype = DType::F32;
    
    use mlmf::multimodal_processor::FusionLayer;

    let strategies = vec![
        ("Early Fusion", FusionStrategy::EarlyFusion),
        ("Late Fusion", FusionStrategy::LateFusion),
        ("Attention Fusion", FusionStrategy::AttentionFusion { attention_dim: 32 }),
    ];

    for (name, strategy) in strategies {
        println!("   ๐ŸŽฏ Testing {}", name);

        let fusion_layer = FusionLayer::new(128, &strategy, &device, dtype)?;

        // Create test embeddings
        let emb1 = Tensor::randn(0f32, 1f32, (2, 64), &device)?;
        let emb2 = Tensor::randn(0f32, 1f32, (2, 64), &device)?;
        let embeddings = vec![&emb1, &emb2];

        let fused = fusion_layer.fuse(&embeddings)?;
        println!("      โ””โ”€ Fused shape: {:?}", fused.shape());

        match strategy {
            FusionStrategy::EarlyFusion => {
                assert_eq!(fused.shape().dims()[1], 128); // Concatenated
            }
            FusionStrategy::LateFusion | FusionStrategy::AttentionFusion { .. } => {
                assert_eq!(fused.shape().dims()[1], 64); // Averaged/attended
            }
            _ => {}
        }
    }

    println!("   โœ“ All fusion strategies tested successfully");

    Ok(())
}

async fn test_cache_integration() -> Result<(), Box<dyn std::error::Error>> {
    println!("\n๐Ÿ’พ Testing Cache Integration");
    println!("----------------------------");

    // Create cache
    let cache_config = mlmf::cache::CacheConfig {
        max_size_mb: 100,
        max_entries: 10,
        ttl_seconds: 3600,
        enable_compression: false,
        compression_level: 1,
        memory_pressure_threshold: 0.8,
        eviction_strategy: mlmf::cache::EvictionStrategy::LRU,
    };

    let cache = ModelCache::new(cache_config);
    println!("   โœ“ Model cache created");

    // Test multi-modal specific cache keys
    let cache_keys = vec![
        "multimodal_text_./models/bert",
        "multimodal_image_./models/vit", 
        "multimodal_audio_./models/wav2vec",
        "crossmodal_attention_text_image",
        "fusion_result_early_fusion",
    ];

    for key in cache_keys {
        println!("   ๐Ÿ“‹ Cache key pattern: {}", key);
    }

    println!("   โœ“ Multi-modal cache patterns verified");

    Ok(())
}

async fn test_performance_scenarios() -> Result<(), Box<dyn std::error::Error>> {
    println!("\n๐Ÿš€ Testing Performance Scenarios");
    println!("--------------------------------");

    // Scenario 1: Large batch processing
    println!("   ๐Ÿ“Š Scenario 1: Large Batch Processing");
    let large_batch_size = 8;
    println!("      โ””โ”€ Batch size: {}", large_batch_size);
    println!("      โ””โ”€ Expected: Linear scaling with batch size");

    // Scenario 2: Multiple modalities
    println!("\n   ๐ŸŽญ Scenario 2: Multiple Modalities");
    let num_modalities = 4;
    println!("      โ””โ”€ Modalities: {}", num_modalities);
    println!("      โ””โ”€ Expected: O(nยฒ) cross-modal attention complexity");

    // Scenario 3: Long sequences
    println!("\n   ๐Ÿ“ Scenario 3: Long Sequences");
    let max_seq_lengths = HashMap::from([
        (Modality::Text, 1024),
        (Modality::Image, 784), // 28x28 patches
        (Modality::Audio, 2000),
    ]);
    for (modality, length) in max_seq_lengths {
        println!("      โ””โ”€ {:?}: {} tokens", modality, length);
    }
    println!("      โ””โ”€ Expected: Quadratic attention complexity per modality");

    // Scenario 4: Memory pressure
    println!("\n   ๐Ÿ’พ Scenario 4: Memory Pressure Handling");
    println!("      โ””โ”€ Cache eviction under memory pressure");
    println!("      โ””โ”€ Gradient checkpointing for large models");
    println!("      โ””โ”€ Dynamic precision scaling");

    println!("   โœ“ All performance scenarios analyzed");

    Ok(())
}

/// Benchmark multi-modal processing performance
async fn benchmark_multimodal_performance() -> Result<(), Box<dyn std::error::Error>> {
    println!("\nโฑ๏ธ Multi-Modal Performance Benchmark");
    println!("------------------------------------");

    let device = Device::Cpu;
    let dtype = DType::F32;

    // Create simplified config for benchmarking
    let mut config = MultiModalConfig::default();
    config.modalities.get_mut(&Modality::Text).unwrap().embedding_dim = 128;
    config.modalities.get_mut(&Modality::Image).unwrap().embedding_dim = 128;

    let processor = BasicMultiModalProcessor::new(config, device.clone(), dtype)?;

    // Benchmark different input sizes
    let test_cases = vec![
        ("Small", 1, 64),
        ("Medium", 4, 256), 
        ("Large", 8, 512),
    ];

    for (name, batch_size, seq_len) in test_cases {
        println!("   ๐Ÿ“Š Testing {} inputs (batch: {}, seq: {})", name, batch_size, seq_len);

        let start_time = std::time::Instant::now();

        // Create test input
        let text_input = Tensor::randint(0, 1000, (batch_size, seq_len), &device)?;
        let image_input = Tensor::randn(0f32, 1f32, (batch_size, 3, 64, 64), &device)?;

        let mut modality_inputs = HashMap::new();
        modality_inputs.insert(Modality::Text, ModalityInput::Text(text_input));
        modality_inputs.insert(Modality::Image, ModalityInput::Image(image_input));

        let multimodal_input = MultiModalInput {
            modality_inputs,
            attention_masks: HashMap::new(),
            batch_size,
        };

        let _output = processor.process(multimodal_input)?;
        let elapsed = start_time.elapsed();

        println!("      โ””โ”€ Processing time: {:.2?}", elapsed);
    }

    Ok(())
}

#[cfg(test)]
mod tests {
    use super::*;

    #[tokio::test]
    async fn test_integration() {
        assert!(main().await.is_ok());
    }

    #[test]
    fn test_modality_types() {
        assert_eq!(Modality::Text.as_str(), "text");
        assert_eq!(Modality::Image.as_str(), "image");
        assert_eq!(Modality::Audio.as_str(), "audio");
        assert_eq!(Modality::Video.as_str(), "video");
    }

    #[test]
    fn test_modality_embedding_dims() {
        assert_eq!(Modality::Text.default_embedding_dim(), Some(768));
        assert_eq!(Modality::Image.default_embedding_dim(), Some(2048));
        assert_eq!(Modality::Audio.default_embedding_dim(), Some(512));
        assert_eq!(Modality::Video.default_embedding_dim(), Some(1024));
    }
}