diff --git a/sgl-router/src/multimodal/mod.rs b/sgl-router/src/multimodal/mod.rs index 76ee5d1c6..172e275f0 100644 --- a/sgl-router/src/multimodal/mod.rs +++ b/sgl-router/src/multimodal/mod.rs @@ -3,6 +3,7 @@ pub mod media; pub mod registry; pub mod tracker; pub mod types; +pub mod vision; pub use error::{MediaConnectorError, MultiModalError, MultiModalResult}; pub use media::{ImageFetchConfig, MediaConnector, MediaConnectorConfig, MediaSource}; diff --git a/sgl-router/src/multimodal/vision/image_processor.rs b/sgl-router/src/multimodal/vision/image_processor.rs new file mode 100644 index 000000000..f8905101f --- /dev/null +++ b/sgl-router/src/multimodal/vision/image_processor.rs @@ -0,0 +1,328 @@ +//! Image processor trait and output types. +//! +//! This module defines the interface for model-specific image processors +//! and the common output format for preprocessed images. + +use std::collections::HashMap; + +use image::DynamicImage; +use ndarray::Array4; + +use super::{preprocessor_config::PreProcessorConfig, transforms::TransformError}; + +/// Model-specific output values that vary by architecture. +/// +/// Different vision models require different auxiliary outputs beyond pixel_values. +/// This enum captures the common types of such outputs. +#[derive(Debug, Clone)] +pub enum ModelSpecificValue { + /// A tensor with shape information (data as flat vec, shape as dims) + Tensor { data: Vec, shape: Vec }, + + /// A tensor of integers (e.g., aspect_ratio_ids) + IntTensor { data: Vec, shape: Vec }, + + /// A tensor of unsigned integers (e.g., image_grid_thw) + UintTensor { data: Vec, shape: Vec }, + + /// Simple integer value + Int(i64), + + /// Simple float value + Float(f64), + + /// List of integers + IntVec(Vec), + + /// List of unsigned integers + UintVec(Vec), + + /// List of floats + FloatVec(Vec), + + /// List of tuples (e.g., image sizes) + TupleVec(Vec<(u32, u32)>), + + /// Boolean flag + Bool(bool), +} + +impl ModelSpecificValue { + /// Create a 1D uint tensor from a vector. + pub fn uint_1d(data: Vec) -> Self { + let len = data.len(); + Self::UintTensor { + data, + shape: vec![len], + } + } + + /// Create a 2D uint tensor. + pub fn uint_2d(data: Vec, rows: usize, cols: usize) -> Self { + Self::UintTensor { + data, + shape: vec![rows, cols], + } + } + + /// Create a 1D int tensor from a vector. + pub fn int_1d(data: Vec) -> Self { + let len = data.len(); + Self::IntTensor { + data, + shape: vec![len], + } + } +} + +/// Preprocessed images ready for model consumption. +/// +/// This struct contains all the outputs needed by the SGLang scheduler +/// to construct `MultimodalInputs` for the model. +#[derive(Debug, Clone)] +pub struct PreprocessedImages { + /// Pixel values as [B, C, H, W] float32 tensor. + /// + /// This is the primary input to the vision encoder. + pub pixel_values: Array4, + + /// Number of image tokens per image in the batch. + /// + /// Used to expand placeholder tokens in the text input. + /// For example, LLaVA with 336x336 and patch_size=14 produces 576 tokens. + pub num_img_tokens: Vec, + + /// Original image sizes as (width, height) before preprocessing. + /// + /// Some models need this for proper attention masking or position encoding. + pub image_sizes: Vec<(u32, u32)>, + + /// Model-specific auxiliary outputs. + /// + /// Examples: + /// - Qwen-VL: `image_grid_thw` for rotary position encoding + /// - LLaMA-Vision: `aspect_ratio_ids`, `aspect_ratio_mask` + /// - Phi3-Vision: `num_img_tokens` per crop + pub model_specific: HashMap, +} + +impl PreprocessedImages { + /// Create a new PreprocessedImages with required fields. + pub fn new( + pixel_values: Array4, + num_img_tokens: Vec, + image_sizes: Vec<(u32, u32)>, + ) -> Self { + Self { + pixel_values, + num_img_tokens, + image_sizes, + model_specific: HashMap::new(), + } + } + + /// Add a model-specific value. + pub fn with_extra(mut self, key: impl Into, value: ModelSpecificValue) -> Self { + self.model_specific.insert(key.into(), value); + self + } + + /// Get the batch size. + pub fn batch_size(&self) -> usize { + self.pixel_values.shape()[0] + } + + /// Get the number of channels. + pub fn channels(&self) -> usize { + self.pixel_values.shape()[1] + } + + /// Get the height of processed images. + pub fn height(&self) -> usize { + self.pixel_values.shape()[2] + } + + /// Get the width of processed images. + pub fn width(&self) -> usize { + self.pixel_values.shape()[3] + } + + /// Get total number of image tokens across all images. + pub fn total_tokens(&self) -> usize { + self.num_img_tokens.iter().sum() + } + + /// Get pixel values as a flat f32 slice (row-major order). + pub fn pixel_values_flat(&self) -> Vec { + self.pixel_values.iter().copied().collect() + } + + /// Get the shape of pixel values as a vector. + pub fn pixel_values_shape(&self) -> Vec { + self.pixel_values.shape().to_vec() + } +} + +/// Trait for model-specific image preprocessors. +/// +/// Each vision model (LLaVA, Qwen-VL, Phi3-Vision, etc.) implements this trait +/// to provide the correct preprocessing pipeline. +pub trait ImagePreProcessor: Send + Sync { + /// Default normalization mean for this model family. + fn default_mean(&self) -> [f64; 3]; + + /// Default normalization std for this model family. + fn default_std(&self) -> [f64; 3]; + + /// Preprocess a batch of images. + /// + /// # Arguments + /// * `images` - Input images to preprocess + /// * `config` - Preprocessor configuration from HuggingFace + /// + /// # Returns + /// Preprocessed images ready for the model, or an error. + fn preprocess( + &self, + images: &[DynamicImage], + config: &PreProcessorConfig, + ) -> Result; + + /// Calculate the number of image tokens for a given image size. + /// + /// This is used to determine how many placeholder tokens to insert + /// in the text input before the image has been fully processed. + /// + /// # Arguments + /// * `width` - Image width after preprocessing + /// * `height` - Image height after preprocessing + /// * `config` - Preprocessor configuration + fn calculate_num_tokens(&self, width: u32, height: u32, config: &PreProcessorConfig) -> usize; + + /// Get the model family name for identification. + fn model_name(&self) -> &'static str; + + /// Get the expected image size after preprocessing. + /// + /// Some models have fixed sizes, others are dynamic. + fn get_processed_size(&self, config: &PreProcessorConfig) -> Option<(u32, u32)> { + config.get_target_size() + } +} + +/// Registry of available image processors. +pub struct ImageProcessorRegistry { + processors: HashMap>, +} + +impl ImageProcessorRegistry { + /// Create a new empty registry. + pub fn new() -> Self { + Self { + processors: HashMap::new(), + } + } + + /// Register a processor for a model pattern. + pub fn register(&mut self, pattern: impl Into, processor: Box) { + self.processors.insert(pattern.into(), processor); + } + + /// Find a processor for the given model ID. + /// + /// Matches by substring containment (case-insensitive). + pub fn find(&self, model_id: &str) -> Option<&dyn ImagePreProcessor> { + let model_lower = model_id.to_lowercase(); + for (pattern, processor) in &self.processors { + if model_lower.contains(&pattern.to_lowercase()) { + return Some(processor.as_ref()); + } + } + None + } + + /// Check if a model has a registered processor. + pub fn has_processor(&self, model_id: &str) -> bool { + self.find(model_id).is_some() + } + + /// Get list of supported model patterns. + pub fn supported_patterns(&self) -> Vec<&str> { + self.processors.keys().map(|s| s.as_str()).collect() + } +} + +impl Default for ImageProcessorRegistry { + fn default() -> Self { + Self::new() + } +} + +#[cfg(test)] +mod tests { + use ndarray::Array4; + + use super::*; + + #[test] + fn test_preprocessed_images_accessors() { + let pixel_values = Array4::::zeros((2, 3, 336, 336)); + let images = + PreprocessedImages::new(pixel_values, vec![576, 576], vec![(640, 480), (800, 600)]); + + assert_eq!(images.batch_size(), 2); + assert_eq!(images.channels(), 3); + assert_eq!(images.height(), 336); + assert_eq!(images.width(), 336); + assert_eq!(images.total_tokens(), 1152); + } + + #[test] + fn test_preprocessed_images_with_extra() { + let pixel_values = Array4::::zeros((1, 3, 224, 224)); + let images = PreprocessedImages::new(pixel_values, vec![196], vec![(224, 224)]) + .with_extra( + "image_grid_thw", + ModelSpecificValue::uint_1d(vec![1, 16, 16]), + ) + .with_extra("aspect_ratio_id", ModelSpecificValue::Int(0)); + + assert!(images.model_specific.contains_key("image_grid_thw")); + assert!(images.model_specific.contains_key("aspect_ratio_id")); + } + + #[test] + fn test_model_specific_value_constructors() { + let uint_1d = ModelSpecificValue::uint_1d(vec![1, 2, 3]); + match uint_1d { + ModelSpecificValue::UintTensor { data, shape } => { + assert_eq!(data, vec![1, 2, 3]); + assert_eq!(shape, vec![3]); + } + _ => panic!("Expected UintTensor"), + } + + let uint_2d = ModelSpecificValue::uint_2d(vec![1, 2, 3, 4], 2, 2); + match uint_2d { + ModelSpecificValue::UintTensor { data, shape } => { + assert_eq!(data, vec![1, 2, 3, 4]); + assert_eq!(shape, vec![2, 2]); + } + _ => panic!("Expected UintTensor"), + } + } + + #[test] + fn test_pixel_values_flat() { + let mut pixel_values = Array4::::zeros((1, 1, 2, 2)); + pixel_values[[0, 0, 0, 0]] = 1.0; + pixel_values[[0, 0, 0, 1]] = 2.0; + pixel_values[[0, 0, 1, 0]] = 3.0; + pixel_values[[0, 0, 1, 1]] = 4.0; + + let images = PreprocessedImages::new(pixel_values, vec![4], vec![(2, 2)]); + let flat = images.pixel_values_flat(); + + assert_eq!(flat, vec![1.0, 2.0, 3.0, 4.0]); + } +} diff --git a/sgl-router/src/multimodal/vision/mod.rs b/sgl-router/src/multimodal/vision/mod.rs new file mode 100644 index 000000000..ce1f98c44 --- /dev/null +++ b/sgl-router/src/multimodal/vision/mod.rs @@ -0,0 +1,39 @@ +//! Pure Rust vision processing module for multimodal models. +//! +//! This module provides image preprocessing pipelines that match HuggingFace processor +//! outputs without requiring Python dependencies. +//! +//! # Architecture +//! +//! The vision module is structured as follows: +//! +//! - `transforms`: Core image transformations (resize, normalize, crop, etc.) +//! - `preprocessor_config`: HuggingFace config parsing +//! - `image_processor`: Trait and output types for processors +//! +//! Model-specific processors will be added in Phase 2. +//! +//! # Usage +//! +//! ```rust,ignore +//! use sgl_model_gateway::multimodal::vision::{ +//! PreProcessorConfig, +//! transforms, +//! }; +//! +//! // Load config from HuggingFace +//! let config = PreProcessorConfig::from_json(config_json)?; +//! +//! // Use transforms directly +//! let tensor = transforms::to_tensor(&image); +//! transforms::normalize(&mut tensor, &mean, &std); +//! ``` + +pub mod image_processor; +pub mod preprocessor_config; +pub mod transforms; + +// Re-export commonly used types +pub use image_processor::{ImagePreProcessor, ModelSpecificValue, PreprocessedImages}; +pub use preprocessor_config::PreProcessorConfig; +pub use transforms::TransformError; diff --git a/sgl-router/src/multimodal/vision/preprocessor_config.rs b/sgl-router/src/multimodal/vision/preprocessor_config.rs new file mode 100644 index 000000000..c422163d2 --- /dev/null +++ b/sgl-router/src/multimodal/vision/preprocessor_config.rs @@ -0,0 +1,367 @@ +//! HuggingFace preprocessor_config.json parsing. +//! +//! This module parses the `preprocessor_config.json` files from HuggingFace model +//! repositories, providing the configuration needed for image preprocessing. + +use std::collections::HashMap; + +use image::imageops::FilterType; +use serde::Deserialize; + +use super::transforms; + +/// HuggingFace preprocessor_config.json structure. +/// +/// This struct captures the common fields across different vision model processors. +/// Model-specific fields are accessed via the flexible `extra` field. +#[derive(Debug, Clone, Deserialize, Default)] +pub struct PreProcessorConfig { + /// Processor class name (e.g., "CLIPImageProcessor", "Qwen2VLImageProcessor") + #[serde(default)] + pub image_processor_type: Option, + + /// Whether to convert to RGB + #[serde(default)] + pub do_convert_rgb: Option, + + /// Whether to normalize with mean/std + #[serde(default)] + pub do_normalize: Option, + + /// Whether to pad images + #[serde(default)] + pub do_pad: Option, + + /// Whether to rescale pixel values (typically by 1/255) + #[serde(default)] + pub do_rescale: Option, + + /// Whether to resize images + #[serde(default)] + pub do_resize: Option, + + /// Whether to center crop after resizing + #[serde(default)] + pub do_center_crop: Option, + + /// Per-channel normalization mean + #[serde(default, alias = "norm_mean")] + pub image_mean: Option>, + + /// Per-channel normalization std + #[serde(default, alias = "norm_std")] + pub image_std: Option>, + + /// Rescale factor (typically 1/255 = 0.00392156862745098) + #[serde(default)] + pub rescale_factor: Option, + + /// PIL resampling filter enum (0=Nearest, 1=Lanczos, 2=Bilinear, 3=Bicubic) + #[serde(default, alias = "resample")] + pub resampling: Option, + + /// Target size for resizing + /// Can be {"height": H, "width": W} or {"shortest_edge": S} + #[serde(default)] + pub size: Option>, + + /// Target size for center cropping + #[serde(default)] + pub crop_size: Option>, + + // ===================== + // Model-specific fields + // ===================== + /// Vision encoder patch size (typically 14 or 16) + #[serde(default)] + pub patch_size: Option, + + /// Qwen-VL: merge size for token reduction + #[serde(default)] + pub merge_size: Option, + + /// Qwen-VL: minimum total pixels + #[serde(default)] + pub min_pixels: Option, + + /// Qwen-VL: maximum total pixels + #[serde(default)] + pub max_pixels: Option, + + /// Qwen-VL: temporal patch size for video + #[serde(default)] + pub temporal_patch_size: Option, + + /// Phi3-Vision: number of image crops + #[serde(default)] + pub num_crops: Option, + + /// LLaMA-Vision: maximum image tiles + #[serde(default)] + pub max_image_tiles: Option, + + /// Fixed number of image tokens (some models use this) + #[serde(default)] + pub num_img_tokens: Option, + + // ===================== + // Special tokens + // ===================== + /// Image start token + #[serde(default)] + pub im_start_token: Option, + + /// Image end token + #[serde(default)] + pub im_end_token: Option, + + /// Slice start token (for multi-crop) + #[serde(default)] + pub slice_start_token: Option, + + /// Slice end token + #[serde(default)] + pub slice_end_token: Option, + + /// Vision start token (alternative naming) + #[serde(default)] + pub vision_start_token: Option, + + /// Vision end token + #[serde(default)] + pub vision_end_token: Option, + + /// Catch-all for model-specific fields not explicitly defined + #[serde(flatten)] + pub extra: HashMap, +} + +impl PreProcessorConfig { + /// Parse from JSON string. + pub fn from_json(json: &str) -> Result { + serde_json::from_str(json) + } + + /// Parse from JSON value. + pub fn from_value(value: serde_json::Value) -> Result { + serde_json::from_value(value) + } + + /// Get image mean as fixed array, with fallback to CLIP defaults. + pub fn get_image_mean(&self) -> [f64; 3] { + self.image_mean + .as_ref() + .and_then(|v| { + if v.len() >= 3 { + Some([v[0], v[1], v[2]]) + } else { + None + } + }) + .unwrap_or(Self::CLIP_MEAN) + } + + /// Get image std as fixed array, with fallback to CLIP defaults. + pub fn get_image_std(&self) -> [f64; 3] { + self.image_std + .as_ref() + .and_then(|v| { + if v.len() >= 3 { + Some([v[0], v[1], v[2]]) + } else { + None + } + }) + .unwrap_or(Self::CLIP_STD) + } + + /// Get target size from various config formats. + /// + /// Handles both `{"height": H, "width": W}` and `{"shortest_edge": S}` formats. + /// Returns (height, width). + pub fn get_target_size(&self) -> Option<(u32, u32)> { + self.size.as_ref().map(|s| { + // Try explicit height/width first + let h = s + .get("height") + .or_else(|| s.get("shortest_edge")) + .copied() + .unwrap_or(224); + let w = s + .get("width") + .or_else(|| s.get("shortest_edge")) + .copied() + .unwrap_or(224); + (h, w) + }) + } + + /// Get crop size. + /// + /// Returns (height, width). + pub fn get_crop_size(&self) -> Option<(u32, u32)> { + self.crop_size.as_ref().map(|s| { + let h = s.get("height").copied().unwrap_or(224); + let w = s.get("width").copied().unwrap_or(224); + (h, w) + }) + } + + /// Get the interpolation filter for resizing. + pub fn get_filter(&self) -> FilterType { + transforms::pil_to_filter(self.resampling) + } + + /// Check if normalization should be applied. + pub fn should_normalize(&self) -> bool { + self.do_normalize.unwrap_or(true) + } + + /// Check if rescaling should be applied. + pub fn should_rescale(&self) -> bool { + self.do_rescale.unwrap_or(false) + } + + /// Check if resizing should be applied. + pub fn should_resize(&self) -> bool { + self.do_resize.unwrap_or(true) + } + + /// Check if center cropping should be applied. + pub fn should_center_crop(&self) -> bool { + self.do_center_crop.unwrap_or(false) + } + + /// Get rescale factor with default. + pub fn get_rescale_factor(&self) -> f64 { + self.rescale_factor.unwrap_or(1.0 / 255.0) + } + + /// Get a typed extra field. + pub fn get_extra(&self, key: &str) -> Option { + self.extra + .get(key) + .and_then(|v| serde_json::from_value(v.clone()).ok()) + } + + // Common default values + pub const CLIP_MEAN: [f64; 3] = [0.48145466, 0.4578275, 0.40821073]; + pub const CLIP_STD: [f64; 3] = [0.26862954, 0.26130258, 0.27577711]; + + pub const IMAGENET_MEAN: [f64; 3] = [0.485, 0.456, 0.406]; + pub const IMAGENET_STD: [f64; 3] = [0.229, 0.224, 0.225]; + + pub const SIGLIP_MEAN: [f64; 3] = [0.5, 0.5, 0.5]; + pub const SIGLIP_STD: [f64; 3] = [0.5, 0.5, 0.5]; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_clip_config() { + let json = r#"{ + "do_center_crop": true, + "do_normalize": true, + "do_resize": true, + "image_mean": [0.48145466, 0.4578275, 0.40821073], + "image_std": [0.26862954, 0.26130258, 0.27577711], + "resample": 3, + "size": {"shortest_edge": 224} + }"#; + + let config = PreProcessorConfig::from_json(json).unwrap(); + + assert!(config.should_normalize()); + assert!(config.should_center_crop()); + assert!(config.should_resize()); + assert_eq!(config.resampling, Some(3)); + + let (h, w) = config.get_target_size().unwrap(); + assert_eq!(h, 224); + assert_eq!(w, 224); + + let mean = config.get_image_mean(); + assert!((mean[0] - 0.48145466).abs() < 1e-6); + } + + #[test] + fn test_parse_qwen_vl_config() { + let json = r#"{ + "do_normalize": true, + "do_rescale": true, + "do_resize": true, + "image_mean": [0.48145466, 0.4578275, 0.40821073], + "image_std": [0.26862954, 0.26130258, 0.27577711], + "min_pixels": 200704, + "max_pixels": 1003520, + "patch_size": 14, + "merge_size": 2, + "temporal_patch_size": 2, + "rescale_factor": 0.00392156862745098 + }"#; + + let config = PreProcessorConfig::from_json(json).unwrap(); + + assert_eq!(config.min_pixels, Some(200704)); + assert_eq!(config.max_pixels, Some(1003520)); + assert_eq!(config.patch_size, Some(14)); + assert_eq!(config.merge_size, Some(2)); + assert!((config.get_rescale_factor() - 1.0 / 255.0).abs() < 1e-10); + } + + #[test] + fn test_parse_size_formats() { + // Height/width format + let json1 = r#"{"size": {"height": 336, "width": 336}}"#; + let config1 = PreProcessorConfig::from_json(json1).unwrap(); + assert_eq!(config1.get_target_size(), Some((336, 336))); + + // Shortest edge format + let json2 = r#"{"size": {"shortest_edge": 224}}"#; + let config2 = PreProcessorConfig::from_json(json2).unwrap(); + assert_eq!(config2.get_target_size(), Some((224, 224))); + } + + #[test] + fn test_defaults() { + let config = PreProcessorConfig::default(); + + // Should use CLIP defaults when not specified + let mean = config.get_image_mean(); + assert!((mean[0] - PreProcessorConfig::CLIP_MEAN[0]).abs() < 1e-6); + + // Default behaviors + assert!(config.should_normalize()); // true by default + assert!(!config.should_rescale()); // false by default + assert!(config.should_resize()); // true by default + assert!(!config.should_center_crop()); // false by default + } + + #[test] + fn test_filter_conversion() { + let json = r#"{"resampling": 3}"#; + let config = PreProcessorConfig::from_json(json).unwrap(); + assert!(matches!(config.get_filter(), FilterType::CatmullRom)); + } + + #[test] + fn test_extra_fields() { + let json = r#"{ + "custom_field": 42, + "nested": {"foo": "bar"} + }"#; + + let config = PreProcessorConfig::from_json(json).unwrap(); + + let custom: Option = config.get_extra("custom_field"); + assert_eq!(custom, Some(42)); + + let nested: Option> = config.get_extra("nested"); + assert_eq!( + nested.as_ref().unwrap().get("foo"), + Some(&"bar".to_string()) + ); + } +} diff --git a/sgl-router/src/multimodal/vision/transforms.rs b/sgl-router/src/multimodal/vision/transforms.rs new file mode 100644 index 000000000..a0cdff2ac --- /dev/null +++ b/sgl-router/src/multimodal/vision/transforms.rs @@ -0,0 +1,395 @@ +//! Image transformation functions for vision preprocessing. +//! +//! This module provides composable transforms that match HuggingFace image processor +//! behavior, enabling pure Rust preprocessing without Python dependencies. + +use image::{imageops::FilterType, DynamicImage, GenericImageView, Rgb, RgbImage}; +use ndarray::{s, Array3, Array4}; +use thiserror::Error; + +/// Errors that can occur during image transformations. +#[derive(Error, Debug)] +pub enum TransformError { + #[error("Invalid tensor shape: expected {expected}, got {actual:?}")] + InvalidShape { + expected: String, + actual: Vec, + }, + + #[error("Image operation failed: {0}")] + ImageError(#[from] image::ImageError), + + #[error("Empty batch: cannot stack zero tensors")] + EmptyBatch, + + #[error("Inconsistent tensor shapes in batch")] + InconsistentShapes, +} + +pub type Result = std::result::Result; + +/// Convert image to tensor [C, H, W] normalized to [0, 1]. +/// +/// This matches the default behavior of `torchvision.transforms.ToTensor()`. +pub fn to_tensor(image: &DynamicImage) -> Array3 { + let rgb = image.to_rgb8(); + let (w, h) = (rgb.width() as usize, rgb.height() as usize); + let mut arr = Array3::::zeros((3, h, w)); + + for (x, y, pixel) in rgb.enumerate_pixels() { + let (x, y) = (x as usize, y as usize); + arr[[0, y, x]] = pixel[0] as f32 / 255.0; + arr[[1, y, x]] = pixel[1] as f32 / 255.0; + arr[[2, y, x]] = pixel[2] as f32 / 255.0; + } + arr +} + +/// Convert image to tensor [C, H, W] without normalization (keeps [0, 255]). +/// +/// Some models expect unnormalized pixel values. +pub fn to_tensor_no_norm(image: &DynamicImage) -> Array3 { + let rgb = image.to_rgb8(); + let (w, h) = (rgb.width() as usize, rgb.height() as usize); + let mut arr = Array3::::zeros((3, h, w)); + + for (x, y, pixel) in rgb.enumerate_pixels() { + let (x, y) = (x as usize, y as usize); + arr[[0, y, x]] = pixel[0] as f32; + arr[[1, y, x]] = pixel[1] as f32; + arr[[2, y, x]] = pixel[2] as f32; + } + arr +} + +/// Normalize tensor per channel: (x - mean) / std. +/// +/// This matches `torchvision.transforms.Normalize(mean, std)`. +/// +/// # Arguments +/// * `tensor` - Input tensor of shape [C, H, W] +/// * `mean` - Per-channel mean values +/// * `std` - Per-channel standard deviation values +pub fn normalize(tensor: &mut Array3, mean: &[f64; 3], std: &[f64; 3]) { + for c in 0..3 { + let mean_c = mean[c] as f32; + let std_c = std[c] as f32; + tensor + .slice_mut(s![c, .., ..]) + .mapv_inplace(|v| (v - mean_c) / std_c); + } +} + +/// Rescale tensor by a constant factor. +/// +/// Used when `do_rescale=True` in HuggingFace configs (typically 1/255). +pub fn rescale(tensor: &mut Array3, factor: f64) { + let factor = factor as f32; + tensor.mapv_inplace(|v| v * factor); +} + +/// Resize image to exact dimensions. +/// +/// # Arguments +/// * `image` - Input image +/// * `width` - Target width +/// * `height` - Target height +/// * `filter` - Interpolation filter (Nearest, Triangle/Bilinear, CatmullRom/Bicubic, Lanczos3) +pub fn resize(image: &DynamicImage, width: u32, height: u32, filter: FilterType) -> DynamicImage { + image.resize_exact(width, height, filter) +} + +/// Resize image preserving aspect ratio, fitting within max dimensions. +pub fn resize_to_fit( + image: &DynamicImage, + max_width: u32, + max_height: u32, + filter: FilterType, +) -> DynamicImage { + image.resize(max_width, max_height, filter) +} + +/// Center crop image to specified dimensions. +/// +/// If the crop size is larger than the image, the image is returned unchanged. +pub fn center_crop(image: &DynamicImage, crop_w: u32, crop_h: u32) -> DynamicImage { + let (w, h) = image.dimensions(); + if crop_w >= w && crop_h >= h { + return image.clone(); + } + let left = (w.saturating_sub(crop_w)) / 2; + let top = (h.saturating_sub(crop_h)) / 2; + let actual_w = crop_w.min(w); + let actual_h = crop_h.min(h); + image.crop_imm(left, top, actual_w, actual_h) +} + +/// Expand image to square by padding with background color. +/// +/// This is used by LLaVA models which expect square inputs. The image is +/// centered and padded with the mean color on the shorter dimension. +pub fn expand_to_square(image: &DynamicImage, background: Rgb) -> DynamicImage { + let (w, h) = image.dimensions(); + match w.cmp(&h) { + std::cmp::Ordering::Equal => image.clone(), + std::cmp::Ordering::Less => { + // Height > Width: pad horizontally + let mut new_image = DynamicImage::from(RgbImage::from_pixel(h, h, background)); + image::imageops::overlay(&mut new_image, image, ((h - w) / 2) as i64, 0); + new_image + } + std::cmp::Ordering::Greater => { + // Width > Height: pad vertically + let mut new_image = DynamicImage::from(RgbImage::from_pixel(w, w, background)); + image::imageops::overlay(&mut new_image, image, 0, ((w - h) / 2) as i64); + new_image + } + } +} + +/// Pad image to specified dimensions with background color. +/// +/// Image is placed at top-left corner. +pub fn pad_to_size( + image: &DynamicImage, + target_w: u32, + target_h: u32, + background: Rgb, +) -> DynamicImage { + let (w, h) = image.dimensions(); + if w >= target_w && h >= target_h { + return image.clone(); + } + let new_w = w.max(target_w); + let new_h = h.max(target_h); + let mut new_image = DynamicImage::from(RgbImage::from_pixel(new_w, new_h, background)); + image::imageops::overlay(&mut new_image, image, 0, 0); + new_image +} + +/// Stack multiple [C, H, W] tensors into [B, C, H, W]. +/// +/// All tensors must have the same shape. +pub fn stack_batch(tensors: &[Array3]) -> Result> { + if tensors.is_empty() { + return Err(TransformError::EmptyBatch); + } + + let shape = tensors[0].shape(); + let (c, h, w) = (shape[0], shape[1], shape[2]); + + // Verify all tensors have the same shape + for tensor in tensors.iter().skip(1) { + if tensor.shape() != shape { + return Err(TransformError::InvalidShape { + expected: format!("[{}, {}, {}]", c, h, w), + actual: tensor.shape().to_vec(), + }); + } + } + + let mut batch = Array4::::zeros((tensors.len(), c, h, w)); + for (i, tensor) in tensors.iter().enumerate() { + batch.slice_mut(s![i, .., .., ..]).assign(tensor); + } + + Ok(batch) +} + +/// Convert PIL/HuggingFace resampling enum to image crate filter. +/// +/// PIL resampling constants: +/// - 0: NEAREST +/// - 1: LANCZOS (also ANTIALIAS) +/// - 2: BILINEAR +/// - 3: BICUBIC +/// - 4: BOX +/// - 5: HAMMING +pub fn pil_to_filter(resampling: Option) -> FilterType { + match resampling { + Some(0) => FilterType::Nearest, + Some(1) => FilterType::Lanczos3, + Some(2) | None => FilterType::Triangle, // Bilinear (default) + Some(3) => FilterType::CatmullRom, // Bicubic + // Box and Hamming don't have direct equivalents, use Triangle + Some(4) | Some(5) => FilterType::Triangle, + _ => FilterType::Triangle, + } +} + +/// Calculate mean color of an image as RGB. +pub fn calculate_mean_color(image: &DynamicImage) -> Rgb { + let rgb = image.to_rgb8(); + let (w, h) = (rgb.width() as u64, rgb.height() as u64); + let total_pixels = w * h; + + if total_pixels == 0 { + return Rgb([128, 128, 128]); + } + + let (mut r_sum, mut g_sum, mut b_sum) = (0u64, 0u64, 0u64); + for pixel in rgb.pixels() { + r_sum += pixel[0] as u64; + g_sum += pixel[1] as u64; + b_sum += pixel[2] as u64; + } + + Rgb([ + (r_sum / total_pixels) as u8, + (g_sum / total_pixels) as u8, + (b_sum / total_pixels) as u8, + ]) +} + +/// Convert normalized mean values [0, 1] to RGB bytes. +pub fn mean_to_rgb(mean: &[f64; 3]) -> Rgb { + Rgb([ + (mean[0] * 255.0).round() as u8, + (mean[1] * 255.0).round() as u8, + (mean[2] * 255.0).round() as u8, + ]) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn create_test_image(width: u32, height: u32, color: Rgb) -> DynamicImage { + DynamicImage::from(RgbImage::from_pixel(width, height, color)) + } + + #[test] + fn test_to_tensor_shape() { + let img = create_test_image(10, 20, Rgb([255, 128, 0])); + let tensor = to_tensor(&img); + assert_eq!(tensor.shape(), &[3, 20, 10]); // [C, H, W] + } + + #[test] + fn test_to_tensor_values() { + let img = create_test_image(2, 2, Rgb([255, 128, 0])); + let tensor = to_tensor(&img); + + // Check normalization to [0, 1] + assert!((tensor[[0, 0, 0]] - 1.0).abs() < 1e-6); // R=255 -> 1.0 + assert!((tensor[[1, 0, 0]] - 0.502).abs() < 0.01); // G=128 -> ~0.5 + assert!((tensor[[2, 0, 0]] - 0.0).abs() < 1e-6); // B=0 -> 0.0 + } + + #[test] + fn test_to_tensor_no_norm() { + let img = create_test_image(2, 2, Rgb([255, 128, 64])); + let tensor = to_tensor_no_norm(&img); + + assert!((tensor[[0, 0, 0]] - 255.0).abs() < 1e-6); + assert!((tensor[[1, 0, 0]] - 128.0).abs() < 1e-6); + assert!((tensor[[2, 0, 0]] - 64.0).abs() < 1e-6); + } + + #[test] + fn test_normalize() { + let mut tensor = Array3::::from_elem((3, 2, 2), 0.5); + let mean = [0.5, 0.5, 0.5]; + let std = [0.5, 0.5, 0.5]; + + normalize(&mut tensor, &mean, &std); + + // (0.5 - 0.5) / 0.5 = 0.0 + for val in tensor.iter() { + assert!(val.abs() < 1e-6); + } + } + + #[test] + fn test_rescale() { + let mut tensor = Array3::::from_elem((3, 2, 2), 255.0); + rescale(&mut tensor, 1.0 / 255.0); + + for val in tensor.iter() { + assert!((val - 1.0).abs() < 1e-6); + } + } + + #[test] + fn test_resize() { + let img = create_test_image(100, 50, Rgb([128, 128, 128])); + let resized = resize(&img, 50, 25, FilterType::Triangle); + + assert_eq!(resized.width(), 50); + assert_eq!(resized.height(), 25); + } + + #[test] + fn test_center_crop() { + let img = create_test_image(100, 100, Rgb([128, 128, 128])); + let cropped = center_crop(&img, 50, 50); + + assert_eq!(cropped.width(), 50); + assert_eq!(cropped.height(), 50); + } + + #[test] + fn test_expand_to_square_horizontal() { + let img = create_test_image(100, 50, Rgb([255, 0, 0])); + let background = Rgb([0, 0, 0]); + let squared = expand_to_square(&img, background); + + assert_eq!(squared.width(), 100); + assert_eq!(squared.height(), 100); + } + + #[test] + fn test_expand_to_square_vertical() { + let img = create_test_image(50, 100, Rgb([255, 0, 0])); + let background = Rgb([0, 0, 0]); + let squared = expand_to_square(&img, background); + + assert_eq!(squared.width(), 100); + assert_eq!(squared.height(), 100); + } + + #[test] + fn test_expand_to_square_already_square() { + let img = create_test_image(100, 100, Rgb([255, 0, 0])); + let background = Rgb([0, 0, 0]); + let squared = expand_to_square(&img, background); + + assert_eq!(squared.width(), 100); + assert_eq!(squared.height(), 100); + } + + #[test] + fn test_stack_batch() { + let t1 = Array3::::zeros((3, 10, 10)); + let t2 = Array3::::ones((3, 10, 10)); + + let batch = stack_batch(&[t1, t2]).unwrap(); + + assert_eq!(batch.shape(), &[2, 3, 10, 10]); + } + + #[test] + fn test_stack_batch_empty() { + let result = stack_batch(&[]); + assert!(matches!(result, Err(TransformError::EmptyBatch))); + } + + #[test] + fn test_pil_to_filter() { + assert!(matches!(pil_to_filter(Some(0)), FilterType::Nearest)); + assert!(matches!(pil_to_filter(Some(1)), FilterType::Lanczos3)); + assert!(matches!(pil_to_filter(Some(2)), FilterType::Triangle)); + assert!(matches!(pil_to_filter(Some(3)), FilterType::CatmullRom)); + assert!(matches!(pil_to_filter(None), FilterType::Triangle)); + } + + #[test] + fn test_mean_to_rgb() { + let mean = [0.5, 0.25, 1.0]; + let rgb = mean_to_rgb(&mean); + + assert_eq!(rgb[0], 128); + assert_eq!(rgb[1], 64); + assert_eq!(rgb[2], 255); + } +}