use std::path::Path;
use std::sync::Mutex;
use ndarray::{Array2, Array4};
use ort::session::Session;
use ort::value::Tensor;
use tokenizers::Tokenizer;
use crate::math::{cosine_similarity, l2_normalize};
use crate::{Error, Image, Result, Scorer};
const IMAGE_SIZE: u32 = 224;
const CLIP_MEAN: [f32; 3] = [0.481_454_7, 0.457_827_5, 0.408_210_7];
const CLIP_STD: [f32; 3] = [0.268_629_5, 0.261_302_6, 0.275_777_1];
const CTX_LEN: usize = 77;
const EOT_TOKEN: i64 = 49407;
pub struct ClipEmbedder {
session: Mutex<Session>,
tokenizer: Tokenizer,
}
impl ClipEmbedder {
pub fn from_files(model: impl AsRef<Path>, tokenizer: impl AsRef<Path>) -> Result<Self> {
let session = Session::builder()
.and_then(|mut b| b.commit_from_file(model.as_ref()))
.map_err(|e| Error::Scorer(format!("loading CLIP model: {e}")))?;
let tokenizer = Tokenizer::from_file(tokenizer.as_ref())
.map_err(|e| Error::Scorer(format!("loading CLIP tokenizer: {e}")))?;
Ok(Self { session: Mutex::new(session), tokenizer })
}
pub fn embed_image(&self, image: &Image) -> Result<Vec<f32>> {
let (ids, mask) = self.tokenize("")?;
let mut outs = self.run(preprocess(image)?, ids, mask, &["image_embeds"])?;
let mut v = outs.remove(0);
l2_normalize(&mut v);
Ok(v)
}
pub fn embed_text(&self, text: &str) -> Result<Vec<f32>> {
let (ids, mask) = self.tokenize(text)?;
let blank = Array4::<f32>::zeros((1, 3, IMAGE_SIZE as usize, IMAGE_SIZE as usize));
let mut outs = self.run(blank, ids, mask, &["text_embeds"])?;
let mut v = outs.remove(0);
l2_normalize(&mut v);
Ok(v)
}
pub fn embed_both(&self, text: &str, image: &Image) -> Result<(Vec<f32>, Vec<f32>)> {
let (ids, mask) = self.tokenize(text)?;
let mut outs = self.run(preprocess(image)?, ids, mask, &["image_embeds", "text_embeds"])?;
let mut txt = outs.remove(1);
let mut img = outs.remove(0);
l2_normalize(&mut img);
l2_normalize(&mut txt);
Ok((img, txt))
}
pub fn image_similarity(&self, a: &Image, b: &Image) -> Result<f32> {
Ok(cosine_similarity(&self.embed_image(a)?, &self.embed_image(b)?))
}
fn tokenize(&self, prompt: &str) -> Result<(Array2<i64>, Array2<i64>)> {
let encoding = self
.tokenizer
.encode(prompt, true)
.map_err(|e| Error::Scorer(format!("tokenizing prompt: {e}")))?;
let ids = encoding.get_ids();
let mut input_ids = Array2::<i64>::from_elem((1, CTX_LEN), EOT_TOKEN);
let mut attention = Array2::<i64>::zeros((1, CTX_LEN));
for (i, &id) in ids.iter().take(CTX_LEN).enumerate() {
input_ids[[0, i]] = id as i64;
attention[[0, i]] = 1;
}
Ok((input_ids, attention))
}
fn run(
&self,
pixel_values: Array4<f32>,
input_ids: Array2<i64>,
attention_mask: Array2<i64>,
wants: &[&str],
) -> Result<Vec<Vec<f32>>> {
let pv = Tensor::from_array(pixel_values)
.map_err(|e| Error::Scorer(format!("pixel tensor: {e}")))?;
let ids = Tensor::from_array(input_ids)
.map_err(|e| Error::Scorer(format!("input_ids tensor: {e}")))?;
let mask = Tensor::from_array(attention_mask)
.map_err(|e| Error::Scorer(format!("attention tensor: {e}")))?;
let mut session =
self.session.lock().map_err(|_| Error::Scorer("CLIP session lock poisoned".into()))?;
let outputs = session
.run(ort::inputs![
"pixel_values" => pv,
"input_ids" => ids,
"attention_mask" => mask,
])
.map_err(|e| Error::Scorer(format!("CLIP inference: {e}")))?;
wants.iter().map(|name| extract_vec(&outputs, name)).collect()
}
}
pub struct ClipScorer {
embedder: ClipEmbedder,
}
impl ClipScorer {
pub fn from_files(model: impl AsRef<Path>, tokenizer: impl AsRef<Path>) -> Result<Self> {
Ok(Self { embedder: ClipEmbedder::from_files(model, tokenizer)? })
}
pub fn embedder(&self) -> &ClipEmbedder {
&self.embedder
}
}
impl Scorer for ClipScorer {
fn score(&self, prompt: &str, image: &Image) -> Result<f32> {
let (image_embed, text_embed) = self.embedder.embed_both(prompt, image)?;
Ok(cosine_similarity(&image_embed, &text_embed))
}
}
fn extract_vec(outputs: &ort::session::SessionOutputs, name: &str) -> Result<Vec<f32>> {
let value = outputs
.get(name)
.ok_or_else(|| Error::Scorer(format!("CLIP model has no `{name}` output")))?;
let (_shape, data) = value
.try_extract_tensor::<f32>()
.map_err(|e| Error::Scorer(format!("extracting `{name}`: {e}")))?;
Ok(data.to_vec())
}
fn preprocess(image: &Image) -> Result<Array4<f32>> {
use image::{DynamicImage, RgbImage, imageops::FilterType};
let rgb = RgbImage::from_raw(image.width, image.height, image.rgb.clone())
.ok_or_else(|| Error::Scorer("image buffer does not match its dimensions".into()))?;
let short = image.width.min(image.height).max(1) as f32;
let scale = IMAGE_SIZE as f32 / short;
let nw = ((image.width as f32 * scale).round() as u32).max(IMAGE_SIZE);
let nh = ((image.height as f32 * scale).round() as u32).max(IMAGE_SIZE);
let resized =
DynamicImage::ImageRgb8(rgb).resize_exact(nw, nh, FilterType::CatmullRom).to_rgb8();
let left = (nw - IMAGE_SIZE) / 2;
let top = (nh - IMAGE_SIZE) / 2;
let mut arr = Array4::<f32>::zeros((1, 3, IMAGE_SIZE as usize, IMAGE_SIZE as usize));
for y in 0..IMAGE_SIZE {
for x in 0..IMAGE_SIZE {
let p = resized.get_pixel(left + x, top + y);
for c in 0..3 {
arr[[0, c, y as usize, x as usize]] =
(p[c] as f32 / 255.0 - CLIP_MEAN[c]) / CLIP_STD[c];
}
}
}
Ok(arr)
}
#[cfg(test)]
mod tests {
use super::*;
fn solid(width: u32, height: u32, value: u8) -> Image {
Image::new(width, height, vec![value; (width * height * 3) as usize]).unwrap()
}
#[test]
fn preprocess_yields_normalized_nchw_tensor() {
let tensor = preprocess(&solid(64, 40, 128)).unwrap();
assert_eq!(tensor.shape(), &[1, 3, 224, 224]);
for &v in tensor.iter() {
assert!(v.is_finite() && v.abs() < 5.0);
}
}
#[test]
fn preprocess_normalization_matches_formula() {
let tensor = preprocess(&solid(224, 224, 255)).unwrap();
let expected = (1.0 - CLIP_MEAN[0]) / CLIP_STD[0];
assert!((tensor[[0, 0, 0, 0]] - expected).abs() < 1e-3);
}
}