147 lines
3.8 KiB
Rust
147 lines
3.8 KiB
Rust
//! 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<usize>,
|
|
pub bottleneck_dim: usize,
|
|
pub output_dim: usize,
|
|
pub num_layers: usize,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct AttentionFusionConfig {
|
|
pub modality_dims: Vec<usize>,
|
|
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<usize>,
|
|
pub hidden_dim: usize,
|
|
pub output_dim: usize,
|
|
pub dropout: f32,
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct HierarchicalFusionConfig {
|
|
pub modality_dims: Vec<usize>,
|
|
pub fusion_stages: Vec<(Vec<usize>, usize)>, // (modality indices, output dim)
|
|
}
|
|
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct GatedFusionConfig {
|
|
pub modality_dims: Vec<usize>,
|
|
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)
|
|
}
|
|
}
|