use std::path::Path;
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
use candle_transformers::models::bert::{BertModel, Config};
use tokenizers::Tokenizer;
#[derive(Debug, thiserror::Error)]
pub enum LocalModelError {
#[error("model io error: {0}")]
Io(#[from] std::io::Error),
#[error("model config error: {0}")]
Config(#[from] serde_json::Error),
#[error("tokenizer error: {0}")]
Tokenizer(String),
#[error("candle error: {0}")]
Candle(#[from] candle_core::Error),
#[error("unsupported model architecture: {0}")]
UnsupportedArch(String),
}
pub struct LocalEmbedder {
model: BertModel,
tokenizer: Tokenizer,
device: Device,
dim: usize,
}
impl LocalEmbedder {
pub fn load(model_dir: &Path) -> Result<Self, LocalModelError> {
let device = Device::Cpu;
let config_bytes = std::fs::read(model_dir.join("config.json"))?;
let config: Config = serde_json::from_slice(&config_bytes)?;
let dim = config.hidden_size;
let tokenizer = Tokenizer::from_file(model_dir.join("tokenizer.json"))
.map_err(|e| LocalModelError::Tokenizer(e.to_string()))?;
let weights = std::fs::read(model_dir.join("model.safetensors"))?;
let vb = VarBuilder::from_buffered_safetensors(weights, DType::F32, &device)?;
let model = BertModel::load(vb, &config)?;
Ok(Self {
model,
tokenizer,
device,
dim,
})
}
#[must_use]
pub fn dim(&self) -> usize {
self.dim
}
pub fn embed(&self, text: &str) -> Result<Vec<f32>, LocalModelError> {
let encoding = self
.tokenizer
.encode(text, true)
.map_err(|e| LocalModelError::Tokenizer(e.to_string()))?;
let ids: Vec<u32> = encoding.get_ids().to_vec();
let mask: Vec<u32> = encoding.get_attention_mask().to_vec();
let input_ids = Tensor::new(ids.as_slice(), &self.device)?.unsqueeze(0)?;
let token_type_ids = input_ids.zeros_like()?;
let attention_mask = Tensor::new(mask.as_slice(), &self.device)?.unsqueeze(0)?;
let hidden = self
.model
.forward(&input_ids, &token_type_ids, Some(&attention_mask))?;
let mask_f = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; let summed = hidden.broadcast_mul(&mask_f)?.sum(1)?; let counts = mask_f.sum(1)?; let mean = summed.broadcast_div(&counts)?;
let norm = mean.sqr()?.sum_keepdim(1)?.sqrt()?;
let normalized = mean.broadcast_div(&norm)?;
Ok(normalized.squeeze(0)?.to_vec1::<f32>()?)
}
}
#[cfg(test)]
mod tests {
use super::{LocalEmbedder, LocalModelError};
fn tmp(name: &str) -> std::path::PathBuf {
let dir = std::env::temp_dir().join(format!("roteiro-embed-{}-{name}", std::process::id()));
std::fs::remove_dir_all(&dir).ok();
dir
}
#[test]
fn load_missing_dir_is_io_error() {
let dir = tmp("missing");
assert!(matches!(
LocalEmbedder::load(&dir),
Err(LocalModelError::Io(_))
));
}
#[test]
fn load_invalid_config_is_config_error() {
let dir = tmp("badcfg");
std::fs::create_dir_all(&dir).expect("mkdir");
std::fs::write(dir.join("config.json"), b"this is not json").expect("write");
assert!(matches!(
LocalEmbedder::load(&dir),
Err(LocalModelError::Config(_))
));
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn load_missing_tokenizer_after_valid_config_errors() {
let dir = tmp("notokenizer");
std::fs::create_dir_all(&dir).expect("mkdir");
std::fs::write(
dir.join("config.json"),
br#"{"hidden_size":16,"num_hidden_layers":1,"num_attention_heads":1,"intermediate_size":16,"vocab_size":8,"max_position_embeddings":8,"type_vocab_size":2}"#,
)
.expect("write");
assert!(LocalEmbedder::load(&dir).is_err());
std::fs::remove_dir_all(&dir).ok();
}
}