use crate::multimodal::*;
use crate::distributed::{DistributedConfig, ShardingStrategy};
use crate::distributed_core::SimpleDistributedManager;
use crate::cache::ModelCache;
use crate::loader::{LoadOptions, LoadedModel};
use crate::error::{Result, Error as MlmfError};
use candle_core::{Device, DType, Tensor};
use std::collections::HashMap;
use std::path::Path;
use std::sync::Arc;
pub struct MultiModalLoader {
config: MultiModalConfig,
base_options: LoadOptions,
distributed_manager: Option<Arc<SimpleDistributedManager>>,
cache: Option<Arc<ModelCache>>,
modality_paths: HashMap<Modality, String>,
}
impl MultiModalLoader {
pub fn new(config: MultiModalConfig, base_options: LoadOptions) -> Self {
Self {
config,
base_options,
distributed_manager: None,
cache: None,
modality_paths: HashMap::new(),
}
}
pub fn with_distributed(mut self, distributed_config: DistributedConfig) -> Result<Self> {
let mut dist_config = distributed_config;
dist_config.sharding_strategy = ShardingStrategy::ModalitySpecific {
modality_assignments: self.create_modality_assignments(),
};
self.distributed_manager = Some(Arc::new(
SimpleDistributedManager::new(dist_config)?
));
Ok(self)
}
pub fn with_cache(mut self, cache: Arc<ModelCache>) -> Self {
self.cache = Some(cache);
self
}
pub fn with_modality_path<P: AsRef<Path>>(mut self, modality: Modality, path: P) -> Self {
self.modality_paths.insert(modality, path.as_ref().to_string_lossy().to_string());
self
}
pub async fn load(&self) -> Result<MultiModalModel> {
let mut modality_models = HashMap::new();
for (modality, path) in &self.modality_paths {
let model = self.load_modality_model(*modality, path).await?;
modality_models.insert(*modality, model);
}
let cross_modal_layers = self.create_cross_modal_layers(&modality_models)?;
let fusion_components = self.create_fusion_components(&modality_models)?;
Ok(MultiModalModel {
config: self.config.clone(),
modality_models,
cross_modal_layers,
fusion_components,
device: self.base_options.device.clone(),
dtype: self.base_options.dtype,
})
}
async fn load_modality_model(&self, modality: Modality, path: &str) -> Result<LoadedModel> {
if let Some(cache) = &self.cache {
let cache_key = format!("multimodal_{}_{}", modality.as_str(), path);
}
let model = if let Some(dist_manager) = &self.distributed_manager {
self.load_distributed_modality(modality, path, dist_manager).await?
} else {
self.load_single_modality(modality, path).await?
};
if let Some(cache) = &self.cache {
let cache_key = format!("multimodal_{}_{}", modality.as_str(), path);
}
Ok(model)
}
async fn load_distributed_modality(
&self,
modality: Modality,
path: &str,
dist_manager: &SimpleDistributedManager,
) -> Result<LoadedModel> {
dist_manager.deploy_model(path, "multimodal_model".to_string(), ShardingStrategy::NoSharding).await?;
crate::loader::load_safetensors_auto(path)
}
async fn load_single_modality(&self, modality: Modality, path: &str) -> Result<LoadedModel> {
if path.ends_with(".safetensors") {
crate::loader::load_safetensors(path, self.base_options.clone_basic())
} else if path.ends_with(".gguf") {
crate::formats::load_gguf(Path::new(path), &self.base_options)
} else {
crate::loader::load_safetensors_auto(path)
}
}
fn create_modality_assignments(&self) -> HashMap<Modality, Vec<String>> {
let mut assignments = HashMap::new();
for modality in self.config.modalities.keys() {
let nodes = match modality {
Modality::Text => vec!["text-node-1".to_string(), "text-node-2".to_string()],
Modality::Image => vec!["image-node-1".to_string(), "image-node-2".to_string()],
Modality::Audio => vec!["audio-node-1".to_string()],
Modality::Video => vec!["video-node-1".to_string(), "video-node-2".to_string()],
Modality::Custom(id) => vec![format!("custom-node-{}", id)],
};
assignments.insert(*modality, nodes);
}
assignments
}
fn create_cross_modal_layers(&self, _modality_models: &HashMap<Modality, LoadedModel>) -> Result<HashMap<(Modality, Modality), Arc<dyn CrossModalLayer>>> {
let mut layers = HashMap::new();
let modalities: Vec<Modality> = self.config.modalities.keys().cloned().collect();
for &mod1 in &modalities {
for &mod2 in &modalities {
if mod1 != mod2 {
let layer = Arc::new(BasicCrossModalLayer::new(
mod1,
mod2,
&self.config,
&self.base_options.device,
self.base_options.dtype,
)?);
layers.insert((mod1, mod2), layer as Arc<dyn CrossModalLayer>);
}
}
}
Ok(layers)
}
fn create_fusion_components(&self, _modality_models: &HashMap<Modality, LoadedModel>) -> Result<Arc<dyn FusionComponent>> {
let total_dim: usize = self.config.modalities.values()
.map(|config| config.embedding_dim)
.sum();
Ok(Arc::new(BasicFusionComponent::new(
total_dim,
&self.config.fusion_strategy,
&self.base_options.device,
self.base_options.dtype,
)?))
}
}
pub struct MultiModalModel {
pub config: MultiModalConfig,
pub modality_models: HashMap<Modality, LoadedModel>,
pub cross_modal_layers: HashMap<(Modality, Modality), Arc<dyn CrossModalLayer>>,
pub fusion_components: Arc<dyn FusionComponent>,
pub device: Device,
pub dtype: DType,
}
impl MultiModalModel {
pub fn infer(&self, input: MultiModalInput) -> Result<MultiModalOutput> {
let mut modality_outputs = HashMap::new();
for (modality, modality_input) in &input.modality_inputs {
if let Some(model) = self.modality_models.get(modality) {
let output = self.process_modality(*modality, modality_input, model)?;
modality_outputs.insert(*modality, output);
}
}
let mut cross_modal_outputs = HashMap::new();
for ((source_mod, target_mod), layer) in &self.cross_modal_layers {
if let (Some(source_output), Some(target_output)) = (
modality_outputs.get(source_mod),
modality_outputs.get(target_mod)
) {
let interaction_output = layer.process(source_output, target_output)?;
cross_modal_outputs.insert((*source_mod, *target_mod), interaction_output);
}
}
let modality_embeddings: Vec<&Tensor> = modality_outputs.values().collect();
let fused_output = self.fusion_components.fuse(&modality_embeddings)?;
Ok(MultiModalOutput {
fused_embeddings: fused_output,
modality_embeddings: modality_outputs,
attention_weights: self.extract_attention_weights(&cross_modal_outputs),
metadata: HashMap::new(),
})
}
fn process_modality(
&self,
modality: Modality,
input: &ModalityInput,
model: &LoadedModel,
) -> Result<Tensor> {
let preprocessed = self.preprocess_for_modality(modality, input.tensor())?;
Ok(preprocessed)
}
fn preprocess_for_modality(&self, modality: Modality, input: &Tensor) -> Result<Tensor> {
let config = self.config.modalities.get(&modality)
.ok_or_else(|| MlmfError::invalid_config(format!("No config for modality: {:?}", modality)))?;
match &config.preprocessing {
PreprocessingConfig::Text { max_length, .. } => {
let shape = input.shape();
if shape.dims().len() >= 2 && shape.dims()[1] > *max_length {
let indices = Tensor::arange(0, *max_length as i64, &self.device)?;
Ok(input.index_select(&indices, 1)?)
} else {
Ok(input.clone())
}
},
PreprocessingConfig::Image { normalize, .. } => {
if *normalize {
Ok(((input - 0.5)? / 0.5)?)
} else {
Ok(input.clone())
}
},
_ => Ok(input.clone()),
}
}
fn extract_attention_weights(&self, cross_modal_outputs: &HashMap<(Modality, Modality), Tensor>) -> HashMap<(Modality, Modality), Tensor> {
cross_modal_outputs.clone()
}
pub fn stats(&self) -> MultiModalModelStats {
let mut modality_sizes = HashMap::new();
let mut total_parameters = 0;
for (modality, model) in &self.modality_models {
let size = model.raw_tensors.values()
.map(|tensor| tensor.elem_count())
.sum::<usize>();
modality_sizes.insert(*modality, size);
total_parameters += size;
}
MultiModalModelStats {
modality_sizes,
total_parameters,
supported_modalities: self.config.modalities.keys().cloned().collect(),
fusion_strategy: self.config.fusion_strategy.clone(),
distributed: self.config.distributed,
}
}
}
#[derive(Debug, Clone)]
pub struct MultiModalModelStats {
pub modality_sizes: HashMap<Modality, usize>,
pub total_parameters: usize,
pub supported_modalities: Vec<Modality>,
pub fusion_strategy: FusionStrategy,
pub distributed: bool,
}
pub trait CrossModalLayer: Send + Sync {
fn process(&self, source: &Tensor, target: &Tensor) -> Result<Tensor>;
fn source_modality(&self) -> Modality;
fn target_modality(&self) -> Modality;
}
pub struct BasicCrossModalLayer {
source_modality: Modality,
target_modality: Modality,
attention_layer: crate::multimodal_processor::CrossModalAttention,
}
impl BasicCrossModalLayer {
pub fn new(
source_modality: Modality,
target_modality: Modality,
config: &MultiModalConfig,
device: &Device,
dtype: DType,
) -> Result<Self> {
let source_config = config.modalities.get(&source_modality)
.ok_or_else(|| MlmfError::invalid_config(format!("No config for source modality: {:?}", source_modality)))?;
let target_config = config.modalities.get(&target_modality)
.ok_or_else(|| MlmfError::invalid_config(format!("No config for target modality: {:?}", target_modality)))?;
let attention_layer = crate::multimodal_processor::CrossModalAttention::new(
source_config.embedding_dim,
target_config.embedding_dim,
&config.cross_modal_attention,
device,
dtype,
)?;
Ok(Self {
source_modality,
target_modality,
attention_layer,
})
}
}
impl CrossModalLayer for BasicCrossModalLayer {
fn process(&self, source: &Tensor, target: &Tensor) -> Result<Tensor> {
let (attended, _weights) = self.attention_layer.forward(source, target, target)?;
Ok(attended)
}
fn source_modality(&self) -> Modality {
self.source_modality
}
fn target_modality(&self) -> Modality {
self.target_modality
}
}
pub trait FusionComponent: Send + Sync {
fn fuse(&self, embeddings: &[&Tensor]) -> Result<Tensor>;
fn fusion_strategy(&self) -> &FusionStrategy;
}
pub struct BasicFusionComponent {
fusion_layer: crate::multimodal_processor::FusionLayer,
strategy: FusionStrategy,
}
impl BasicFusionComponent {
pub fn new(
input_dim: usize,
strategy: &FusionStrategy,
device: &Device,
dtype: DType,
) -> Result<Self> {
let fusion_layer = crate::multimodal_processor::FusionLayer::new(
input_dim,
strategy,
device,
dtype,
)?;
Ok(Self {
fusion_layer,
strategy: strategy.clone(),
})
}
}
impl FusionComponent for BasicFusionComponent {
fn fuse(&self, embeddings: &[&Tensor]) -> Result<Tensor> {
self.fusion_layer.fuse(embeddings)
}
fn fusion_strategy(&self) -> &FusionStrategy {
&self.strategy
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
#[test]
fn test_multimodal_loader_creation() {
let config = MultiModalConfig::default();
let options = LoadOptions {
device: Device::Cpu,
dtype: DType::F32,
use_mmap: false,
validate_cuda: false,
progress: None,
smart_mapping_oracle: None,
};
let loader = MultiModalLoader::new(config, options);
assert_eq!(loader.modality_paths.len(), 0);
}
}