use ort::session::Session;
use std::collections::{HashMap, HashSet};
use tokenizers::Tokenizer;
use crate::{
init::{HasMaxLength, InitOptionsWithLength},
models::sparse::SparseModel,
TokenizerFiles,
};
use super::DEFAULT_MAX_LENGTH;
impl HasMaxLength for SparseModel {
const MAX_LENGTH: usize = DEFAULT_MAX_LENGTH;
}
pub type SparseInitOptions = InitOptionsWithLength<SparseModel>;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct UserDefinedSparseModel {
pub onnx_file: Vec<u8>,
pub tokenizer_files: TokenizerFiles,
pub model: SparseModel,
pub idf_file: Option<Vec<u8>>,
}
impl UserDefinedSparseModel {
pub fn new(onnx_file: Vec<u8>, tokenizer_files: TokenizerFiles, model: SparseModel) -> Self {
Self {
onnx_file,
tokenizer_files,
model,
idf_file: None,
}
}
pub fn with_idf_file(mut self, idf_file: Vec<u8>) -> Self {
self.idf_file = Some(idf_file);
self
}
}
pub struct SparseTextEmbedding {
pub tokenizer: Tokenizer,
pub(crate) session: Session,
pub(crate) need_token_type_ids: bool,
pub(crate) model: SparseModel,
pub(crate) special_token_ids: HashSet<usize>,
pub(crate) token_id_to_idf: Option<HashMap<usize, f32>>,
}