use std::path::Path;
use crate::{
ComputeUnits, DataType, Model, MultiArray,
model::contract::{Checked, Dim, FeatureContract, LoadContract, StateContract},
};
use tokenizers::{
PostProcessor, Tokenizer, TruncationDirection, TruncationParams, TruncationStrategy,
normalizers::{Lowercase, NormalizerWrapper, Sequence as NormalizerSequence},
};
use crate::embeddings::siglip::{
embedding::{EMBEDDING_DIM, Embedding, check_finite_output},
error::{
ArtifactTokenizerIdentity, ArtifactTokenizerRead, Error, OutputShape, Result,
SpecialTokenOverhead, TokenCount, contract_violation,
},
};
mod names {
pub const INPUT_IDS: &str = "input_ids";
pub const TEXT_FEATURES: &str = "text_features";
}
mod contract {
pub const TOKENIZER_SHA256_HEX: &str =
"58a1696e79c9d97937389ed116f552a15c84811d7b8023918b86f4bc5775b1b0";
}
const PLACEHOLDER_SENTINEL: &[u8] =
b"PLACEHOLDER_REPLACE_WITH_SOURCE_REVISION_GEMMA_TOKENIZER_IN_WAVE_B";
fn ensure_not_placeholder(bytes: &[u8]) -> Result<()> {
if bytes.len() < 1_000_000
&& bytes
.windows(PLACEHOLDER_SENTINEL.len())
.any(|w| w == PLACEHOLDER_SENTINEL)
{
return Err(Error::TokenizerPlaceholder);
}
Ok(())
}
fn artifact_tokenizer_path(model_path: &Path) -> std::path::PathBuf {
model_path
.parent()
.unwrap_or_else(|| Path::new(""))
.join(crate::embeddings::siglip::TOKENIZER_FILE_NAME)
}
fn ensure_pinned_identity(bytes: &[u8], path: &Path) -> Result<()> {
use sha2::{Digest, Sha256};
let actual: String = Sha256::digest(bytes)
.iter()
.map(|b| format!("{b:02x}"))
.collect();
if actual == contract::TOKENIZER_SHA256_HEX {
return Ok(());
}
Err(Error::ArtifactTokenizerIdentity(
ArtifactTokenizerIdentity::new(path.to_path_buf(), contract::TOKENIZER_SHA256_HEX, actual),
))
}
pub const DEFAULT_TEXT_COMPUTE: ComputeUnits = ComputeUnits::CpuAndGpu;
#[cfg(feature = "serde")]
fn default_text_compute() -> ComputeUnits {
DEFAULT_TEXT_COMPUTE
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct TextEmbedderOptions {
#[cfg_attr(feature = "serde", serde(default = "default_text_compute"))]
compute: ComputeUnits,
}
impl Default for TextEmbedderOptions {
fn default() -> Self {
Self::new()
}
}
impl TextEmbedderOptions {
pub const fn new() -> Self {
Self {
compute: DEFAULT_TEXT_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, PartialEq, Eq)]
pub(crate) enum PadSide {
Right,
#[allow(dead_code)]
Left,
}
#[derive(Debug)]
pub struct TextEmbedder {
model: Checked,
tokenizer: Tokenizer,
pad_id: i32,
pad_side: PadSide,
max_tokens: usize,
}
impl TextEmbedder {
pub fn load(model_path: impl AsRef<Path>, options: TextEmbedderOptions) -> Result<Self> {
let model_path = model_path.as_ref();
let tokenizer_path = artifact_tokenizer_path(model_path);
let bytes = std::fs::read(&tokenizer_path).map_err(|source| {
Error::ArtifactTokenizerRead(ArtifactTokenizerRead::new(tokenizer_path.clone(), source))
})?;
ensure_not_placeholder(&bytes)?;
ensure_pinned_identity(&bytes, &tokenizer_path)?;
let tokenizer = Tokenizer::from_bytes(&bytes).map_err(Error::TokenizerLoad)?;
Self::from_parts(model_path, tokenizer, options)
}
pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self> {
Self::load(model_path, TextEmbedderOptions::new())
}
pub fn from_files(
model_path: impl AsRef<Path>,
tokenizer_json_path: impl AsRef<Path>,
options: TextEmbedderOptions,
) -> Result<Self> {
let tokenizer =
Tokenizer::from_file(tokenizer_json_path.as_ref()).map_err(Error::TokenizerLoad)?;
Self::from_parts(model_path, tokenizer, options)
}
pub fn from_memory(
model_path: impl AsRef<Path>,
tokenizer_json_bytes: &[u8],
options: TextEmbedderOptions,
) -> Result<Self> {
ensure_not_placeholder(tokenizer_json_bytes)?;
let tokenizer = Tokenizer::from_bytes(tokenizer_json_bytes).map_err(Error::TokenizerLoad)?;
Self::from_parts(model_path, tokenizer, options)
}
fn from_parts(
model_path: impl AsRef<Path>,
mut tokenizer: Tokenizer,
options: TextEmbedderOptions,
) -> Result<Self> {
let model = Model::load(model_path, options.compute())?;
let model = Checked::new(model, &text_contract()).map_err(contract_violation)?;
let max_tokens = model
.description()
.input(names::INPUT_IDS)
.and_then(|declared| declared.shape().get(1).copied())
.expect("the contract names `input_ids` at rank 2 and the check passed");
configure_tokenizer(&mut tokenizer, max_tokens)?;
let pad_id = resolve_pad_id(&tokenizer)?;
Ok(Self {
model,
tokenizer,
pad_id,
pad_side: PadSide::Right,
max_tokens,
})
}
#[inline]
pub const fn max_tokens(&self) -> usize {
self.max_tokens
}
pub fn token_ids(&self, text: &str) -> Result<Vec<i32>> {
if text.is_empty() {
return Err(Error::EmptyText);
}
let encoding = self.tokenizer.encode(text, true).map_err(Error::Tokenize)?;
build_window(
encoding.get_ids(),
self.pad_id,
self.pad_side,
self.max_tokens,
)
}
pub fn embed(&self, text: &str) -> Result<Embedding> {
let ids = self.token_ids(text)?;
let ids_tensor = MultiArray::from_slice(&[1, self.max_tokens], &ids)?;
let mut outputs = self
.model
.predict_with(&[(names::INPUT_IDS, &ids_tensor)])?;
let feats = outputs
.take(names::TEXT_FEATURES)
.ok_or_else(|| crate::PredictionError::MissingOutput(names::TEXT_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<()> {
self.embed("warmup")?;
Ok(())
}
}
fn text_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
names::INPUT_IDS,
DataType::I32,
vec![Dim::Exactly(1), Dim::AnyFixed],
)],
vec![FeatureContract::new(
names::TEXT_FEATURES,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(EMBEDDING_DIM)],
)],
StateContract::None,
)
}
fn configure_tokenizer(tokenizer: &mut Tokenizer, max_tokens: usize) -> Result<()> {
crate::embeddings::tokenizer_guard::check_post_processor(tokenizer)
.map_err(Error::PostProcessorTemplate)?;
let added = tokenizer
.get_post_processor()
.map_or(0, |post| post.added_tokens(false));
if added >= max_tokens {
return Err(Error::SpecialTokenOverhead(SpecialTokenOverhead::new(
added, max_tokens,
)));
}
let lowercased: NormalizerWrapper = match tokenizer.get_normalizer() {
Some(existing) => NormalizerSequence::new(vec![Lowercase.into(), existing.clone()]).into(),
None => Lowercase.into(),
};
tokenizer
.with_normalizer(Some(lowercased))
.map_err(Error::TokenizerConfig)?;
tokenizer
.with_truncation(Some(TruncationParams {
max_length: max_tokens,
strategy: TruncationStrategy::LongestFirst,
stride: 0,
direction: TruncationDirection::Right,
}))
.map_err(Error::TokenizerConfig)?;
tokenizer.with_padding(None);
Ok(())
}
fn resolve_pad_id(tokenizer: &Tokenizer) -> Result<i32> {
tokenizer.token_to_id("<pad>").map_or(Ok(0), |id| {
i32::try_from(id).map_err(|_| Error::TokenIdRange(id))
})
}
fn build_window(
ids: &[u32],
pad_id: i32,
pad_side: PadSide,
max_tokens: usize,
) -> Result<Vec<i32>> {
if ids.len() > max_tokens {
return Err(Error::TokenCount(TokenCount::new(ids.len(), max_tokens)));
}
let mut window = vec![pad_id; max_tokens];
let offset = match pad_side {
PadSide::Right => 0,
PadSide::Left => max_tokens - ids.len(),
};
for (i, &id) in ids.iter().enumerate() {
window[offset + i] = i32::try_from(id).map_err(|_| Error::TokenIdRange(id))?;
}
Ok(window)
}
#[doc(hidden)]
pub fn configured_tokenizer_from_bytes(bytes: &[u8], max_tokens: usize) -> Result<Tokenizer> {
let mut tokenizer = Tokenizer::from_bytes(bytes).map_err(Error::TokenizerLoad)?;
configure_tokenizer(&mut tokenizer, max_tokens)?;
Ok(tokenizer)
}
#[cfg(test)]
mod tests;