Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 58 additions & 29 deletions crates/multimodal/src/vision/processors/llama4_vision.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@

use std::collections::HashSet;

use image::{imageops::FilterType, DynamicImage, GenericImageView, Rgb, RgbImage};
use image::{imageops::FilterType, DynamicImage, GenericImageView};
use ndarray::{s, Array3, Array4};

use crate::vision::{
Expand Down Expand Up @@ -258,28 +258,55 @@ impl Llama4VisionProcessor {
}
}

/// Pad image to target dimensions with black padding.
#[expect(
clippy::unused_self,
reason = "method logically belongs to the processor; keeps API consistent"
)]
fn pad_image(&self, image: &DynamicImage, target_w: u32, target_h: u32) -> DynamicImage {
let (w, h) = image.dimensions();
if w == target_w && h == target_h {
return image.clone();
}

// Create black background (LLaMA 4 uses 0 for padding)
let black = Rgb([0u8, 0, 0]);
let mut padded = RgbImage::from_pixel(target_w, target_h, black);

// Copy image to top-left — avoid to_rgb8() if already RGB8
match image {
DynamicImage::ImageRgb8(rgb) => image::imageops::overlay(&mut padded, rgb, 0, 0),
_ => image::imageops::overlay(&mut padded, &image.to_rgb8(), 0, 0),
/// Build a padded [C, H, W] f32 tensor from a smaller image.
///
/// The image is placed at top-left, and the remaining canvas is filled with
/// the normalized value of black (0). This fuses pad + tensor conversion
/// into one step, avoiding an intermediate padded `RgbImage` allocation.
fn pad_and_normalize_to_tensor(
&self,
image: &DynamicImage,
canvas_w: usize,
canvas_h: usize,
) -> Array3<f32> {
let (img_w, img_h, raw) = transforms::rgb_bytes(image);
let canvas_pixels = canvas_h * canvas_w;

// Precompute fused scale/bias: (pixel/255 - mean) / std
let scale: [f32; 3] = std::array::from_fn(|c| 1.0 / (255.0 * self.std[c] as f32));
let bias: [f32; 3] = std::array::from_fn(|c| -(self.mean[c] as f32) / (self.std[c] as f32));
Comment thread
CatherineSue marked this conversation as resolved.

let mut data = vec![0.0f32; 3 * canvas_pixels];
let (r_plane, rest) = data.split_at_mut(canvas_pixels);
let (g_plane, b_plane) = rest.split_at_mut(canvas_pixels);

// Pre-fill with normalized black: 0 * scale + bias = bias
r_plane.fill(bias[0]);
g_plane.fill(bias[1]);
b_plane.fill(bias[2]);

// Overwrite image region row-by-row using the shared block-optimized helper
let rw = img_w.min(canvas_w);
let rh = img_h.min(canvas_h);
for y in 0..rh {
let src_row = &raw[y * img_w * 3..y * img_w * 3 + rw * 3];
let dst_offset = y * canvas_w;
transforms::deinterleave_rgb_to_planes(
src_row,
&mut r_plane[dst_offset..dst_offset + rw],
&mut g_plane[dst_offset..dst_offset + rw],
&mut b_plane[dst_offset..dst_offset + rw],
scale,
bias,
);
}

DynamicImage::ImageRgb8(padded)
#[expect(
clippy::expect_used,
reason = "data has exactly 3*canvas_h*canvas_w elements by construction"
)]
Array3::from_shape_vec((3, canvas_h, canvas_w), data)
.expect("shape matches pre-allocated buffer")
}

