memra-engine 0.86.0

From-scratch CUDA LLM inference engine for NVIDIA RTX 50-series (sm_120a) and Hopper (sm_90a) - custom kernels, no frameworks
Documentation
//! Host preprocessor for vision input (lane/vision): bytes -> ViT patch rows.
//!
//! Qwen2VLImageProcessorFast semantics: smart_resize to multiples of
//! factor = patch(16) * merge(2) = 32 with the pixel-area budget, rescale 1/255,
//! normalize mean/std 0.5 -> [-1, 1], patchify to [gh*gw, 3*2*16*16 = 1536] rows in
//! row-major grid order with (c, t, ph, pw) inner order — the flatten of the conv
//! weight [1152, 3, 2, 16, 16], so `VisionTower::forward` consumes rows directly.
//! Images duplicate their frame across temporal_patch 2; videos fill the pair with
//! consecutive sampled frames.
//!
//! Resize filter: CatmullRom (Keys bicubic a=-0.5, PIL-BICUBIC family). The HF fast
//! processor runs torch bicubic (a=-0.75) antialias — close but not bit-equal; the
//! merger-cosine parity gate arbitrates whether the difference matters.

use crate::vision::{V_MERGE, V_PATCH, V_PATCH_IN, V_TEMPORAL};
use base64::Engine as _;
use image::RgbImage;
use image::imageops::FilterType;

/// Area budget (pixels) from preprocessor_config: shortest_edge / longest_edge.
pub const MIN_PIXELS: usize = 65536;
pub const MAX_PIXELS: usize = 16_777_216;
const FACTOR: usize = V_PATCH * V_MERGE; // 32

pub struct PreppedImage {
    /// [gh*gw, 1536] row-major grid order.
    pub patches: Vec<f32>,
    pub gh: usize,
    pub gw: usize,
}

impl PreppedImage {
    /// Trunk tokens this image occupies (after 2x2 merge).
    pub fn n_tokens(&self) -> usize {
        self.gh * self.gw / (V_MERGE * V_MERGE)
    }
}

/// smart_resize (HF): round each side to a multiple of 32 preserving aspect ratio,
/// then scale into the [MIN_PIXELS, MAX_PIXELS] area budget.
pub fn smart_resize(h: usize, w: usize) -> Result<(usize, usize), String> {
    if h < 2 || w < 2 {
        return Err(format!("image too small: {w}x{h}"));
    }
    let ar = h.max(w) as f64 / h.min(w) as f64;
    if ar > 200.0 {
        return Err(format!("aspect ratio {ar:.0} exceeds 200"));
    }
    let f = FACTOR as f64;
    let (hf, wf) = (h as f64, w as f64);
    let mut h_bar = ((hf / f).round() * f).max(f);
    let mut w_bar = ((wf / f).round() * f).max(f);
    if h_bar * w_bar > MAX_PIXELS as f64 {
        let beta = (hf * wf / MAX_PIXELS as f64).sqrt();
        h_bar = ((hf / beta / f).floor() * f).max(f);
        w_bar = ((wf / beta / f).floor() * f).max(f);
    } else if h_bar * w_bar < MIN_PIXELS as f64 {
        let beta = (MIN_PIXELS as f64 / (hf * wf)).sqrt();
        h_bar = (hf * beta / f).ceil() * f;
        w_bar = (wf * beta / f).ceil() * f;
    }
    Ok((h_bar as usize, w_bar as usize))
}

/// Decode + resize one image to its target grid. Returns the resized RGB frame and
/// the patch grid (gh, gw) in 16px patches (both even — factor 32 guarantees it).
fn decode_frame(bytes: &[u8]) -> Result<(RgbImage, usize, usize), String> {
    let img = image::load_from_memory(bytes).map_err(|e| format!("image decode: {e}"))?;
    let rgb = img.to_rgb8();
    let (w, h) = (rgb.width() as usize, rgb.height() as usize);
    let (rh, rw) = smart_resize(h, w)?;
    let resized = image::imageops::resize(&rgb, rw as u32, rh as u32, FilterType::CatmullRom);
    Ok((resized, rh / V_PATCH, rw / V_PATCH))
}

/// Fill patch rows for one temporal slot `t` from a frame. Rows are row-major over
/// the (gh, gw) grid; inner order (c, t, ph, pw).
fn fill_slot(rows: &mut [f32], frame: &RgbImage, gh: usize, gw: usize, t: usize) {
    let inv = 1.0f32 / 127.5;
    for py in 0..gh {
        for px in 0..gw {
            let row = &mut rows[(py * gw + px) * V_PATCH_IN..(py * gw + px + 1) * V_PATCH_IN];
            for c in 0..3 {
                let base = c * V_TEMPORAL * V_PATCH * V_PATCH + t * V_PATCH * V_PATCH;
                for ph in 0..V_PATCH {
                    for pw in 0..V_PATCH {
                        let p =
                            frame.get_pixel((px * V_PATCH + pw) as u32, (py * V_PATCH + ph) as u32);
                        row[base + ph * V_PATCH + pw] = p.0[c] as f32 * inv - 1.0;
                    }
                }
            }
        }
    }
}

