//! Multimodal fusion transformer implementation. use crate::Result; use rtx_tensor::{Device, Tensor}; use rtx_transformers::architectures::TransformerConfig; use rtx_transformers::layers::Layer; use rtx_transformers::prelude::{LayerNorm, TransformerBlock}; use super::config::MultimodalFusionConfig; pub struct MultimodalFusionTransformer { blocks: Vec, ln_f: LayerNorm, config: MultimodalFusionConfig, device: Device, } impl MultimodalFusionTransformer { pub fn new(config: &MultimodalFusionConfig, device: &Device) -> Result { let transformer_config = TransformerConfig { vocab_size: 32000, // Default vocab size d_model: config.embed_dim, num_layers: 1, // We'll create the layers manually num_heads: config.num_heads, num_key_value_heads: Some(config.num_heads), // Same as num_heads for MHA use_mqa: false, // Multi-head attention, not MQA d_ff: config.embed_dim * 4, // Standard 4x expansion for FFN max_seq_len: 4096, // Default max sequence length dropout: config.dropout as f64, layer_norm_eps: 1e-5, bias: true, activation: "gelu".to_string(), }; let mut blocks = Vec::new(); for _ in 0..config.num_layers { blocks.push(TransformerBlock::new(transformer_config.clone(), device)?); } let ln_f = LayerNorm::new(config.embed_dim, 1e-5, true, device)?; Ok(Self { blocks, ln_f, config: config.clone(), device: device.clone(), }) } pub fn forward(&self, vision: &Tensor, audio: &Tensor, text: &Tensor) -> Result { // Concatenate all modalities along sequence dimension let fused = Tensor::cat(&[vision.clone(), audio.clone(), text.clone()], 1)?; // Pass through transformer blocks let mut x = fused; for block in &self.blocks { x = block.forward(&x)?; } // Final layer norm let x = self.ln_f.forward(&x)?; Ok(x) } }