//! Configuration types for multimodal fusion. use super::strategies::ModalityFusionStrategy; use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CrossModalAttentionConfig { pub embed_dim: usize, pub num_heads: usize, pub dropout: f32, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MultimodalFusionConfig { pub embed_dim: usize, pub num_heads: usize, pub num_layers: usize, pub dropout: f32, pub max_vision_tokens: usize, pub max_audio_tokens: usize, pub max_text_tokens: usize, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct TensorFusionConfig { pub vision_dim: usize, pub audio_dim: usize, pub text_dim: usize, pub output_dim: usize, pub hidden_dim: usize, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct BottleneckFusionConfig { pub input_dims: Vec, pub bottleneck_dim: usize, pub output_dim: usize, pub num_layers: usize, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AttentionFusionConfig { pub modality_dims: Vec, pub hidden_dim: usize, pub num_heads: usize, pub dropout: f32, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ContrastiveLearningConfig { pub embed_dim: usize, pub temperature: f32, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SequentialFusionConfig { pub input_dims: Vec, pub hidden_dim: usize, pub output_dim: usize, pub dropout: f32, } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct HierarchicalFusionConfig { pub modality_dims: Vec, pub fusion_stages: Vec<(Vec, usize)>, // (modality indices, output dim) } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct GatedFusionConfig { pub modality_dims: Vec, pub hidden_dim: usize, pub output_dim: usize, } /// Revolutionary Multimodal Fusion Configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct FusionConfig { /// Vision modality dimension pub vision_dim: usize, /// Audio modality dimension pub audio_dim: usize, /// Text modality dimension pub text_dim: usize, /// Output unified dimension pub output_dim: usize, /// Number of attention heads for cross-modal fusion pub num_heads: usize, /// Number of fusion layers pub num_fusion_layers: usize, /// Dropout probability pub dropout: f32, /// Enable Flash Attention pub use_flash_attention: bool, /// Fusion strategy pub fusion_strategy: ModalityFusionStrategy, /// Cross-modal alignment loss weight pub alignment_loss_weight: f32, /// Temperature for contrastive learning pub contrastive_temperature: f32, } impl FusionConfig { /// Create a new fusion configuration pub fn new( vision_dim: usize, audio_dim: usize, text_dim: usize, output_dim: usize, num_heads: usize, ) -> Self { Self { vision_dim, audio_dim, text_dim, output_dim, num_heads, num_fusion_layers: 6, dropout: 0.1, use_flash_attention: true, fusion_strategy: ModalityFusionStrategy::HierarchicalAttention, alignment_loss_weight: 0.1, contrastive_temperature: 0.07, } } /// Configuration optimized for inference pub fn for_inference( vision_dim: usize, audio_dim: usize, text_dim: usize, output_dim: usize, num_heads: usize, ) -> Self { let mut config = Self::new(vision_dim, audio_dim, text_dim, output_dim, num_heads); config.dropout = 0.0; config.use_flash_attention = true; config } } impl Default for FusionConfig { fn default() -> Self { Self::new(768, 768, 768, 768, 12) } }