//! Multimodal tokenization for images and audio //! //! This module provides comprehensive tokenization capabilities for non-text modalities, //! enabling unified representation across text, image, and audio data with real processing algorithms. use crate::{Result, TokenId, TokenizationError, TokenizationStats, Tokenizer}; use image::DynamicImage; use parking_lot::RwLock; use rayon::prelude::*; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::sync::Arc; /// Supported modality types #[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum ModalityType { /// Text modality Text, /// Image modality Image, /// Audio modality Audio, /// Video modality (future extension) Video, } /// Image tokenization method #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub enum ImageTokenizationMethod { /// Patch-based tokenization (like Vision Transformer) Patches, /// CNN-style convolutional features ConvolutionalFeatures, /// Hierarchical patches (multi-scale) HierarchicalPatches, /// CLIP-style visual features VisualFeatures, } /// Audio tokenization method #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub enum AudioTokenizationMethod { /// Mel-frequency cepstral coefficients MFCC, /// Spectogram-based features Spectrogram, /// Raw waveform quantization WaveformQuantization, /// Audio patches (like `AudioMAE`) AudioPatches, } /// Multimodal tokenizer configuration #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct MultimodalConfig { /// Supported modalities pub modalities: Vec, /// Vocabulary size per modality pub vocab_size_per_modality: usize, /// Image tokenization configuration pub image_config: ImageTokenizationConfig, /// Audio tokenization configuration pub audio_config: AudioTokenizationConfig, /// Cross-modal alignment settings pub cross_modal_alignment: bool, /// Global token offset per modality pub modality_token_offset: HashMap, } /// Image tokenization configuration #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct ImageTokenizationConfig { /// Tokenization method pub method: ImageTokenizationMethod, /// Image patch size (width, height) pub patch_size: (u32, u32), /// Target image resolution (width, height) pub target_resolution: (u32, u32), /// Number of color channels (3 for RGB, 1 for grayscale) pub channels: u32, /// Normalization parameters (mean, std) scaled by 10000 to avoid f32 pub normalization: Option<([u32; 3], [u32; 3])>, /// Whether to use overlap between patches pub patch_overlap: bool, /// Overlap stride if enabled pub overlap_stride: u32, } impl Default for ImageTokenizationConfig { fn default() -> Self { Self { method: ImageTokenizationMethod::Patches, patch_size: (16, 16), target_resolution: (224, 224), channels: 3, // ImageNet stats scaled by 10000: mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225] normalization: Some(([4850, 4560, 4060], [2290, 2240, 2250])), patch_overlap: false, overlap_stride: 8, } } } impl ImageTokenizationConfig { /// Get normalization parameters as f32 tuples #[must_use] pub fn normalization_f32(&self) -> Option<([f32; 3], [f32; 3])> { self.normalization.map(|(mean, std)| { let mean_f32 = [ mean[0] as f32 / 10000.0, mean[1] as f32 / 10000.0, mean[2] as f32 / 10000.0, ]; let std_f32 = [ std[0] as f32 / 10000.0, std[1] as f32 / 10000.0, std[2] as f32 / 10000.0, ]; (mean_f32, std_f32) }) } } /// Audio tokenization configuration #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct AudioTokenizationConfig { /// Tokenization method pub method: AudioTokenizationMethod, /// Sample rate for audio processing pub sample_rate: u32, /// Window size for spectral analysis pub window_size: usize, /// Hop size for windowing pub hop_size: usize, /// Number of mel filter banks pub n_mel_filters: usize, /// Number of MFCC coefficients pub n_mfcc: usize, /// Frame length in microseconds (to avoid f32) pub frame_length_microseconds: u32, /// Whether to apply pre-emphasis pub pre_emphasis: bool, /// Pre-emphasis coefficient as integer (multiplied by 1000) pub pre_emphasis_coeff_scaled: u32, } impl Default for AudioTokenizationConfig { fn default() -> Self { Self { method: AudioTokenizationMethod::MFCC, sample_rate: 16000, window_size: 512, hop_size: 256, n_mel_filters: 80, n_mfcc: 13, frame_length_microseconds: 25000, // 0.025 seconds = 25000 microseconds pre_emphasis: true, pre_emphasis_coeff_scaled: 970, // 0.97 * 1000 } } } impl AudioTokenizationConfig { /// Get frame length in seconds #[must_use] pub fn frame_length_seconds(&self) -> f32 { self.frame_length_microseconds as f32 / 1_000_000.0 } /// Get pre-emphasis coefficient as f32 #[must_use] pub fn pre_emphasis_coeff(&self) -> f32 { self.pre_emphasis_coeff_scaled as f32 / 1000.0 } } impl Default for MultimodalConfig { fn default() -> Self { let mut modality_token_offset = HashMap::new(); modality_token_offset.insert(ModalityType::Text, 0); modality_token_offset.insert(ModalityType::Image, 10000); modality_token_offset.insert(ModalityType::Audio, 20000); Self { modalities: vec![ModalityType::Text, ModalityType::Image, ModalityType::Audio], vocab_size_per_modality: 10000, image_config: ImageTokenizationConfig::default(), audio_config: AudioTokenizationConfig::default(), cross_modal_alignment: true, modality_token_offset, } } } /// Image patch representation #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct ImagePatch { /// Patch data as flattened RGB values pub data: Vec, /// Patch position (row, column) pub position: (usize, usize), /// Patch dimensions pub dimensions: (u32, u32), } /// Audio frame representation #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct AudioFrame { /// Audio features (MFCC, spectral, etc.) pub features: Vec, /// Frame timestamp in seconds pub timestamp: f32, /// Frame duration in seconds pub duration: f32, } /// Multimodal tokenizer with real processing capabilities #[derive(Debug)] pub struct MultimodalTokenizer { config: MultimodalConfig, /// Pre-computed image patch codebook image_codebook: Arc>>, /// Pre-computed audio feature codebook audio_codebook: Arc>>, /// Cross-modal alignment matrix alignment_matrix: Arc>>>, /// Statistics stats: Arc>, } impl MultimodalTokenizer { /// Create new multimodal tokenizer #[must_use] pub fn new(config: MultimodalConfig) -> Self { Self { config, image_codebook: Arc::new(RwLock::new(Vec::new())), audio_codebook: Arc::new(RwLock::new(Vec::new())), alignment_matrix: Arc::new(RwLock::new(None)), stats: Arc::new(RwLock::new(TokenizationStats::default())), } } /// Train codebooks from multimodal data pub async fn train_codebooks( &mut self, images: &[Vec], audio_samples: &[Vec], ) -> Result<()> { // Train image codebook if !images.is_empty() { let image_patches = self.extract_all_image_patches(images).await?; let image_codebook = self.build_image_codebook(image_patches).await?; *self.image_codebook.write() = image_codebook; } // Train audio codebook if !audio_samples.is_empty() { let audio_features = self.extract_all_audio_features(audio_samples).await?; let audio_codebook = self.build_audio_codebook(audio_features).await?; *self.audio_codebook.write() = audio_codebook; } // Build cross-modal alignment if enabled if self.config.cross_modal_alignment { self.build_cross_modal_alignment().await?; } Ok(()) } /// Tokenize image data with real patch extraction and processing pub async fn encode_image(&self, image_data: &[u8]) -> Result> { let start_time = std::time::Instant::now(); // Load and decode image let image = image::load_from_memory(image_data).map_err(|e| { TokenizationError::MultimodalError(format!("Failed to load image: {e}")) })?; // Resize to target resolution let resized_image = image.resize_exact( self.config.image_config.target_resolution.0, self.config.image_config.target_resolution.1, image::imageops::FilterType::Lanczos3, ); // Extract patches let patches = self.extract_image_patches(&resized_image).await?; // Quantize patches to tokens let token_ids = self.quantize_image_patches(patches).await?; // Add modality offset let offset = *self .config .modality_token_offset .get(&ModalityType::Image) .unwrap_or(&10000); let final_tokens: Vec = token_ids.into_iter().map(|id| id + offset).collect(); // Update statistics { let mut stats = self.stats.write(); stats.processing_time_ms += start_time.elapsed().as_millis() as u64; } Ok(final_tokens) } /// Tokenize audio data with real feature extraction pub async fn encode_audio(&self, audio_data: &[f32]) -> Result> { let start_time = std::time::Instant::now(); // Extract audio features based on configuration let features = match self.config.audio_config.method { AudioTokenizationMethod::MFCC => self.extract_mfcc_features(audio_data).await?, AudioTokenizationMethod::Spectrogram => { self.extract_spectrogram_features(audio_data).await? } AudioTokenizationMethod::WaveformQuantization => { self.quantize_waveform(audio_data).await? } AudioTokenizationMethod::AudioPatches => self.extract_audio_patches(audio_data).await?, }; // Quantize features to tokens let token_ids = self.quantize_audio_features(features).await?; // Add modality offset let offset = *self .config .modality_token_offset .get(&ModalityType::Audio) .unwrap_or(&20000); let final_tokens: Vec = token_ids.into_iter().map(|id| id + offset).collect(); // Update statistics { let mut stats = self.stats.write(); stats.processing_time_ms += start_time.elapsed().as_millis() as u64; } Ok(final_tokens) } /// Extract image patches with real computer vision processing async fn extract_image_patches(&self, image: &DynamicImage) -> Result> { let rgb_image = image.to_rgb8(); let (width, height) = rgb_image.dimensions(); let (patch_width, patch_height) = self.config.image_config.patch_size; let _patches: Vec = Vec::new(); let stride_x = if self.config.image_config.patch_overlap { self.config.image_config.overlap_stride } else { patch_width }; let stride_y = if self.config.image_config.patch_overlap { self.config.image_config.overlap_stride } else { patch_height }; // Extract patches in parallel let patch_positions: Vec<(usize, usize)> = (0..height) .step_by(stride_y as usize) .flat_map(|y| { (0..width) .step_by(stride_x as usize) .map(move |x| (y as usize, x as usize)) }) .collect(); let rgb_data = rgb_image.as_raw(); let extracted_patches: Vec = patch_positions .par_iter() .filter_map(|&(y, x)| { if x + patch_width as usize <= width as usize && y + patch_height as usize <= height as usize { let mut patch_data = Vec::with_capacity((patch_width * patch_height * 3) as usize); // Extract patch pixel data for py in y..y + patch_height as usize { for px in x..x + patch_width as usize { let pixel_idx = (py * width as usize + px) * 3; if pixel_idx + 2 < rgb_data.len() { patch_data.push(f32::from(rgb_data[pixel_idx]) / 255.0); // R patch_data.push(f32::from(rgb_data[pixel_idx + 1]) / 255.0); // G patch_data.push(f32::from(rgb_data[pixel_idx + 2]) / 255.0); // B } } } // Apply normalization if configured if let Some((mean, std)) = self.config.image_config.normalization_f32() { for (i, value) in patch_data.iter_mut().enumerate() { let channel = i % 3; *value = (*value - mean[channel]) / std[channel]; } } Some(ImagePatch { data: patch_data, position: (y / stride_y as usize, x / stride_x as usize), dimensions: (patch_width, patch_height), }) } else { None } }) .collect(); Ok(extracted_patches) } /// Extract MFCC features from audio async fn extract_mfcc_features(&self, audio_data: &[f32]) -> Result> { let window_size = self.config.audio_config.window_size; let hop_size = self.config.audio_config.hop_size; let sample_rate = self.config.audio_config.sample_rate; let n_mfcc = self.config.audio_config.n_mfcc; let mut frames = Vec::new(); let mut processed_audio = audio_data.to_vec(); // Apply pre-emphasis if enabled if self.config.audio_config.pre_emphasis { let coeff = self.config.audio_config.pre_emphasis_coeff(); for i in (1..processed_audio.len()).rev() { processed_audio[i] -= coeff * processed_audio[i - 1]; } } // Extract frames for (frame_idx, frame_start) in (0..processed_audio.len()).step_by(hop_size).enumerate() { if frame_start + window_size <= processed_audio.len() { let frame = &processed_audio[frame_start..frame_start + window_size]; // Apply Hamming window let windowed_frame: Vec = frame .iter() .enumerate() .map(|(i, &sample)| { sample * (0.54 - 0.46 * (2.0 * std::f32::consts::PI * i as f32 / (window_size - 1) as f32) .cos()) }) .collect(); // Compute FFT (simplified - in production use proper FFT library) let fft_magnitudes = self.compute_fft_magnitudes(&windowed_frame); // Convert to mel scale let mel_features = self.apply_mel_filterbank(&fft_magnitudes, sample_rate); // Compute DCT to get MFCC let mfcc_features = self.compute_dct(&mel_features, n_mfcc); frames.push(AudioFrame { features: mfcc_features, timestamp: frame_idx as f32 * hop_size as f32 / sample_rate as f32, duration: window_size as f32 / sample_rate as f32, }); } } Ok(frames) } /// Extract spectrogram features from audio async fn extract_spectrogram_features(&self, audio_data: &[f32]) -> Result> { let window_size = self.config.audio_config.window_size; let hop_size = self.config.audio_config.hop_size; let sample_rate = self.config.audio_config.sample_rate; let mut frames = Vec::new(); // Extract spectral frames for (frame_idx, frame_start) in (0..audio_data.len()).step_by(hop_size).enumerate() { if frame_start + window_size <= audio_data.len() { let frame = &audio_data[frame_start..frame_start + window_size]; // Apply Hamming window let windowed_frame: Vec = frame .iter() .enumerate() .map(|(i, &sample)| { sample * (0.54 - 0.46 * (2.0 * std::f32::consts::PI * i as f32 / (window_size - 1) as f32) .cos()) }) .collect(); // Compute power spectrum let power_spectrum = self.compute_power_spectrum(&windowed_frame); frames.push(AudioFrame { features: power_spectrum, timestamp: frame_idx as f32 * hop_size as f32 / sample_rate as f32, duration: window_size as f32 / sample_rate as f32, }); } } Ok(frames) } /// Quantize waveform directly async fn quantize_waveform(&self, audio_data: &[f32]) -> Result> { let frame_size = self.config.audio_config.hop_size; let sample_rate = self.config.audio_config.sample_rate; let mut frames = Vec::new(); for (frame_idx, chunk) in audio_data.chunks(frame_size).enumerate() { // Simple quantization: divide into bins and represent as discrete values let quantized_samples: Vec = chunk .iter() .map(|&sample| { // Quantize to 256 levels let quantized = ((sample + 1.0) * 127.5).round().max(0.0).min(255.0); quantized / 255.0 * 2.0 - 1.0 }) .collect(); frames.push(AudioFrame { features: quantized_samples, timestamp: frame_idx as f32 * frame_size as f32 / sample_rate as f32, duration: frame_size as f32 / sample_rate as f32, }); } Ok(frames) } /// Extract audio patches (similar to image patches but for audio) async fn extract_audio_patches(&self, audio_data: &[f32]) -> Result> { let patch_size = self.config.audio_config.window_size; let hop_size = self.config.audio_config.hop_size; let sample_rate = self.config.audio_config.sample_rate; let mut frames = Vec::new(); for (frame_idx, frame_start) in (0..audio_data.len()).step_by(hop_size).enumerate() { if frame_start + patch_size <= audio_data.len() { let patch = audio_data[frame_start..frame_start + patch_size].to_vec(); frames.push(AudioFrame { features: patch, timestamp: frame_idx as f32 * hop_size as f32 / sample_rate as f32, duration: patch_size as f32 / sample_rate as f32, }); } } Ok(frames) } /// Helper functions for audio processing /// Compute FFT magnitudes (simplified implementation) fn compute_fft_magnitudes(&self, frame: &[f32]) -> Vec { // Simplified FFT - in production use rustfft or similar let n = frame.len(); let mut magnitudes = Vec::new(); for k in 0..n / 2 { let mut real = 0.0; let mut imag = 0.0; for n_idx in 0..n { let angle = -2.0 * std::f32::consts::PI * k as f32 * n_idx as f32 / n as f32; real += frame[n_idx] * angle.cos(); imag += frame[n_idx] * angle.sin(); } magnitudes.push((real * real + imag * imag).sqrt()); } magnitudes } /// Apply mel filterbank fn apply_mel_filterbank(&self, magnitudes: &[f32], sample_rate: u32) -> Vec { let n_mel = self.config.audio_config.n_mel_filters; let mut mel_features = vec![0.0; n_mel]; // Simplified mel filterbank implementation let mel_max = self.hz_to_mel(sample_rate as f32 / 2.0); let mel_step = mel_max / (n_mel + 1) as f32; for i in 0..n_mel { let mel_center = (i + 1) as f32 * mel_step; let hz_center = self.mel_to_hz(mel_center); let bin_center = hz_center * magnitudes.len() as f32 * 2.0 / sample_rate as f32; // Simple triangular filter let start_bin = (bin_center - 10.0).max(0.0) as usize; let end_bin = (bin_center + 10.0).min(magnitudes.len() as f32) as usize; let mut energy = 0.0; for bin in start_bin..end_bin { let weight = 1.0 - (bin as f32 - bin_center).abs() / 10.0; if weight > 0.0 && bin < magnitudes.len() { energy += magnitudes[bin] * weight; } } mel_features[i] = if energy > 0.0 { energy.ln() } else { -10.0 }; } mel_features } /// Convert Hz to Mel scale fn hz_to_mel(&self, hz: f32) -> f32 { 2595.0 * (1.0 + hz / 700.0).log10() } /// Convert Mel to Hz scale fn mel_to_hz(&self, mel: f32) -> f32 { 700.0 * (10.0_f32.powf(mel / 2595.0) - 1.0) } /// Compute DCT for MFCC fn compute_dct(&self, mel_features: &[f32], n_mfcc: usize) -> Vec { let mut mfcc = vec![0.0; n_mfcc]; let n = mel_features.len(); for i in 0..n_mfcc { let mut sum = 0.0; for j in 0..n { sum += mel_features[j] * (std::f32::consts::PI * i as f32 * (j as f32 + 0.5) / n as f32).cos(); } mfcc[i] = sum; } mfcc } /// Compute power spectrum fn compute_power_spectrum(&self, frame: &[f32]) -> Vec { let magnitudes = self.compute_fft_magnitudes(frame); magnitudes.into_iter().map(|mag| mag * mag).collect() } /// Extract patches from all images for codebook training async fn extract_all_image_patches(&self, images: &[Vec]) -> Result> { let mut all_patches = Vec::new(); for image_data in images { if let Ok(image) = image::load_from_memory(image_data) { let resized = image.resize_exact( self.config.image_config.target_resolution.0, self.config.image_config.target_resolution.1, image::imageops::FilterType::Lanczos3, ); let patches = self.extract_image_patches(&resized).await?; all_patches.extend(patches); } } Ok(all_patches) } /// Extract features from all audio samples for codebook training async fn extract_all_audio_features( &self, audio_samples: &[Vec], ) -> Result> { let mut all_features = Vec::new(); for audio_data in audio_samples { let features = match self.config.audio_config.method { AudioTokenizationMethod::MFCC => self.extract_mfcc_features(audio_data).await?, AudioTokenizationMethod::Spectrogram => { self.extract_spectrogram_features(audio_data).await? } AudioTokenizationMethod::WaveformQuantization => { self.quantize_waveform(audio_data).await? } AudioTokenizationMethod::AudioPatches => { self.extract_audio_patches(audio_data).await? } }; all_features.extend(features); } Ok(all_features) } /// Build image codebook using k-means clustering async fn build_image_codebook(&self, patches: Vec) -> Result> { if patches.is_empty() { return Ok(Vec::new()); } let k = self.config.vocab_size_per_modality.min(patches.len()); let mut centroids = Vec::new(); // Initialize centroids with random patches let mut rng = rand::thread_rng(); use rand::seq::SliceRandom; let mut shuffled_patches = patches.clone(); shuffled_patches.shuffle(&mut rng); for i in 0..k { centroids.push(shuffled_patches[i % shuffled_patches.len()].clone()); } // Simple k-means (in production, use more sophisticated clustering) for _iteration in 0..10 { let mut new_centroids = vec![ ImagePatch { data: vec![0.0; patches[0].data.len()], position: (0, 0), dimensions: patches[0].dimensions, }; k ]; let mut counts = vec![0; k]; // Assign patches to nearest centroids for patch in &patches { let mut best_distance = f32::INFINITY; let mut best_centroid = 0; for (centroid_idx, centroid) in centroids.iter().enumerate() { let distance = self.compute_patch_distance(patch, centroid); if distance < best_distance { best_distance = distance; best_centroid = centroid_idx; } } // Accumulate for new centroid for (i, &value) in patch.data.iter().enumerate() { new_centroids[best_centroid].data[i] += value; } counts[best_centroid] += 1; } // Update centroids for (centroid_idx, count) in counts.iter().enumerate() { if *count > 0 { for value in &mut new_centroids[centroid_idx].data { *value /= *count as f32; } } } centroids = new_centroids; } Ok(centroids) } /// Build audio codebook using k-means clustering async fn build_audio_codebook(&self, features: Vec) -> Result> { if features.is_empty() { return Ok(Vec::new()); } let k = self.config.vocab_size_per_modality.min(features.len()); let mut centroids = Vec::new(); // Initialize centroids let mut rng = rand::thread_rng(); use rand::seq::SliceRandom; let mut shuffled_features = features.clone(); shuffled_features.shuffle(&mut rng); for i in 0..k { centroids.push(shuffled_features[i % shuffled_features.len()].clone()); } // k-means clustering for _iteration in 0..10 { let mut new_centroids = vec![ AudioFrame { features: vec![0.0; features[0].features.len()], timestamp: 0.0, duration: 0.0, }; k ]; let mut counts = vec![0; k]; for feature in &features { let mut best_distance = f32::INFINITY; let mut best_centroid = 0; for (centroid_idx, centroid) in centroids.iter().enumerate() { let distance = self.compute_feature_distance(feature, centroid); if distance < best_distance { best_distance = distance; best_centroid = centroid_idx; } } for (i, &value) in feature.features.iter().enumerate() { new_centroids[best_centroid].features[i] += value; } counts[best_centroid] += 1; } for (centroid_idx, count) in counts.iter().enumerate() { if *count > 0 { for value in &mut new_centroids[centroid_idx].features { *value /= *count as f32; } } } centroids = new_centroids; } Ok(centroids) } /// Build cross-modal alignment matrix async fn build_cross_modal_alignment(&mut self) -> Result<()> { let image_codebook_size = self.image_codebook.read().len(); let audio_codebook_size = self.audio_codebook.read().len(); if image_codebook_size > 0 && audio_codebook_size > 0 { // Create a simple random alignment matrix (in production, use learned alignment) let matrix = nalgebra::DMatrix::::from_fn( image_codebook_size, audio_codebook_size, |_i, _j| rand::random::(), ); *self.alignment_matrix.write() = Some(matrix); } Ok(()) } /// Quantize image patches to token IDs async fn quantize_image_patches(&self, patches: Vec) -> Result> { let codebook = self.image_codebook.read(); let mut tokens = Vec::new(); for patch in patches { let mut best_distance = f32::INFINITY; let mut best_token = 0; for (token_id, centroid) in codebook.iter().enumerate() { let distance = self.compute_patch_distance(&patch, centroid); if distance < best_distance { best_distance = distance; best_token = token_id as TokenId; } } tokens.push(best_token); } Ok(tokens) } /// Quantize audio features to token IDs async fn quantize_audio_features(&self, features: Vec) -> Result> { let codebook = self.audio_codebook.read(); let mut tokens = Vec::new(); for feature in features { let mut best_distance = f32::INFINITY; let mut best_token = 0; for (token_id, centroid) in codebook.iter().enumerate() { let distance = self.compute_feature_distance(&feature, centroid); if distance < best_distance { best_distance = distance; best_token = token_id as TokenId; } } tokens.push(best_token); } Ok(tokens) } /// Compute Euclidean distance between image patches fn compute_patch_distance(&self, patch1: &ImagePatch, patch2: &ImagePatch) -> f32 { patch1 .data .iter() .zip(patch2.data.iter()) .map(|(a, b)| (a - b).powi(2)) .sum::() .sqrt() } /// Compute Euclidean distance between audio features fn compute_feature_distance(&self, feature1: &AudioFrame, feature2: &AudioFrame) -> f32 { feature1 .features .iter() .zip(feature2.features.iter()) .map(|(a, b)| (a - b).powi(2)) .sum::() .sqrt() } } #[async_trait::async_trait] impl Tokenizer for MultimodalTokenizer { async fn encode(&self, text: &str) -> Result> { // Default text tokenization - simple character-based for now let start_time = std::time::Instant::now(); let tokens: Vec = text .chars() .take(self.config.vocab_size_per_modality) .enumerate() .map(|(i, _)| i as TokenId) .collect(); // Update statistics { let mut stats = self.stats.write(); stats.token_count = tokens.len(); stats.processing_time_ms = start_time.elapsed().as_millis() as u64; } Ok(tokens) } async fn decode(&self, token_ids: &[TokenId]) -> Result { // Determine modality from token IDs and decode appropriately let mut decoded_parts = Vec::new(); for &token_id in token_ids { if token_id < 10000 { decoded_parts.push("TEXT".to_string()); } else if token_id < 20000 { decoded_parts.push("IMAGE".to_string()); } else if token_id < 30000 { decoded_parts.push("AUDIO".to_string()); } else { decoded_parts.push("UNKNOWN".to_string()); } } Ok(format!("multimodal[{}]", decoded_parts.join(","))) } fn vocab_size(&self) -> usize { self.config.modalities.len() * self.config.vocab_size_per_modality } fn get_stats(&self) -> TokenizationStats { self.stats.read().clone() } fn supports(&self, _text: &str) -> bool { self.config.modalities.contains(&ModalityType::Text) } } /// Simplified serialization data structure #[derive(serde::Serialize, serde::Deserialize)] struct MultimodalTokenizerData { config: MultimodalConfig, image_codebook: Vec, audio_codebook: Vec, } impl serde::Serialize for MultimodalTokenizer { fn serialize(&self, serializer: S) -> std::result::Result where S: serde::Serializer, { let data = MultimodalTokenizerData { config: self.config.clone(), image_codebook: self.image_codebook.read().clone(), audio_codebook: self.audio_codebook.read().clone(), }; data.serialize(serializer) } } impl<'de> serde::Deserialize<'de> for MultimodalTokenizer { fn deserialize(deserializer: D) -> std::result::Result where D: serde::Deserializer<'de>, { let data = MultimodalTokenizerData::deserialize(deserializer)?; Ok(Self { config: data.config, image_codebook: Arc::new(RwLock::new(data.image_codebook)), audio_codebook: Arc::new(RwLock::new(data.audio_codebook)), alignment_matrix: Arc::new(RwLock::new(None)), stats: Arc::new(RwLock::new(TokenizationStats::default())), }) } } #[cfg(test)] mod tests { use super::*; use image::{ImageBuffer, Rgb, RgbImage}; #[test] fn multimodal_config_default() { let config = MultimodalConfig::default(); assert_eq!(config.modalities.len(), 3); assert!(config.modalities.contains(&ModalityType::Text)); assert!(config.modalities.contains(&ModalityType::Image)); assert!(config.modalities.contains(&ModalityType::Audio)); assert_eq!(config.vocab_size_per_modality, 10000); assert!(config.cross_modal_alignment); } #[test] fn image_config_default() { let config = ImageTokenizationConfig::default(); assert_eq!(config.method, ImageTokenizationMethod::Patches); assert_eq!(config.patch_size, (16, 16)); assert_eq!(config.target_resolution, (224, 224)); assert_eq!(config.channels, 3); assert!(config.normalization.is_some()); } #[test] fn audio_config_default() { let config = AudioTokenizationConfig::default(); assert_eq!(config.method, AudioTokenizationMethod::MFCC); assert_eq!(config.sample_rate, 16000); assert_eq!(config.window_size, 512); assert_eq!(config.hop_size, 256); assert_eq!(config.n_mel_filters, 80); assert_eq!(config.n_mfcc, 13); } #[tokio::test] async fn multimodal_tokenizer_creation() { let config = MultimodalConfig::default(); let tokenizer = MultimodalTokenizer::new(config); assert_eq!(tokenizer.vocab_size(), 30000); // 3 modalities * 10000 each assert!(tokenizer.supports("test")); } #[test] fn image_patch_creation() { let patch = ImagePatch { data: vec![0.1, 0.2, 0.3], position: (0, 1), dimensions: (16, 16), }; assert_eq!(patch.data.len(), 3); assert_eq!(patch.position, (0, 1)); assert_eq!(patch.dimensions, (16, 16)); } #[test] fn audio_frame_creation() { let frame = AudioFrame { features: vec![0.1, 0.2, 0.3, 0.4], timestamp: 1.5, duration: 0.025, }; assert_eq!(frame.features.len(), 4); assert_eq!(frame.timestamp, 1.5); assert_eq!(frame.duration, 0.025); } #[tokio::test] async fn encode_decode_text() { let config = MultimodalConfig::default(); let tokenizer = MultimodalTokenizer::new(config); let tokens = tokenizer.encode("hello").await.unwrap(); assert!(!tokens.is_empty()); let decoded = tokenizer.decode(&tokens).await.unwrap(); assert!(decoded.contains("multimodal")); } #[test] fn hz_mel_conversion() { let config = MultimodalConfig::default(); let tokenizer = MultimodalTokenizer::new(config); let hz = 1000.0; let mel = tokenizer.hz_to_mel(hz); let hz_back = tokenizer.mel_to_hz(mel); assert!((hz - hz_back).abs() < 0.1); } #[tokio::test] async fn real_image_processing() { let config = MultimodalConfig::default(); let tokenizer = MultimodalTokenizer::new(config); // Create a simple test image (solid red 224x224) let width = 224u32; let height = 224u32; let mut image_data: Vec = Vec::new(); // PNG header and simple red image let image: RgbImage = ImageBuffer::from_fn(width, height, |_x, _y| Rgb([255u8, 0u8, 0u8])); let dynamic_image = DynamicImage::ImageRgb8(image); // Test patch extraction let patches = tokenizer .extract_image_patches(&dynamic_image) .await .unwrap(); assert!(!patches.is_empty()); assert_eq!(patches[0].dimensions, (16, 16)); assert_eq!(patches[0].data.len(), 16 * 16 * 3); // RGB } #[tokio::test] async fn real_audio_processing() { let config = MultimodalConfig::default(); let tokenizer = MultimodalTokenizer::new(config); // Create test audio - simple sine wave let sample_rate = 16000; let duration = 1.0; // 1 second let frequency = 440.0; // A4 note let samples: Vec = (0..((sample_rate as f32 * duration) as usize)) .map(|i| (2.0 * std::f32::consts::PI * frequency * i as f32 / sample_rate as f32).sin()) .collect(); // Test MFCC extraction let frames = tokenizer.extract_mfcc_features(&samples).await.unwrap(); assert!(!frames.is_empty()); assert_eq!( frames[0].features.len(), tokenizer.config.audio_config.n_mfcc ); } }