A compact multimodal embedding model that unifies text, image, and audio representations in a shared semantic space. Part of the MoM (Mixture of Models) family.
Model Description
multi-modal-embed-small is a lightweight multimodal encoder (~120M parameters) supporting:
Text encoding via MiniLM-L6-v2 (22M params)
Image encoding via SigLIP-base-patch16-512 (86M params)
Audio encoding via Whisper-tiny encoder (8M params)
Cross-modal fusion via 2-layer transformer attention
2DMSE: Two-Dimensional Matryoshka Sentence Embeddings for adaptive compute
MRL: Matryoshka Representation Learning for flexible embedding dimensions
Key Features
Feature
Description
Embedding Dimension
384 (supports MRL truncation to 32, 64, 128, 256)
Image Resolution
512×512
Audio Input
Up to 30s, 16kHz (Whisper Mel spectrogram)
Modalities
Text, Image, Audio, Multimodal fusion
2DMSE Support
Early exit at any encoder layer
Languages
English
Installation
pip install torch transformers pillow safetensors
Usage
Load Model
Two checkpoint formats are available:
model.pt (932 MB) - PyTorch format
model.safetensors (1.35 GB) - SafeTensors format
python
1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4from transformers import AutoModel, AutoTokenizer, SiglipModel, SiglipProcessor, WhisperModel, WhisperFeatureExtractor
5from huggingface_hub import hf_hub_download
67classMultiModalEmbedder(nn.Module):8"""Standalone multimodal embedder - no external dependencies."""910def__init__(self):11super().__init__()12# Text encoder (384d, no projection needed)13 self.text_tokenizer = AutoTokenizer.from_pretrained("sentence-transformers/all-MiniLM-L6-v2")14 self.text_encoder = AutoModel.from_pretrained("sentence-transformers/all-MiniLM-L6-v2")1516# Image encoder (768d -> 384d projection)17 self.image_processor = SiglipProcessor.from_pretrained("google/siglip-base-patch16-512")18 self.image_encoder = SiglipModel.from_pretrained("google/siglip-base-patch16-512").vision_model
19 self.image_proj = nn.Linear(768,384)2021# Audio encoder (384d, no projection needed)22 self.audio_processor = WhisperFeatureExtractor.from_pretrained("openai/whisper-tiny")23 self.audio_encoder = WhisperModel.from_pretrained("openai/whisper-tiny").encoder
2425defencode_text(self, texts):26ifisinstance(texts,str):27 texts =[texts]28 inputs = self.text_tokenizer(texts, padding=True, truncation=True, return_tensors="pt")29 inputs ={k: v.to(next(self.parameters()).device)for k, v in inputs.items()}30 outputs = self.text_encoder(**inputs)31 embeddings = outputs.last_hidden_state.mean(dim=1)# Mean pooling32return F.normalize(embeddings, p=2, dim=-1)3334defencode_image(self, images):35 inputs = self.image_processor(images=images, return_tensors="pt")36 inputs ={k: v.to(next(self.parameters()).device)for k, v in inputs.items()}37 outputs = self.image_encoder(**inputs)38 embeddings = outputs.pooler_output
39 embeddings = self.image_proj(embeddings)# 768 -> 38440return F.normalize(embeddings, p=2, dim=-1)4142defencode_audio(self, waveform):43# waveform: numpy array or tensor at 16kHz44ifisinstance(waveform, torch.Tensor):45 waveform = waveform.squeeze().numpy()46 inputs = self.audio_processor(waveform, sampling_rate=16000, return_tensors="pt")47 inputs ={k: v.to(next(self.parameters()).device)for k, v in inputs.items()}48 outputs = self.audio_encoder(**inputs)49 embeddings = outputs.last_hidden_state.mean(dim=1)# Mean pooling50return F.normalize(embeddings, p=2, dim=-1)5152# Load model53model = MultiModalEmbedder()5455# Download and load trained weights56checkpoint_path = hf_hub_download(57 repo_id="llm-semantic-router/multi-modal-embed-small",58 filename="model.pt"59)60state_dict = torch.load(checkpoint_path, map_location="cpu", weights_only=False)6162# Load text encoder weights63model.text_encoder.load_state_dict({64 k.replace("text_encoder.encoder.",""): v
65for k, v in state_dict.items()66if k.startswith("text_encoder.encoder.")67})6869# Load image encoder and projection weights70model.image_encoder.load_state_dict({71 k.replace("image_encoder.vision_encoder.",""): v
72for k, v in state_dict.items()73if k.startswith("image_encoder.vision_encoder.")74})75model.image_proj.load_state_dict({76 k.replace("image_encoder.projection.",""): v
77for k, v in state_dict.items()78if k.startswith("image_encoder.projection.")79})8081# Load audio encoder weights82model.audio_encoder.load_state_dict({83 k.replace("audio_encoder.encoder.",""): v
84for k, v in state_dict.items()85if k.startswith("audio_encoder.encoder.")86})8788model.eval()89print("Model loaded successfully!")
Text Embedding
python
1import torch.nn.functional as F
23# Single text4text_embedding = model.encode_text("A photo of a cat")# Shape: [1, 384]56# Batch of texts7texts =["A fluffy orange cat","A golden retriever dog","A red sports car"]8text_embeddings = model.encode_text(texts)# Shape: [3, 384]910# Compute similarity11similarities = F.cosine_similarity(text_embeddings[0:1], text_embeddings[1:], dim=-1)12print(f"Cat vs Dog: {similarities[0]:.3f}")13print(f"Cat vs Car: {similarities[1]:.3f}")
1# Image-to-text retrieval2image = Image.open("cat.jpg").convert('RGB')3image_emb = model.encode_image(image)45captions =[6"A cat sleeping on a bed",7"A dog playing in the park",8"A car driving on the highway",9]10text_embs = model.encode_text(captions)1112similarities = F.cosine_similarity(image_emb, text_embs)13best_idx = similarities.argmax().item()14print(f"Best match: {captions[best_idx]} ({similarities[best_idx]:.3f})")