/// Split image tensor into tiles.
Expand Down Expand Up @@ -326,7 +353,6 @@ impl Llama4VisionProcessor {
let (target_h, target_w) = target_size;

// Step 2: Compute resize target - limit upscaling if not resize_to_max_canvas
// This limits how much we resize the image, but we still pad to target_size
let resize_target = if self.resize_to_max_canvas {
target_size
} else {
Expand All @@ -342,22 +368,23 @@ impl Llama4VisionProcessor {

let resized = transforms::resize(image, new_w, new_h, FilterType::Triangle);

// Step 4: Pad to target_size (the canvas from get_best_fit, not resize_target)
let padded = self.pad_image(&resized, target_w, target_h);

// Step 5: Convert to tensor and normalize
let tensor = transforms::to_tensor_and_normalize(&padded, &self.mean, &self.std);
// Fused pad + tensor: build the padded f32 tensor directly from the
// resized RGB bytes, avoiding an intermediate padded RgbImage allocation.
let tensor = if new_w != target_w || new_h != target_h {
self.pad_and_normalize_to_tensor(&resized, target_w as usize, target_h as usize)
} else {
transforms::to_tensor_and_normalize(&resized, &self.mean, &self.std)
};

// Step 6: Calculate tile counts based on target_size (canvas size)
let tile = self.tile_size as usize;
let num_tiles_h = target_h as usize / tile;
let num_tiles_w = target_w as usize / tile;

// Step 7: Split into tiles
// Step 7: Split into tiles + global tile
let tiles = self.split_to_tiles(&tensor, num_tiles_h, num_tiles_w);
let num_tiles = num_tiles_h * num_tiles_w;

// Step 8: Add global tile if there are multiple tiles
let output = if num_tiles > 1 {
let global_tile = self.create_global_image(image);
let mut combined = Array4::<f32>::zeros((num_tiles + 1, 3, tile, tile));
Expand Down Expand Up @@ -508,6 +535,8 @@ impl ImagePreProcessor for Llama4VisionProcessor {

#[cfg(test)]
mod tests {
use image::{Rgb, RgbImage};

use super::*;

fn create_test_image(width: u32, height: u32, color: Rgb<u8>) -> DynamicImage {
Expand Down
55 changes: 49 additions & 6 deletions crates/multimodal/src/vision/transforms.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@ pub type Result<T> = std::result::Result<T, TransformError>;

/// Extract RGB pixel data from a DynamicImage, avoiding a copy when already RGB8.
/// Returns (width, height, raw_bytes) where raw_bytes is interleaved R,G,B,R,G,B,...
fn rgb_bytes(image: &DynamicImage) -> (usize, usize, std::borrow::Cow<'_, [u8]>) {
pub fn rgb_bytes(image: &DynamicImage) -> (usize, usize, std::borrow::Cow<'_, [u8]>) {
match image {
DynamicImage::ImageRgb8(rgb) => (
rgb.width() as usize,
Expand All @@ -54,6 +54,53 @@ fn rgb_bytes(image: &DynamicImage) -> (usize, usize, std::borrow::Cow<'_, [u8]>)
}
}

/// Deinterleave interleaved RGB bytes into separate R, G, B f32 planes with
/// per-channel `scale` and `bias`: `plane[c][i] = rgb[i*3 + c] * scale[c] + bias[c]`.
///
/// Processes 8 pixels at a time so the compiler can unroll and auto-vectorize
/// the stride-3 gather pattern.
pub fn deinterleave_rgb_to_planes(
rgb: &[u8],
r_plane: &mut [f32],
g_plane: &mut [f32],
b_plane: &mut [f32],
scale: [f32; 3],
bias: [f32; 3],
) {
let pixels = r_plane.len();
debug_assert_eq!(pixels, g_plane.len());
debug_assert_eq!(pixels, b_plane.len());
debug_assert!(rgb.len() >= pixels * 3);

let full_blocks = pixels / 8;
let remainder = pixels % 8;

for block in 0..full_blocks {
let dst = block * 8;
let src_base = dst * 3;
let src = &rgb[src_base..src_base + 24];
let rd = &mut r_plane[dst..dst + 8];
let gd = &mut g_plane[dst..dst + 8];
let bd = &mut b_plane[dst..dst + 8];

for i in 0..8 {
let s = i * 3;
rd[i] = src[s] as f32 * scale[0] + bias[0];
gd[i] = src[s + 1] as f32 * scale[1] + bias[1];
bd[i] = src[s + 2] as f32 * scale[2] + bias[2];
}
}

let tail_dst = full_blocks * 8;
let tail_src = tail_dst * 3;
for i in 0..remainder {
let s = tail_src + i * 3;
r_plane[tail_dst + i] = rgb[s] as f32 * scale[0] + bias[0];
g_plane[tail_dst + i] = rgb[s + 1] as f32 * scale[1] + bias[1];
b_plane[tail_dst + i] = rgb[s + 2] as f32 * scale[2] + bias[2];
}
}

/// Build a [C, H, W] f32 tensor from interleaved RGB bytes with per-channel
/// `scale` and `bias`: `output[c][i] = raw[i*3 + c] * scale[c] + bias[c]`.
fn build_planar_tensor(
Expand All @@ -68,11 +115,7 @@ fn build_planar_tensor(
let (r_plane, rest) = data.split_at_mut(pixels);
let (g_plane, b_plane) = rest.split_at_mut(pixels);

for (i, chunk) in raw.chunks_exact(3).enumerate() {
r_plane[i] = chunk[0] as f32 * scale[0] + bias[0];
g_plane[i] = chunk[1] as f32 * scale[1] + bias[1];
b_plane[i] = chunk[2] as f32 * scale[2] + bias[2];
}
deinterleave_rgb_to_planes(raw, r_plane, g_plane, b_plane, scale, bias);

#[expect(
clippy::expect_used,
Expand Down
14 changes: 13 additions & 1 deletion model_gateway/src/routers/grpc/multimodal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -409,6 +409,7 @@ async fn process_multimodal_parts(

debug!(
image_count = images.len(),
image_sizes = ?images.iter().map(|f| (f.image.width(), f.image.height())).collect::<Vec<_>>(),
"Fetched images for multimodal processing"
);

Expand Down Expand Up @@ -684,7 +685,18 @@ fn serialize_pixel_values(preprocessed: &PreprocessedImages) -> (Vec<u8>, Vec<u3
.as_slice()
.or_else(|| preprocessed.pixel_values.as_slice_memory_order())
Comment thread
CatherineSue marked this conversation as resolved.
{
pixel_slice.iter().flat_map(|v| v.to_le_bytes()).collect()
// Zero-copy reinterpret: &[f32] → &[u8] on little-endian (x86).
// This replaces the per-element flat_map(to_le_bytes) which was the
// #1 CPU hotspot (13% of SMG CPU in profiling).
#[cfg(target_endian = "little")]
{
let byte_slice: &[u8] = bytemuck::cast_slice(pixel_slice);
byte_slice.to_vec()
}
#[cfg(not(target_endian = "little"))]
{
pixel_slice.iter().flat_map(|v| v.to_le_bytes()).collect()
}
Comment thread
CatherineSue marked this conversation as resolved.
} else {
// Fallback for non-contiguous arrays
preprocessed
Expand Down
Loading