mirror of
https://github.com/ruvnet/RuView
synced 2026-08-02 19:11:46 +00:00
Merge commit 'd803bfe2b1fe7f5e219e50ac20d6801a0a58ac75' as 'vendor/ruvector'
This commit is contained in:
@@ -0,0 +1,152 @@
|
||||
//! Thread-safe model caching with lazy loading
|
||||
|
||||
use dashmap::DashMap;
|
||||
use fastembed::{EmbeddingModel as FastEmbedModel, InitOptions, TextEmbedding};
|
||||
use parking_lot::RwLock;
|
||||
|
||||
use super::models::EmbeddingModel;
|
||||
|
||||
/// Global model cache for lazy loading and reuse
|
||||
pub struct ModelCache {
|
||||
/// Cached embedding models (using RwLock for interior mutability)
|
||||
models: DashMap<EmbeddingModel, RwLock<TextEmbedding>>,
|
||||
/// Default model setting
|
||||
default_model: RwLock<EmbeddingModel>,
|
||||
}
|
||||
|
||||
impl ModelCache {
|
||||
/// Create a new model cache
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
models: DashMap::new(),
|
||||
default_model: RwLock::new(EmbeddingModel::default()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get or load a model and generate embeddings
|
||||
pub fn embed(&self, model: EmbeddingModel, texts: Vec<&str>) -> Result<Vec<Vec<f32>>, String> {
|
||||
// Check if already cached
|
||||
if let Some(cached) = self.models.get(&model) {
|
||||
let mut embedding = cached.write();
|
||||
return embedding
|
||||
.embed(texts, None)
|
||||
.map_err(|e| format!("Embedding failed: {}", e));
|
||||
}
|
||||
|
||||
// Load the model
|
||||
let embedding = self.load_model(model)?;
|
||||
|
||||
// Generate embeddings first
|
||||
let mut embedding_model = embedding;
|
||||
let result = embedding_model
|
||||
.embed(texts, None)
|
||||
.map_err(|e| format!("Embedding failed: {}", e));
|
||||
|
||||
// Cache the model
|
||||
self.models.insert(model, RwLock::new(embedding_model));
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// Load a model from fastembed
|
||||
fn load_model(&self, model: EmbeddingModel) -> Result<TextEmbedding, String> {
|
||||
let fastembed_model = match model {
|
||||
EmbeddingModel::AllMiniLmL6V2 => FastEmbedModel::AllMiniLML6V2,
|
||||
EmbeddingModel::BgeSmallEnV15 => FastEmbedModel::BGESmallENV15,
|
||||
EmbeddingModel::BgeBaseEnV15 => FastEmbedModel::BGEBaseENV15,
|
||||
EmbeddingModel::BgeLargeEnV15 => FastEmbedModel::BGELargeENV15,
|
||||
EmbeddingModel::AllMpnetBaseV2 => FastEmbedModel::AllMiniLML6V2, // Fallback
|
||||
EmbeddingModel::NomicEmbedTextV15 => FastEmbedModel::NomicEmbedTextV15,
|
||||
};
|
||||
|
||||
let options = InitOptions::new(fastembed_model).with_show_download_progress(false);
|
||||
|
||||
TextEmbedding::try_new(options)
|
||||
.map_err(|e| format!("Failed to load model '{}': {}", model.name(), e))
|
||||
}
|
||||
|
||||
/// Pre-load a model into the cache
|
||||
pub fn preload(&self, model: EmbeddingModel) -> Result<(), String> {
|
||||
if self.models.contains_key(&model) {
|
||||
return Ok(());
|
||||
}
|
||||
let embedding = self.load_model(model)?;
|
||||
self.models.insert(model, RwLock::new(embedding));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Check if a model is loaded
|
||||
pub fn is_loaded(&self, model: EmbeddingModel) -> bool {
|
||||
self.models.contains_key(&model)
|
||||
}
|
||||
|
||||
/// Get list of loaded models
|
||||
pub fn loaded_models(&self) -> Vec<EmbeddingModel> {
|
||||
self.models.iter().map(|r| *r.key()).collect()
|
||||
}
|
||||
|
||||
/// Unload a model from cache
|
||||
pub fn unload(&self, model: EmbeddingModel) -> bool {
|
||||
self.models.remove(&model).is_some()
|
||||
}
|
||||
|
||||
/// Clear all cached models
|
||||
pub fn clear(&self) {
|
||||
self.models.clear();
|
||||
}
|
||||
|
||||
/// Get the default model
|
||||
pub fn default_model(&self) -> EmbeddingModel {
|
||||
*self.default_model.read()
|
||||
}
|
||||
|
||||
/// Set the default model
|
||||
pub fn set_default_model(&self, model: EmbeddingModel) {
|
||||
*self.default_model.write() = model;
|
||||
}
|
||||
|
||||
/// Get memory usage estimate in bytes
|
||||
pub fn estimated_memory_usage(&self) -> usize {
|
||||
self.models
|
||||
.iter()
|
||||
.map(|r| r.key().memory_mb() * 1024 * 1024)
|
||||
.sum()
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for ModelCache {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
// Global singleton cache
|
||||
lazy_static::lazy_static! {
|
||||
pub static ref GLOBAL_CACHE: ModelCache = ModelCache::new();
|
||||
}
|
||||
|
||||
/// Get the global model cache
|
||||
pub fn global_cache() -> &'static ModelCache {
|
||||
&GLOBAL_CACHE
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_cache_creation() {
|
||||
let cache = ModelCache::new();
|
||||
assert!(!cache.is_loaded(EmbeddingModel::AllMiniLmL6V2));
|
||||
assert!(cache.loaded_models().is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_model() {
|
||||
let cache = ModelCache::new();
|
||||
assert_eq!(cache.default_model(), EmbeddingModel::AllMiniLmL6V2);
|
||||
|
||||
cache.set_default_model(EmbeddingModel::BgeSmallEnV15);
|
||||
assert_eq!(cache.default_model(), EmbeddingModel::BgeSmallEnV15);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,368 @@
|
||||
//! SQL function implementations for embedding generation
|
||||
|
||||
use pgrx::prelude::*;
|
||||
|
||||
use super::cache::global_cache;
|
||||
use super::models::{EmbeddingModel, ModelInfo};
|
||||
use super::{MAX_BATCH_SIZE, MAX_TEXT_LENGTH};
|
||||
|
||||
// ============================================================================
|
||||
// Core Embedding Functions
|
||||
// ============================================================================
|
||||
|
||||
/// Generate an embedding vector from text
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `text` - The text to embed
|
||||
/// * `model_name` - Optional model name (defaults to 'all-MiniLM-L6-v2')
|
||||
///
|
||||
/// # Returns
|
||||
/// A vector of f32 values representing the text embedding
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_embed('Hello world');
|
||||
/// SELECT ruvector_embed('Hello world', 'bge-small');
|
||||
/// ```
|
||||
#[pg_extern(immutable, parallel_safe)]
|
||||
pub fn ruvector_embed(text: &str, model_name: default!(&str, "'all-MiniLM-L6-v2'")) -> Vec<f32> {
|
||||
// Validate text length
|
||||
if text.len() > MAX_TEXT_LENGTH {
|
||||
pgrx::error!(
|
||||
"Text length {} exceeds maximum {} characters",
|
||||
text.len(),
|
||||
MAX_TEXT_LENGTH
|
||||
);
|
||||
}
|
||||
|
||||
// Parse model name
|
||||
let model = EmbeddingModel::from_name(model_name).unwrap_or_else(|| {
|
||||
pgrx::warning!("Unknown model '{}', using default", model_name);
|
||||
EmbeddingModel::default()
|
||||
});
|
||||
|
||||
// Generate embedding using cached model
|
||||
let documents = vec![text];
|
||||
match global_cache().embed(model, documents) {
|
||||
Ok(embeddings) => {
|
||||
if let Some(embedding) = embeddings.into_iter().next() {
|
||||
embedding
|
||||
} else {
|
||||
pgrx::error!("No embedding generated");
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
pgrx::error!("Embedding generation failed: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Generate embeddings for multiple texts in batch
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `texts` - Array of texts to embed
|
||||
/// * `model_name` - Optional model name (defaults to 'all-MiniLM-L6-v2')
|
||||
///
|
||||
/// # Returns
|
||||
/// A 2D array of embeddings (one per input text)
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_embed_batch(ARRAY['Hello', 'World', 'Test']);
|
||||
/// ```
|
||||
#[pg_extern(immutable, parallel_safe)]
|
||||
pub fn ruvector_embed_batch(
|
||||
texts: Vec<String>,
|
||||
model_name: default!(&str, "'all-MiniLM-L6-v2'"),
|
||||
) -> Vec<Vec<f32>> {
|
||||
// Validate batch size
|
||||
if texts.len() > MAX_BATCH_SIZE {
|
||||
pgrx::error!(
|
||||
"Batch size {} exceeds maximum {}",
|
||||
texts.len(),
|
||||
MAX_BATCH_SIZE
|
||||
);
|
||||
}
|
||||
|
||||
// Validate text lengths
|
||||
for (i, text) in texts.iter().enumerate() {
|
||||
if text.len() > MAX_TEXT_LENGTH {
|
||||
pgrx::error!(
|
||||
"Text at index {} exceeds maximum {} characters",
|
||||
i,
|
||||
MAX_TEXT_LENGTH
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Parse model name
|
||||
let model = EmbeddingModel::from_name(model_name).unwrap_or_else(|| {
|
||||
pgrx::warning!("Unknown model '{}', using default", model_name);
|
||||
EmbeddingModel::default()
|
||||
});
|
||||
|
||||
// Generate embeddings using cached model
|
||||
let documents: Vec<&str> = texts.iter().map(|s| s.as_str()).collect();
|
||||
match global_cache().embed(model, documents) {
|
||||
Ok(embeddings) => embeddings,
|
||||
Err(e) => {
|
||||
pgrx::error!("Batch embedding generation failed: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Model Management Functions
|
||||
// ============================================================================
|
||||
|
||||
/// List all available embedding models
|
||||
///
|
||||
/// # Returns
|
||||
/// Table with model name, dimensions, and description
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT * FROM ruvector_embedding_models();
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_embedding_models() -> TableIterator<
|
||||
'static,
|
||||
(
|
||||
name!(name, String),
|
||||
name!(dimensions, i32),
|
||||
name!(description, String),
|
||||
name!(speed, i32),
|
||||
name!(quality, i32),
|
||||
name!(memory_mb, i32),
|
||||
name!(loaded, bool),
|
||||
),
|
||||
> {
|
||||
let cache = global_cache();
|
||||
let rows: Vec<_> = EmbeddingModel::all()
|
||||
.iter()
|
||||
.map(|model| {
|
||||
(
|
||||
model.name().to_string(),
|
||||
model.dimensions() as i32,
|
||||
model.description().to_string(),
|
||||
model.speed_rating() as i32,
|
||||
model.quality_rating() as i32,
|
||||
model.memory_mb() as i32,
|
||||
cache.is_loaded(*model),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
TableIterator::new(rows)
|
||||
}
|
||||
|
||||
/// Pre-load an embedding model into cache
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model_name` - Name of the model to load
|
||||
///
|
||||
/// # Returns
|
||||
/// true if loaded successfully
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_load_model('bge-small');
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_load_model(model_name: &str) -> bool {
|
||||
let model = match EmbeddingModel::from_name(model_name) {
|
||||
Some(m) => m,
|
||||
None => {
|
||||
pgrx::warning!("Unknown model: {}", model_name);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
match global_cache().preload(model) {
|
||||
Ok(_) => {
|
||||
pgrx::info!("Model '{}' loaded successfully", model.name());
|
||||
true
|
||||
}
|
||||
Err(e) => {
|
||||
pgrx::warning!("Failed to load model '{}': {}", model_name, e);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Unload an embedding model from cache
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model_name` - Name of the model to unload
|
||||
///
|
||||
/// # Returns
|
||||
/// true if the model was unloaded
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_unload_model('bge-small');
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_unload_model(model_name: &str) -> bool {
|
||||
let model = match EmbeddingModel::from_name(model_name) {
|
||||
Some(m) => m,
|
||||
None => {
|
||||
pgrx::warning!("Unknown model: {}", model_name);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
global_cache().unload(model)
|
||||
}
|
||||
|
||||
/// Get information about a specific model
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model_name` - Name of the model
|
||||
///
|
||||
/// # Returns
|
||||
/// JSON object with model information
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_model_info('all-MiniLM-L6-v2');
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_model_info(model_name: &str) -> pgrx::JsonB {
|
||||
let model = match EmbeddingModel::from_name(model_name) {
|
||||
Some(m) => m,
|
||||
None => {
|
||||
return pgrx::JsonB(serde_json::json!({
|
||||
"error": format!("Unknown model: {}", model_name),
|
||||
"available_models": EmbeddingModel::all().iter().map(|m| m.name()).collect::<Vec<_>>()
|
||||
}));
|
||||
}
|
||||
};
|
||||
|
||||
let cache = global_cache();
|
||||
let mut info = ModelInfo::from(model);
|
||||
info.loaded = cache.is_loaded(model);
|
||||
|
||||
pgrx::JsonB(serde_json::to_value(info).unwrap_or_default())
|
||||
}
|
||||
|
||||
/// Set the default embedding model
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model_name` - Name of the model to set as default
|
||||
///
|
||||
/// # Returns
|
||||
/// true if set successfully
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_set_default_model('bge-small');
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_set_default_model(model_name: &str) -> bool {
|
||||
let model = match EmbeddingModel::from_name(model_name) {
|
||||
Some(m) => m,
|
||||
None => {
|
||||
pgrx::warning!("Unknown model: {}", model_name);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
global_cache().set_default_model(model);
|
||||
pgrx::info!("Default model set to '{}'", model.name());
|
||||
true
|
||||
}
|
||||
|
||||
/// Get the current default embedding model name
|
||||
///
|
||||
/// # Returns
|
||||
/// Name of the default model
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_default_model();
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_default_model() -> String {
|
||||
global_cache().default_model().name().to_string()
|
||||
}
|
||||
|
||||
/// Get embedding cache statistics
|
||||
///
|
||||
/// # Returns
|
||||
/// JSON object with cache statistics
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_embedding_stats();
|
||||
/// ```
|
||||
#[pg_extern]
|
||||
pub fn ruvector_embedding_stats() -> pgrx::JsonB {
|
||||
let cache = global_cache();
|
||||
let loaded_models = cache.loaded_models();
|
||||
|
||||
pgrx::JsonB(serde_json::json!({
|
||||
"loaded_model_count": loaded_models.len(),
|
||||
"loaded_models": loaded_models.iter().map(|m| m.name()).collect::<Vec<_>>(),
|
||||
"estimated_memory_mb": cache.estimated_memory_usage() / (1024 * 1024),
|
||||
"default_model": cache.default_model().name(),
|
||||
"available_model_count": EmbeddingModel::all().len(),
|
||||
}))
|
||||
}
|
||||
|
||||
/// Get embedding dimensions for a model
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model_name` - Name of the model
|
||||
///
|
||||
/// # Returns
|
||||
/// Number of dimensions, or -1 if model unknown
|
||||
///
|
||||
/// # Example
|
||||
/// ```sql
|
||||
/// SELECT ruvector_embedding_dims('all-MiniLM-L6-v2'); -- Returns 384
|
||||
/// ```
|
||||
#[pg_extern(immutable, parallel_safe)]
|
||||
pub fn ruvector_embedding_dims(model_name: &str) -> i32 {
|
||||
match EmbeddingModel::from_name(model_name) {
|
||||
Some(m) => m.dimensions() as i32,
|
||||
None => -1,
|
||||
}
|
||||
}
|
||||
|
||||
// ============================================================================
|
||||
// Tests
|
||||
// ============================================================================
|
||||
|
||||
#[cfg(feature = "pg_test")]
|
||||
#[pg_schema]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[pg_test]
|
||||
fn test_embedding_models_list() {
|
||||
let models: Vec<_> = ruvector_embedding_models().collect();
|
||||
assert!(!models.is_empty());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_model_info() {
|
||||
let info = ruvector_model_info("all-MiniLM-L6-v2");
|
||||
let json = info.0;
|
||||
assert!(json.get("name").is_some());
|
||||
assert!(json.get("dimensions").is_some());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_default_model() {
|
||||
let name = ruvector_default_model();
|
||||
assert!(!name.is_empty());
|
||||
}
|
||||
|
||||
#[pg_test]
|
||||
fn test_embedding_dims() {
|
||||
assert_eq!(ruvector_embedding_dims("all-MiniLM-L6-v2"), 384);
|
||||
assert_eq!(ruvector_embedding_dims("bge-base"), 768);
|
||||
assert_eq!(ruvector_embedding_dims("unknown"), -1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
//! Local embedding generation module for ruvector-postgres
|
||||
//!
|
||||
//! Provides text-to-embedding functionality using fastembed-rs (ONNX-based models).
|
||||
//! Supports multiple embedding models with lazy loading and thread-safe caching.
|
||||
//!
|
||||
//! # Features
|
||||
//!
|
||||
//! - Local embedding generation (no external API calls)
|
||||
//! - Multiple model support (MiniLM, BGE, MPNet)
|
||||
//! - Lazy model loading (loads on first use)
|
||||
//! - Thread-safe model cache
|
||||
//! - Batch embedding for efficiency
|
||||
//!
|
||||
//! # SQL Functions
|
||||
//!
|
||||
//! ```sql
|
||||
//! -- Generate embedding from text
|
||||
//! SELECT ruvector_embed('Hello world');
|
||||
//! SELECT ruvector_embed('Hello world', 'all-MiniLM-L6-v2');
|
||||
//!
|
||||
//! -- Batch embedding
|
||||
//! SELECT ruvector_embed_batch(ARRAY['text1', 'text2']);
|
||||
//!
|
||||
//! -- List available models
|
||||
//! SELECT * FROM ruvector_embedding_models();
|
||||
//!
|
||||
//! -- Model management
|
||||
//! SELECT ruvector_load_model('all-MiniLM-L6-v2');
|
||||
//! SELECT ruvector_model_info('all-MiniLM-L6-v2');
|
||||
//! ```
|
||||
|
||||
mod cache;
|
||||
mod functions;
|
||||
mod models;
|
||||
|
||||
pub use cache::ModelCache;
|
||||
pub use functions::*;
|
||||
pub use models::{EmbeddingModel, ModelInfo};
|
||||
|
||||
/// Default embedding model
|
||||
pub const DEFAULT_MODEL: &str = "all-MiniLM-L6-v2";
|
||||
|
||||
/// Maximum batch size for embedding generation
|
||||
pub const MAX_BATCH_SIZE: usize = 256;
|
||||
|
||||
/// Maximum text length (in characters) for embedding
|
||||
pub const MAX_TEXT_LENGTH: usize = 8192;
|
||||
@@ -0,0 +1,207 @@
|
||||
//! Embedding model definitions and metadata
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Supported embedding models
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub enum EmbeddingModel {
|
||||
/// all-MiniLM-L6-v2: Fast, good quality, 384 dimensions
|
||||
AllMiniLmL6V2,
|
||||
/// BAAI/bge-small-en-v1.5: Fast, high quality, 384 dimensions
|
||||
BgeSmallEnV15,
|
||||
/// BAAI/bge-base-en-v1.5: Medium speed, higher quality, 768 dimensions
|
||||
BgeBaseEnV15,
|
||||
/// sentence-transformers/all-mpnet-base-v2: Medium speed, high quality, 768 dimensions
|
||||
AllMpnetBaseV2,
|
||||
/// nomic-ai/nomic-embed-text-v1.5: Good quality, 768 dimensions
|
||||
NomicEmbedTextV15,
|
||||
/// BAAI/bge-large-en-v1.5: Slower, highest quality, 1024 dimensions
|
||||
BgeLargeEnV15,
|
||||
}
|
||||
|
||||
impl EmbeddingModel {
|
||||
/// Parse model name string to enum
|
||||
pub fn from_name(name: &str) -> Option<Self> {
|
||||
match name.to_lowercase().as_str() {
|
||||
"all-minilm-l6-v2" | "minilm" | "default" => Some(Self::AllMiniLmL6V2),
|
||||
"bge-small-en-v1.5" | "bge-small" | "baai/bge-small-en-v1.5" => {
|
||||
Some(Self::BgeSmallEnV15)
|
||||
}
|
||||
"bge-base-en-v1.5" | "bge-base" | "baai/bge-base-en-v1.5" => Some(Self::BgeBaseEnV15),
|
||||
"bge-large-en-v1.5" | "bge-large" | "baai/bge-large-en-v1.5" => {
|
||||
Some(Self::BgeLargeEnV15)
|
||||
}
|
||||
"all-mpnet-base-v2" | "mpnet" | "sentence-transformers/all-mpnet-base-v2" => {
|
||||
Some(Self::AllMpnetBaseV2)
|
||||
}
|
||||
"nomic-embed-text-v1.5" | "nomic" | "nomic-ai/nomic-embed-text-v1.5" => {
|
||||
Some(Self::NomicEmbedTextV15)
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the canonical name for this model
|
||||
pub fn name(&self) -> &'static str {
|
||||
match self {
|
||||
Self::AllMiniLmL6V2 => "all-MiniLM-L6-v2",
|
||||
Self::BgeSmallEnV15 => "BAAI/bge-small-en-v1.5",
|
||||
Self::BgeBaseEnV15 => "BAAI/bge-base-en-v1.5",
|
||||
Self::BgeLargeEnV15 => "BAAI/bge-large-en-v1.5",
|
||||
Self::AllMpnetBaseV2 => "sentence-transformers/all-mpnet-base-v2",
|
||||
Self::NomicEmbedTextV15 => "nomic-ai/nomic-embed-text-v1.5",
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the embedding dimensions for this model
|
||||
pub fn dimensions(&self) -> usize {
|
||||
match self {
|
||||
Self::AllMiniLmL6V2 => 384,
|
||||
Self::BgeSmallEnV15 => 384,
|
||||
Self::BgeBaseEnV15 => 768,
|
||||
Self::BgeLargeEnV15 => 1024,
|
||||
Self::AllMpnetBaseV2 => 768,
|
||||
Self::NomicEmbedTextV15 => 768,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get a description of this model
|
||||
pub fn description(&self) -> &'static str {
|
||||
match self {
|
||||
Self::AllMiniLmL6V2 => "Fast general-purpose model, good for most use cases",
|
||||
Self::BgeSmallEnV15 => "High quality small model from BAAI, great for semantic search",
|
||||
Self::BgeBaseEnV15 => "Higher quality base model from BAAI, better accuracy",
|
||||
Self::BgeLargeEnV15 => "Highest quality large model from BAAI, best accuracy",
|
||||
Self::AllMpnetBaseV2 => "High quality model from sentence-transformers",
|
||||
Self::NomicEmbedTextV15 => "Modern model with good quality from Nomic AI",
|
||||
}
|
||||
}
|
||||
|
||||
/// Get model speed rating (1-5, higher is faster)
|
||||
pub fn speed_rating(&self) -> u8 {
|
||||
match self {
|
||||
Self::AllMiniLmL6V2 => 5,
|
||||
Self::BgeSmallEnV15 => 5,
|
||||
Self::BgeBaseEnV15 => 3,
|
||||
Self::BgeLargeEnV15 => 1,
|
||||
Self::AllMpnetBaseV2 => 3,
|
||||
Self::NomicEmbedTextV15 => 3,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get model quality rating (1-5, higher is better)
|
||||
pub fn quality_rating(&self) -> u8 {
|
||||
match self {
|
||||
Self::AllMiniLmL6V2 => 3,
|
||||
Self::BgeSmallEnV15 => 4,
|
||||
Self::BgeBaseEnV15 => 4,
|
||||
Self::BgeLargeEnV15 => 5,
|
||||
Self::AllMpnetBaseV2 => 4,
|
||||
Self::NomicEmbedTextV15 => 4,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get approximate memory usage in MB
|
||||
pub fn memory_mb(&self) -> usize {
|
||||
match self {
|
||||
Self::AllMiniLmL6V2 => 90,
|
||||
Self::BgeSmallEnV15 => 130,
|
||||
Self::BgeBaseEnV15 => 440,
|
||||
Self::BgeLargeEnV15 => 1340,
|
||||
Self::AllMpnetBaseV2 => 440,
|
||||
Self::NomicEmbedTextV15 => 550,
|
||||
}
|
||||
}
|
||||
|
||||
/// Get all supported models
|
||||
pub fn all() -> &'static [EmbeddingModel] {
|
||||
&[
|
||||
Self::AllMiniLmL6V2,
|
||||
Self::BgeSmallEnV15,
|
||||
Self::BgeBaseEnV15,
|
||||
Self::BgeLargeEnV15,
|
||||
Self::AllMpnetBaseV2,
|
||||
Self::NomicEmbedTextV15,
|
||||
]
|
||||
}
|
||||
|
||||
/// Get the default model
|
||||
pub fn default_model() -> Self {
|
||||
Self::AllMiniLmL6V2
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for EmbeddingModel {
|
||||
fn default() -> Self {
|
||||
Self::default_model()
|
||||
}
|
||||
}
|
||||
|
||||
/// Model information for SQL queries
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ModelInfo {
|
||||
pub name: String,
|
||||
pub dimensions: i32,
|
||||
pub description: String,
|
||||
pub speed_rating: i32,
|
||||
pub quality_rating: i32,
|
||||
pub memory_mb: i32,
|
||||
pub loaded: bool,
|
||||
}
|
||||
|
||||
impl From<EmbeddingModel> for ModelInfo {
|
||||
fn from(model: EmbeddingModel) -> Self {
|
||||
Self {
|
||||
name: model.name().to_string(),
|
||||
dimensions: model.dimensions() as i32,
|
||||
description: model.description().to_string(),
|
||||
speed_rating: model.speed_rating() as i32,
|
||||
quality_rating: model.quality_rating() as i32,
|
||||
memory_mb: model.memory_mb() as i32,
|
||||
loaded: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_model_parsing() {
|
||||
assert_eq!(
|
||||
EmbeddingModel::from_name("all-minilm-l6-v2"),
|
||||
Some(EmbeddingModel::AllMiniLmL6V2)
|
||||
);
|
||||
assert_eq!(
|
||||
EmbeddingModel::from_name("minilm"),
|
||||
Some(EmbeddingModel::AllMiniLmL6V2)
|
||||
);
|
||||
assert_eq!(
|
||||
EmbeddingModel::from_name("default"),
|
||||
Some(EmbeddingModel::AllMiniLmL6V2)
|
||||
);
|
||||
assert_eq!(
|
||||
EmbeddingModel::from_name("bge-small"),
|
||||
Some(EmbeddingModel::BgeSmallEnV15)
|
||||
);
|
||||
assert_eq!(EmbeddingModel::from_name("unknown"), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_model_dimensions() {
|
||||
assert_eq!(EmbeddingModel::AllMiniLmL6V2.dimensions(), 384);
|
||||
assert_eq!(EmbeddingModel::BgeBaseEnV15.dimensions(), 768);
|
||||
assert_eq!(EmbeddingModel::BgeLargeEnV15.dimensions(), 1024);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_all_models() {
|
||||
let models = EmbeddingModel::all();
|
||||
assert!(models.len() >= 4);
|
||||
for model in models {
|
||||
assert!(!model.name().is_empty());
|
||||
assert!(model.dimensions() > 0);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user