use tokenizers::Tokenizer;
use ort::session::Session;
use ndarray::{Array1, Array2};
use ort::value::Tensor;
use std::convert::TryFrom;
use std::collections::HashMap;
use crate::classifier::error::ClassifierError;
use super::utils::normalize_vector;
pub(crate) trait TextEmbedding {
fn tokenizer(&self) -> Option<&Tokenizer>;
fn session(&self) -> Option<&Session>;
fn tokenize(&self, text: &str) -> Result<Vec<u32>, ClassifierError> {
let tokenizer = self.tokenizer()
.ok_or_else(|| ClassifierError::TokenizerError("Tokenizer not initialized".into()))?;
let encoding = tokenizer.encode(text, false)
.map_err(|e| ClassifierError::TokenizerError(e.to_string()))?;
let token_ids = encoding.get_ids();
let safe_tokens: Result<Vec<u32>, _> = token_ids.iter()
.map(|&id| u32::try_from(id))
.collect();
safe_tokens.map_err(|_| ClassifierError::TokenizerError("Invalid token ID encountered".into()))
}
fn embed_text(&self, text: &str) -> Result<Array1<f32>, ClassifierError> {
let tokens = self.tokenize(text)?;
self.get_embedding(&tokens)
}
fn get_embedding(&self, tokens: &[u32]) -> Result<Array1<f32>, ClassifierError> {
let session = self.session()
.ok_or_else(|| ClassifierError::ModelError("Session not initialized".into()))?;
let input_array = Array2::from_shape_vec((1, tokens.len()),
tokens.iter().map(|&x| x as i64).collect())
.map_err(|e| ClassifierError::ModelError(format!("Failed to create input array: {}", e)))?;
let input_dyn = input_array.into_dyn();
let input_ids = input_dyn.as_standard_layout();
let mask_array = Array2::from_shape_vec((1, tokens.len()),
tokens.iter().map(|&x| if x == 0 { 0i64 } else { 1i64 }).collect())
.map_err(|e| ClassifierError::ModelError(format!("Failed to create mask array: {}", e)))?;
let mask_dyn = mask_array.into_dyn();
let attention_mask = mask_dyn.as_standard_layout();
let mut input_tensors = HashMap::new();
input_tensors.insert("input_ids", Tensor::from_array(&input_ids)
.map_err(|e| ClassifierError::ModelError(format!("Failed to create input tensor: {}", e)))?);
input_tensors.insert("attention_mask", Tensor::from_array(&attention_mask)
.map_err(|e| ClassifierError::ModelError(format!("Failed to create mask tensor: {}", e)))?);
let outputs = session.run(input_tensors)
.map_err(|e| ClassifierError::ModelError(format!("Failed to run model: {}", e)))?;
let output_tensor = outputs[0].try_extract_tensor::<f32>()
.map_err(|e| ClassifierError::ModelError(format!("Failed to extract output tensor: {}", e)))?;
let mut embedding = Array1::zeros(output_tensor.shape()[2]);
let embedding_slice = output_tensor.slice(ndarray::s![0, 0, ..]);
embedding.assign(&Array1::from_iter(embedding_slice.iter().cloned()));
Ok(normalize_vector(&embedding))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{ModelManager, BuiltinModel};
use tokenizers::Tokenizer;
use ort::session::Session;
const TEST_CASES: &[(&str, usize)] = &[
("Hello, world!", 3),
("This is a test", 4),
("", 0),
("A", 1),
("A B C", 3),
];
struct TestEmbedding {
session: Option<Session>,
tokenizer: Option<Tokenizer>,
}
impl TextEmbedding for TestEmbedding {
fn tokenizer(&self) -> Option<&Tokenizer> {
self.tokenizer.as_ref()
}
fn session(&self) -> Option<&Session> {
self.session.as_ref()
}
}
#[tokio::test]
async fn test_embedding() -> Result<(), Box<dyn std::error::Error>> {
let manager = ModelManager::new_default()?;
let model = BuiltinModel::MiniLM;
if !manager.is_model_downloaded(model) {
manager.download_model(model).await?;
}
let model_path = manager.get_model_path(model);
let tokenizer_path = manager.get_tokenizer_path(model);
let session = Session::builder()
.map_err(|e| ClassifierError::ModelError(e.to_string()))?
.commit_from_file(&model_path)
.map_err(|e| ClassifierError::ModelError(e.to_string()))?;
let tokenizer = Tokenizer::from_file(&tokenizer_path)
.map_err(|e| ClassifierError::TokenizerError(e.to_string()))?;
let embedding = TestEmbedding {
session: Some(session),
tokenizer: Some(tokenizer),
};
for (text, _) in TEST_CASES.iter() {
let result = embedding.tokenize(text);
assert!(result.is_ok());
}
Ok(())
}
mod token_validation {
use super::*;
#[test]
fn test_rejects_missing_tokenizer() {
let embedding = TestEmbedding {
session: None,
tokenizer: None,
};
let result = embedding.tokenize("test text");
assert!(matches!(result, Err(ClassifierError::TokenizerError(_))));
}
#[test]
fn test_validates_sequence_length() {
let embedding = TestEmbedding {
session: None,
tokenizer: None,
};
let long_text = "this is a very long text that should be rejected";
let result = embedding.tokenize(long_text);
assert!(matches!(result, Err(ClassifierError::TokenizerError(_))));
}
}
mod embedding_generation {
use super::*;
#[tokio::test]
async fn test_embedding() -> Result<(), Box<dyn std::error::Error>> {
let manager = ModelManager::new_default()?;
let model = BuiltinModel::MiniLM;
if !manager.is_model_downloaded(model) {
manager.download_model(model).await?;
}
let model_path = manager.get_model_path(model);
let tokenizer_path = manager.get_tokenizer_path(model);
let session = Session::builder()
.map_err(|e| ClassifierError::ModelError(e.to_string()))?
.commit_from_file(&model_path)
.map_err(|e| ClassifierError::ModelError(e.to_string()))?;
let tokenizer = Tokenizer::from_file(&tokenizer_path)
.map_err(|e| ClassifierError::TokenizerError(e.to_string()))?;
let embedding = TestEmbedding {
session: Some(session),
tokenizer: Some(tokenizer),
};
for (text, _) in TEST_CASES.iter() {
let result = embedding.tokenize(text);
assert!(result.is_ok());
}
Ok(())
}
}
}