use std::path::Path;
use crate::{
ComputeUnits, DataType, Model, MultiArray,
model::contract::{Checked, Dim, FeatureContract, LoadContract, StateContract},
};
use tokenizers::{
PostProcessor, Tokenizer, TruncationDirection, TruncationParams, TruncationStrategy,
};
use crate::embeddings::clap::{
embedding::{EMBEDDING_DIM, Embedding, check_finite_output},
error::{Error, OutputShape, Result, SpecialTokenOverhead, TokenCount, contract_violation},
};
mod names {
pub const INPUT_IDS: &str = "input_ids";
pub const ATTENTION_MASK: &str = "attention_mask";
pub const TEXT_EMBEDS: &str = "text_embeds";
}
pub const TEXT_MAX_TOKENS: usize = 512;
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 TextEncoderOptions {
#[cfg_attr(feature = "serde", serde(default = "default_text_compute"))]
compute: ComputeUnits,
}
impl Default for TextEncoderOptions {
fn default() -> Self {
Self::new()
}
}
impl TextEncoderOptions {
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)]
pub struct TextEncoder {
model: Checked,
tokenizer: Tokenizer,
pad_id: i32,
}
impl TextEncoder {
pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self> {
Self::from_bundled_tokenizer(model_path, TextEncoderOptions::new())
}
pub fn from_bundled_tokenizer(
model_path: impl AsRef<Path>,
options: TextEncoderOptions,
) -> Result<Self> {
let tokenizer = Tokenizer::from_bytes(crate::embeddings::clap::BUNDLED_TOKENIZER)
.map_err(Error::TokenizerLoad)?;
Self::from_parts(model_path, tokenizer, options)
}
pub fn from_files(
model_path: impl AsRef<Path>,
tokenizer_json_path: impl AsRef<Path>,
options: TextEncoderOptions,
) -> 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: TextEncoderOptions,
) -> Result<Self> {
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: TextEncoderOptions,
) -> Result<Self> {
configure_tokenizer(&mut tokenizer)?;
let pad_id = resolve_pad_id(&tokenizer)?;
let model = Model::load(model_path, options.compute())?;
let model = Checked::new(model, &text_contract()).map_err(contract_violation)?;
Ok(Self {
model,
tokenizer,
pad_id,
})
}
pub fn token_ids(&self, text: &str) -> Result<Vec<u32>> {
if text.is_empty() {
return Err(Error::EmptyText);
}
let encoding = self.tokenizer.encode(text, true).map_err(Error::Tokenize)?;
Ok(encoding.get_ids().to_vec())
}
pub fn embed(&self, text: &str) -> Result<Embedding> {
let ids = self.token_ids(text)?;
let (input_ids, attention_mask) = build_window(&ids, self.pad_id)?;
let ids_tensor = MultiArray::from_slice(&[1, TEXT_MAX_TOKENS], &input_ids)?;
let mask_tensor = MultiArray::from_slice(&[1, TEXT_MAX_TOKENS], &attention_mask)?;
let mut outputs = self.model.predict_with(&[
(names::INPUT_IDS, &ids_tensor),
(names::ATTENTION_MASK, &mask_tensor),
])?;
let embeds = outputs
.take(names::TEXT_EMBEDS)
.ok_or_else(|| crate::PredictionError::MissingOutput(names::TEXT_EMBEDS.to_string()))?;
if embeds.shape() != [1, EMBEDDING_DIM] {
return Err(Error::OutputShape(OutputShape::new(
embeds.shape().to_vec(),
vec![1, EMBEDDING_DIM],
)));
}
let mut row = [0.0f32; EMBEDDING_DIM];
embeds.copy_into::<f32>(&mut row)?;
check_finite_output(&row)?;
Embedding::from_slice_normalizing(&row)
}
pub fn prewarm(&self) -> Result<()> {
self.embed("warmup")?;
Ok(())
}
}
fn configure_tokenizer(tokenizer: &mut Tokenizer) -> 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 >= TEXT_MAX_TOKENS {
return Err(Error::SpecialTokenOverhead(SpecialTokenOverhead::new(
added,
TEXT_MAX_TOKENS,
)));
}
tokenizer
.with_truncation(Some(TruncationParams {
max_length: TEXT_MAX_TOKENS,
strategy: TruncationStrategy::LongestFirst,
stride: 0,
direction: TruncationDirection::Right,
}))
.map_err(Error::TokenizerConfig)?;
tokenizer.with_padding(None);
Ok(())
}
const FALLBACK_PAD_ID: i32 = 1;
fn resolve_pad_id(tokenizer: &Tokenizer) -> Result<i32> {
tokenizer
.token_to_id("<pad>")
.map_or(Ok(FALLBACK_PAD_ID), |id| {
i32::try_from(id).map_err(|_| Error::TokenIdRange(id))
})
}
fn build_window(
ids: &[u32],
pad_id: i32,
) -> Result<([i32; TEXT_MAX_TOKENS], [i32; TEXT_MAX_TOKENS])> {
if ids.len() > TEXT_MAX_TOKENS {
return Err(Error::TokenCount(TokenCount::new(
ids.len(),
TEXT_MAX_TOKENS,
)));
}
let mut input_ids = [pad_id; TEXT_MAX_TOKENS];
let mut attention_mask = [0i32; TEXT_MAX_TOKENS];
for (i, &id) in ids.iter().enumerate() {
input_ids[i] = i32::try_from(id).map_err(|_| Error::TokenIdRange(id))?;
attention_mask[i] = 1;
}
Ok((input_ids, attention_mask))
}
#[cfg(test)]
pub(crate) fn configured_tokenizer_from_bytes(bytes: &[u8]) -> Result<Tokenizer> {
let mut tokenizer = Tokenizer::from_bytes(bytes).map_err(Error::TokenizerLoad)?;
configure_tokenizer(&mut tokenizer)?;
Ok(tokenizer)
}
fn text_contract() -> LoadContract {
let window = vec![Dim::Exactly(1), Dim::Exactly(TEXT_MAX_TOKENS)];
LoadContract::new(
vec![
FeatureContract::new(names::INPUT_IDS, DataType::I32, window.clone()),
FeatureContract::new(names::ATTENTION_MASK, DataType::I32, window),
],
vec![FeatureContract::new(
names::TEXT_EMBEDS,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(EMBEDDING_DIM)],
)],
StateContract::None,
)
}
#[cfg(test)]
mod tests;