mod preprocess;
use core::fmt;
use std::path::Path;
use crate::{
ComputeUnits, DataType, Model, ModelDescription, MultiArray,
model::contract::{Checked, Dim, FeatureContract, LoadContract, StateContract},
};
use crate::embeddings::siglip::{
embedding::{EMBEDDING_DIM, Embedding, check_finite_output},
error::{
ContractMismatch, Error, ImageDataLength, ImageDimensions, OutputShape, PatchBudgetMismatch,
PreprocessedLength, PreprocessedMaskValue, PreprocessedNonFinite, PreprocessedPadNonZero,
Result, contract_violation,
},
image::preprocess::{parse_base_pos_grid, preprocess_image},
};
pub use preprocess::{MAX_IMAGE_AXIS, PATCH_DIM};
mod names {
pub const PIXEL_VALUES: &str = "pixel_values";
pub const POSITION_EMBEDDINGS: &str = "position_embeddings";
pub const ATTENTION_MASK: &str = "attention_mask";
pub const IMAGE_FEATURES: &str = "image_features";
}
pub const DEFAULT_IMAGE_COMPUTE: ComputeUnits = ComputeUnits::CpuAndGpu;
#[cfg(feature = "serde")]
fn default_image_compute() -> ComputeUnits {
DEFAULT_IMAGE_COMPUTE
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ImageEmbedderOptions {
#[cfg_attr(feature = "serde", serde(default = "default_image_compute"))]
compute: ComputeUnits,
}
impl Default for ImageEmbedderOptions {
fn default() -> Self {
Self::new()
}
}
impl ImageEmbedderOptions {
pub const fn new() -> Self {
Self {
compute: DEFAULT_IMAGE_COMPUTE,
}
}
#[inline]
pub const fn compute(&self) -> ComputeUnits {
self.compute
}
#[must_use]
#[inline]
pub const fn with_compute(mut self, compute: ComputeUnits) -> Self {
self.set_compute(compute);
self
}
#[inline]
pub const fn set_compute(&mut self, compute: ComputeUnits) -> &mut Self {
self.compute = compute;
self
}
}
#[derive(Debug, Clone, Copy)]
pub struct Rgb8Image<'a> {
data: &'a [u8],
width: usize,
height: usize,
}
impl<'a> Rgb8Image<'a> {
pub fn new(data: &'a [u8], width: usize, height: usize) -> Result<Self> {
if width == 0 || height == 0 {
return Err(Error::ImageDimensions(ImageDimensions::new(width, height)));
}
let expected = width
.checked_mul(height)
.and_then(|hw| hw.checked_mul(3))
.ok_or(Error::ImageDimensions(ImageDimensions::new(width, height)))?;
if data.len() != expected {
return Err(Error::ImageDataLength(ImageDataLength::new(
data.len(),
expected,
)));
}
Ok(Self {
data,
width,
height,
})
}
#[inline]
pub const fn width(&self) -> usize {
self.width
}
#[inline]
pub const fn height(&self) -> usize {
self.height
}
#[inline]
pub const fn data(&self) -> &'a [u8] {
self.data
}
}
#[derive(Clone)]
pub struct PreprocessedImage {
pixel_values: Vec<f32>,
position_embeddings: Vec<f32>,
attention_mask: Vec<f32>,
max_num_patches: usize,
}
impl PreprocessedImage {
pub fn try_new(
pixel_values: Vec<f32>,
position_embeddings: Vec<f32>,
attention_mask: Vec<f32>,
max_num_patches: usize,
) -> Result<Self> {
validate_budget_and_lengths(
&pixel_values,
&position_embeddings,
&attention_mask,
max_num_patches,
)?;
check_tensor_finite(names::PIXEL_VALUES, &pixel_values)?;
check_tensor_finite(names::POSITION_EMBEDDINGS, &position_embeddings)?;
let num_real = validate_mask(&attention_mask)?;
validate_pad_rows(names::PIXEL_VALUES, &pixel_values, num_real, PATCH_DIM)?;
validate_pad_rows(
names::POSITION_EMBEDDINGS,
&position_embeddings,
num_real,
EMBEDDING_DIM,
)?;
Ok(Self {
pixel_values,
position_embeddings,
attention_mask,
max_num_patches,
})
}
fn from_pipeline(
pixel_values: Vec<f32>,
position_embeddings: Vec<f32>,
attention_mask: Vec<f32>,
max_num_patches: usize,
) -> Self {
debug_assert!(
validate_structural(
&pixel_values,
&position_embeddings,
&attention_mask,
max_num_patches
)
.is_ok(),
"internal NaFlex pipeline emitted a structurally invalid tensor bundle"
);
Self {
pixel_values,
position_embeddings,
attention_mask,
max_num_patches,
}
}
#[inline]
pub const fn max_num_patches(&self) -> usize {
self.max_num_patches
}
#[inline]
pub fn pixel_values(&self) -> &[f32] {
&self.pixel_values
}
#[inline]
pub fn position_embeddings(&self) -> &[f32] {
&self.position_embeddings
}
#[inline]
pub fn attention_mask(&self) -> &[f32] {
&self.attention_mask
}
}
impl fmt::Debug for PreprocessedImage {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let num_real = self
.attention_mask
.iter()
.take_while(|&&v| v == 1.0)
.count();
f.debug_struct("PreprocessedImage")
.field("max_num_patches", &self.max_num_patches)
.field("num_real_patches", &num_real)
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub struct ImageEmbedder {
model: Checked,
base_pos_embed: Vec<f32>,
max_num_patches: usize,
}
impl ImageEmbedder {
pub fn load(
model_path: impl AsRef<Path>,
pos_embed_path: impl AsRef<Path>,
options: ImageEmbedderOptions,
) -> Result<Self> {
let bytes = std::fs::read(pos_embed_path.as_ref()).map_err(Error::PosEmbedLoad)?;
let base_pos_embed = parse_base_pos_grid(&bytes)?;
Self::from_parts(model_path, base_pos_embed, options)
}
pub fn from_files(
model_path: impl AsRef<Path>,
pos_embed_path: impl AsRef<Path>,
) -> Result<Self> {
Self::load(model_path, pos_embed_path, ImageEmbedderOptions::new())
}
pub fn from_memory(
model_path: impl AsRef<Path>,
pos_embed_bytes: &[u8],
options: ImageEmbedderOptions,
) -> Result<Self> {
let base_pos_embed = parse_base_pos_grid(pos_embed_bytes)?;
Self::from_parts(model_path, base_pos_embed, options)
}
fn from_parts(
model_path: impl AsRef<Path>,
base_pos_embed: Vec<f32>,
options: ImageEmbedderOptions,
) -> Result<Self> {
let model = Model::load(model_path, options.compute())?;
let declared = declared_patch_budget(model.description())?;
let model = Checked::new(model, &image_contract(declared)).map_err(contract_violation)?;
let max_num_patches = declared_patch_budget(model.description())
.expect("the load contract established `pixel_values` and its rank");
Ok(Self {
model,
base_pos_embed,
max_num_patches,
})
}
#[inline]
pub const fn max_num_patches(&self) -> usize {
self.max_num_patches
}
pub fn preprocess(&self, image: Rgb8Image<'_>) -> Result<PreprocessedImage> {
let inputs = preprocess_image(
image.data(),
image.width(),
image.height(),
&self.base_pos_embed,
self.max_num_patches,
)?;
Ok(PreprocessedImage::from_pipeline(
inputs.pixel_values,
inputs.position_embeddings,
inputs.attention_mask,
self.max_num_patches,
))
}
pub fn embed(&self, image: Rgb8Image<'_>) -> Result<Embedding> {
let inputs = self.preprocess(image)?;
self.embed_preprocessed(&inputs)
}
pub fn embed_preprocessed(&self, inputs: &PreprocessedImage) -> Result<Embedding> {
check_patch_budget(inputs.max_num_patches(), self.max_num_patches)?;
self.predict_embedding(
inputs.pixel_values(),
inputs.position_embeddings(),
inputs.attention_mask(),
)
}
fn predict_embedding(
&self,
pixel_values: &[f32],
position_embeddings: &[f32],
attention_mask: &[f32],
) -> Result<Embedding> {
let pixel_values = MultiArray::from_slice(&[1, self.max_num_patches, PATCH_DIM], pixel_values)?;
let position_embeddings = MultiArray::from_slice(
&[1, self.max_num_patches, EMBEDDING_DIM],
position_embeddings,
)?;
let attention_mask = MultiArray::from_slice(&[1, self.max_num_patches], attention_mask)?;
let mut outputs = self.model.predict_with(&[
(names::PIXEL_VALUES, &pixel_values),
(names::POSITION_EMBEDDINGS, &position_embeddings),
(names::ATTENTION_MASK, &attention_mask),
])?;
let feats = outputs
.take(names::IMAGE_FEATURES)
.ok_or_else(|| crate::PredictionError::MissingOutput(names::IMAGE_FEATURES.to_string()))?;
if feats.shape() != [1, EMBEDDING_DIM] {
return Err(Error::OutputShape(OutputShape::new(
feats.shape().to_vec(),
vec![1, EMBEDDING_DIM],
)));
}
let mut row = [0.0f32; EMBEDDING_DIM];
feats.copy_into::<f32>(&mut row)?;
check_finite_output(&row)?;
Embedding::from_slice_normalizing(&row)
}
pub fn prewarm(&self) -> Result<()> {
let data = vec![128u8; 64 * 64 * 3];
let image = Rgb8Image::new(&data, 64, 64)?;
self.embed(image)?;
Ok(())
}
}
fn image_contract(p: usize) -> LoadContract {
LoadContract::new(
vec![
FeatureContract::new(
names::PIXEL_VALUES,
DataType::F32,
vec![Dim::Exactly(1), Dim::AnyFixed, Dim::Exactly(PATCH_DIM)],
),
FeatureContract::new(
names::POSITION_EMBEDDINGS,
DataType::F32,
vec![
Dim::Exactly(1),
Dim::Exactly(p),
Dim::Exactly(EMBEDDING_DIM),
],
),
FeatureContract::new(
names::ATTENTION_MASK,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(p)],
),
],
vec![FeatureContract::new(
names::IMAGE_FEATURES,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(EMBEDDING_DIM)],
)],
StateContract::None,
)
}
fn declared_patch_budget(description: &ModelDescription) -> Result<usize> {
let expected = format!("[1, P, {PATCH_DIM}] float32 with P >= 1");
let declared = description.input(names::PIXEL_VALUES).ok_or_else(|| {
Error::ContractMismatch(ContractMismatch::new(
names::PIXEL_VALUES,
expected.clone(),
"missing".to_string(),
))
})?;
match declared.shape() {
[_, p, _] if *p > 0 => Ok(*p),
shape => Err(Error::ContractMismatch(ContractMismatch::new(
names::PIXEL_VALUES,
expected,
format!("{shape:?}"),
))),
}
}
fn validate_budget_and_lengths(
pixel_values: &[f32],
position_embeddings: &[f32],
attention_mask: &[f32],
max_num_patches: usize,
) -> Result<()> {
const MAX_ROW_DIM: usize = if PATCH_DIM > EMBEDDING_DIM {
PATCH_DIM
} else {
EMBEDDING_DIM
};
if max_num_patches == 0 || max_num_patches > usize::MAX / MAX_ROW_DIM {
return Err(Error::PreprocessedPatchBudget(max_num_patches));
}
check_len(
names::PIXEL_VALUES,
pixel_values.len(),
max_num_patches * PATCH_DIM,
)?;
check_len(
names::POSITION_EMBEDDINGS,
position_embeddings.len(),
max_num_patches * EMBEDDING_DIM,
)?;
check_len(names::ATTENTION_MASK, attention_mask.len(), max_num_patches)?;
Ok(())
}
fn check_len(feature: &'static str, got: usize, expected: usize) -> Result<()> {
if got != expected {
return Err(Error::PreprocessedLength(PreprocessedLength::new(
feature, got, expected,
)));
}
Ok(())
}
fn check_tensor_finite(feature: &'static str, values: &[f32]) -> Result<()> {
if let Some(index) = values.iter().position(|v| !v.is_finite()) {
return Err(Error::PreprocessedNonFinite(PreprocessedNonFinite::new(
feature, index,
)));
}
Ok(())
}
fn validate_mask(mask: &[f32]) -> Result<usize> {
let mut num_real = 0usize;
let mut in_pad = false;
for (index, &value) in mask.iter().enumerate() {
if value == 1.0 {
if in_pad {
return Err(Error::PreprocessedMaskOrder(index));
}
num_real += 1;
} else if value == 0.0 {
in_pad = true;
} else {
return Err(Error::PreprocessedMaskValue(PreprocessedMaskValue::new(
index, value,
)));
}
}
if num_real == 0 {
return Err(Error::PreprocessedMaskEmpty);
}
Ok(num_real)
}
fn validate_pad_rows(
feature: &'static str,
values: &[f32],
num_real: usize,
row_dim: usize,
) -> Result<()> {
let pad_start = num_real * row_dim;
if let Some(offset) = values[pad_start..].iter().position(|&v| v != 0.0) {
return Err(Error::PreprocessedPadNonZero(PreprocessedPadNonZero::new(
feature,
pad_start + offset,
)));
}
Ok(())
}
fn validate_structural(
pixel_values: &[f32],
position_embeddings: &[f32],
attention_mask: &[f32],
max_num_patches: usize,
) -> Result<()> {
validate_budget_and_lengths(
pixel_values,
position_embeddings,
attention_mask,
max_num_patches,
)?;
let num_real = validate_mask(attention_mask)?;
validate_pad_rows(names::PIXEL_VALUES, pixel_values, num_real, PATCH_DIM)?;
validate_pad_rows(
names::POSITION_EMBEDDINGS,
position_embeddings,
num_real,
EMBEDDING_DIM,
)?;
Ok(())
}
fn check_patch_budget(input: usize, model: usize) -> Result<()> {
if input != model {
return Err(Error::PatchBudgetMismatch(PatchBudgetMismatch::new(
input, model,
)));
}
Ok(())
}
#[cfg(test)]
mod tests;