/// Image bytes (png/jpeg/webp/gif/bmp) -> patch rows. The single frame fills both
/// temporal slots (HF: images are tiled to temporal_patch_size).
pub fn prep_image_bytes(bytes: &[u8]) -> Result<PreppedImage, String> {
    let (frame, gh, gw) = decode_frame(bytes)?;
    let mut patches = vec![0f32; gh * gw * V_PATCH_IN];
    for t in 0..V_TEMPORAL {
        fill_slot(&mut patches, &frame, gh, gw, t);
    }
    Ok(PreppedImage { patches, gh, gw })
}

/// `data:image/...;base64,<payload>` -> patch rows.
pub fn prep_data_uri(uri: &str) -> Result<PreppedImage, String> {
    let bytes = decode_data_uri(uri)?;
    prep_image_bytes(&bytes)
}

/// One pad-run unit crossing the API boundary: a standalone image, or one temporal
/// group of a video. Units with the same `video` index are consecutive and forward
/// TOGETHER through `forward_seq` (one attention span per video).
pub struct VisionUnit {
    pub prep: PreppedImage,
    /// Some(video_idx) for video groups; None for standalone images.
    pub video: Option<usize>,
}

/// One prepared VIDEO: temporal groups as PreppedImage units (each = one pad run of
/// `gh*gw/4` tokens) + per-group timestamps for the HF placeholder format
/// (`<t.t seconds>` before each group's pad run). Groups forward TOGETHER through
/// `VisionTower::forward_seq` — one attention span per video, the HF cu_seqlens law.
pub struct PreppedVideo {
    pub groups: Vec<PreppedImage>,
    pub timestamps: Vec<f32>,
}

/// Serving cap on total video patches (groups*gh*gw): sdpa_naive keys the whole span in
/// shared memory, so the pixel budget stays well under the HF default. Env-tunable.
pub fn video_max_pixels() -> usize {
    std::env::var("MEMRA_VIDEO_MAX_PIXELS")
        .ok()
        .and_then(|v| v.parse().ok())
        .unwrap_or(2_097_152)
}
pub const VID_MIN_PIXELS: usize = 4096;
/// Sampled frame cap (2 frames per temporal group).
pub const VID_MAX_FRAMES: usize = 32;

/// HF Qwen3VL video smart_resize: the pixel budget covers t_bar*h*w — ALL frames.
fn smart_resize_video(frames: usize, h: usize, w: usize) -> Result<(usize, usize), String> {
    if h < 2 || w < 2 {
        return Err(format!("frame too small: {w}x{h}"));
    }
    let ar = h.max(w) as f64 / h.min(w) as f64;
    if ar > 200.0 {
        return Err(format!("aspect ratio {ar:.0} exceeds 200"));
    }
    let f = FACTOR as f64;
    let (hf, wf) = (h as f64, w as f64);
    let t_bar = ((frames as f64 / V_TEMPORAL as f64).round() * V_TEMPORAL as f64).max(2.0);
    let mut h_bar = ((hf / f).round() * f).max(f);
    let mut w_bar = ((wf / f).round() * f).max(f);
    let (min_px, max_px) = (VID_MIN_PIXELS as f64, video_max_pixels() as f64);
    if t_bar * h_bar * w_bar > max_px {
        let beta = (frames as f64 * hf * wf / max_px).sqrt();
        h_bar = ((hf / beta / f).floor() * f).max(f);
        w_bar = ((wf / beta / f).floor() * f).max(f);
    } else if t_bar * h_bar * w_bar < min_px {
        let beta = (min_px / (frames as f64 * hf * wf)).sqrt();
        h_bar = (hf * beta / f).ceil() * f;
        w_bar = (wf * beta / f).ceil() * f;
    }
    Ok((h_bar as usize, w_bar as usize))
}

