129 lines
3.2 KiB
Python
129 lines
3.2 KiB
Python
"""
|
|
Embedding model client for text vectorization
|
|
|
|
Uses sentence-transformers for generating embeddings.
|
|
"""
|
|
import logging
|
|
from typing import List, Optional
|
|
from sentence_transformers import SentenceTransformer
|
|
from src.config import get_settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
settings = get_settings()
|
|
|
|
|
|
class EmbeddingClient:
|
|
"""Client for generating text embeddings"""
|
|
|
|
def __init__(self, model_name: Optional[str] = None):
|
|
"""
|
|
Initialize embedding client
|
|
|
|
Args:
|
|
model_name: Optional model name, defaults to config
|
|
"""
|
|
self.model_name = model_name or settings.embedding_model
|
|
self.dimension = settings.embedding_dimension
|
|
self._model: Optional[SentenceTransformer] = None
|
|
logger.info(f"Initializing EmbeddingClient with model: {self.model_name}")
|
|
|
|
def _load_model(self) -> SentenceTransformer:
|
|
"""
|
|
Lazy load the embedding model
|
|
|
|
Returns:
|
|
Loaded SentenceTransformer model
|
|
"""
|
|
if self._model is None:
|
|
logger.info(f"Loading embedding model: {self.model_name}")
|
|
self._model = SentenceTransformer(self.model_name)
|
|
logger.info(f"Model loaded successfully. Embedding dimension: {self.dimension}")
|
|
return self._model
|
|
|
|
def embed_text(self, text: str) -> List[float]:
|
|
"""
|
|
Generate embedding for a single text
|
|
|
|
Args:
|
|
text: Input text to embed
|
|
|
|
Returns:
|
|
List of floats representing the embedding vector
|
|
"""
|
|
model = self._load_model()
|
|
embedding = model.encode(text, convert_to_numpy=True)
|
|
return embedding.tolist()
|
|
|
|
def embed_batch(self, texts: List[str]) -> List[List[float]]:
|
|
"""
|
|
Generate embeddings for multiple texts
|
|
|
|
Args:
|
|
texts: List of input texts
|
|
|
|
Returns:
|
|
List of embedding vectors
|
|
"""
|
|
model = self._load_model()
|
|
embeddings = model.encode(
|
|
texts,
|
|
batch_size=settings.embedding_batch_size,
|
|
convert_to_numpy=True,
|
|
show_progress_bar=False
|
|
)
|
|
return embeddings.tolist()
|
|
|
|
def get_dimension(self) -> int:
|
|
"""
|
|
Get embedding dimension
|
|
|
|
Returns:
|
|
Embedding vector dimension
|
|
"""
|
|
return self.dimension
|
|
|
|
|
|
# Global instance
|
|
_embedding_client: Optional[EmbeddingClient] = None
|
|
|
|
|
|
def get_embedding_client() -> EmbeddingClient:
|
|
"""
|
|
Get or create global embedding client instance
|
|
|
|
Returns:
|
|
EmbeddingClient instance
|
|
"""
|
|
global _embedding_client
|
|
if _embedding_client is None:
|
|
_embedding_client = EmbeddingClient()
|
|
return _embedding_client
|
|
|
|
|
|
async def embed_text_async(text: str) -> List[float]:
|
|
"""
|
|
Async wrapper for embedding text
|
|
|
|
Args:
|
|
text: Input text
|
|
|
|
Returns:
|
|
Embedding vector
|
|
"""
|
|
client = get_embedding_client()
|
|
return client.embed_text(text)
|
|
|
|
|
|
async def embed_batch_async(texts: List[str]) -> List[List[float]]:
|
|
"""
|
|
Async wrapper for batch embedding
|
|
|
|
Args:
|
|
texts: List of input texts
|
|
|
|
Returns:
|
|
List of embedding vectors
|
|
"""
|
|
client = get_embedding_client()
|
|
return client.embed_batch(texts)
|