Initial commit
This commit is contained in:
@@ -0,0 +1,146 @@
|
||||
//! 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user