/// Animated GIF -> prepared video: decode frames + delays, uniform-sample to an even
/// count <= VID_MAX_FRAMES, resize on the total-pixel budget, patchify CONSECUTIVE
/// frame pairs into temporal groups (frame 2g fills t=0, 2g+1 fills t=1). Timestamps
/// come from the GIF's own delays at the sampled indices (HF `_calculate_timestamps`).
pub fn prep_video_gif(bytes: &[u8]) -> Result<PreppedVideo, String> {
    use image::AnimationDecoder;
    let dec = image::codecs::gif::GifDecoder::new(std::io::Cursor::new(bytes))
        .map_err(|e| format!("gif decode: {e}"))?;
    let mut frames: Vec<(RgbImage, f32)> = Vec::new(); // (frame, start_seconds)
    let mut t = 0f32;
    for fr in dec.into_frames() {
        let fr = fr.map_err(|e| format!("gif frame: {e}"))?;
        let (num, den) = fr.delay().numer_denom_ms();
        let dt = if den == 0 {
            100.0
        } else {
            num as f32 / den as f32
        } / 1000.0;
        frames.push((
            image::DynamicImage::ImageRgba8(fr.into_buffer()).to_rgb8(),
            t,
        ));
        t += dt.max(0.01);
    }
    if frames.is_empty() {
        return Err("gif has no frames".into());
    }
    // still gif: duplicate the frame so one temporal group forms
    if frames.len() == 1 {
        let f0 = frames[0].clone();
        frames.push((f0.0, f0.1));
    }
    // uniform sample to an even count <= VID_MAX_FRAMES
    let total = frames.len();
    let take = total.min(VID_MAX_FRAMES) & !1;
    let picked: Vec<usize> = (0..take)
        .map(|i| i * total / take) // floor spacing, strictly increasing for take <= total
        .collect();
    let (h, w) = (frames[0].0.height() as usize, frames[0].0.width() as usize);
    let (rh, rw) = smart_resize_video(take, h, w)?;
    let (gh, gw) = (rh / V_PATCH, rw / V_PATCH);
    let mut groups = Vec::with_capacity(take / 2);
    let mut timestamps = Vec::with_capacity(take / 2);
    for g in 0..take / 2 {
        let (a, b) = (picked[2 * g], picked[2 * g + 1]);
        let mut patches = vec![0f32; gh * gw * V_PATCH_IN];
        for (slot, idx) in [(0usize, a), (1usize, b)] {
            let resized = image::imageops::resize(
                &frames[idx].0,
                rw as u32,
                rh as u32,
                FilterType::CatmullRom,
            );
            fill_slot(&mut patches, &resized, gh, gw, slot);
        }
        groups.push(PreppedImage { patches, gh, gw });
        timestamps.push(frames[a].1);
    }
    Ok(PreppedVideo { groups, timestamps })
}

/// Parse a base64 data URI into raw bytes (any `data:*;base64,` media type).
pub fn decode_data_uri(uri: &str) -> Result<Vec<u8>, String> {
    let rest = uri
        .strip_prefix("data:")
        .ok_or_else(|| "expected data: URI (http fetch requires MEMRA_FETCH_URLS=1)".to_string())?;
    let (meta, payload) = rest
        .split_once(',')
        .ok_or_else(|| "malformed data URI: no comma".to_string())?;
    if !meta.ends_with(";base64") {
        return Err("data URI must be base64-encoded".into());
    }
    base64::engine::general_purpose::STANDARD
        .decode(payload.trim())
        .map_err(|e| format!("base64 decode: {e}"))
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn smart_resize_multiples_and_budget() {
        // typical photo
        let (h, w) = smart_resize(1080, 1920).unwrap();
        assert_eq!(h % 32, 0);
        assert_eq!(w % 32, 0);
        assert!(h * w >= MIN_PIXELS && h * w <= MAX_PIXELS);
        // tiny icon scales UP to the floor
        let (h, w) = smart_resize(64, 64).unwrap();
        assert!(h * w >= MIN_PIXELS);
        // huge pano scales DOWN under the cap
        let (h, w) = smart_resize(8000, 12000).unwrap();
        assert!(h * w <= MAX_PIXELS);
        assert!(smart_resize(10, 4000).is_err()); // ar > 200
    }

    #[test]
    fn patchify_shape_and_order() {
        // 2x2-patch (32x32 px) synthetic image, distinct channel values
        let mut img = RgbImage::new(64, 64);
        for (x, y, p) in img.enumerate_pixels_mut() {
            *p = image::Rgb([x as u8, y as u8, 200]);
        }
        let mut buf = std::io::Cursor::new(Vec::new());
        img.write_to(&mut buf, image::ImageFormat::Png).unwrap();
        let prep = prep_image_bytes(buf.get_ref()).unwrap();
        assert_eq!(prep.patches.len(), prep.gh * prep.gw * V_PATCH_IN);
        assert_eq!(prep.gh % V_MERGE, 0);
        assert_eq!(prep.gw % V_MERGE, 0);
        // temporal slots identical for still images
        let row = &prep.patches[0..V_PATCH_IN];
        let slot = V_PATCH * V_PATCH;
        for c in 0..3 {
            let b = c * V_TEMPORAL * slot;
            assert_eq!(row[b..b + slot], row[b + slot..b + 2 * slot]);
        }
        // values in [-1, 1]
        assert!(prep.patches.iter().all(|v| (-1.0..=1.0).contains(v)));
    }

    #[test]
    fn data_uri_roundtrip() {
        let png = {
            let img = RgbImage::new(32, 32);
            let mut buf = std::io::Cursor::new(Vec::new());
            img.write_to(&mut buf, image::ImageFormat::Png).unwrap();
            buf.into_inner()
        };
        let uri = format!(
            "data:image/png;base64,{}",
            base64::engine::general_purpose::STANDARD.encode(&png)
        );
        let prep = prep_data_uri(&uri).unwrap();
        assert_eq!(prep.n_tokens(), prep.gh * prep.gw / 4);
        assert!(decode_data_uri("http://x/y.png").is_err());
    }
}