101 lines
3.3 KiB
Python
101 lines
3.3 KiB
Python
"""
|
|
Embeddings abstraction for the memory system.
|
|
|
|
Provides an interface for generating embeddings with different backends.
|
|
|
|
IMPORTANT: All embeddings must produce 384-dimensional vectors to match
|
|
the database schema (pgvector column defined as vector(384)).
|
|
"""
|
|
from abc import ABC, abstractmethod
|
|
from typing import List
|
|
import logging
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# Fixed embedding dimension required by database schema
|
|
EMBEDDING_DIMENSION = 384
|
|
|
|
|
|
class Embeddings(ABC):
|
|
"""
|
|
Abstract base class for embedding generation.
|
|
|
|
All implementations MUST generate 384-dimensional embeddings to match
|
|
the database schema.
|
|
"""
|
|
|
|
@abstractmethod
|
|
def encode(self, texts: List[str]) -> List[List[float]]:
|
|
"""
|
|
Generate 384-dimensional embeddings for a list of texts.
|
|
|
|
Args:
|
|
texts: List of text strings to encode
|
|
|
|
Returns:
|
|
List of 384-dimensional embedding vectors (each is a list of floats)
|
|
"""
|
|
pass
|
|
|
|
|
|
class SentenceTransformersEmbeddings(Embeddings):
|
|
"""
|
|
Embeddings implementation using SentenceTransformers.
|
|
|
|
Uses lazy import so sentence-transformers is not required if another
|
|
embedding backend is used.
|
|
|
|
Default model is BAAI/bge-small-en-v1.5 which produces 384-dimensional
|
|
embeddings matching the database schema.
|
|
"""
|
|
|
|
def __init__(self, model_name: str = "BAAI/bge-small-en-v1.5"):
|
|
"""
|
|
Initialize SentenceTransformers embeddings.
|
|
|
|
Args:
|
|
model_name: Name of the SentenceTransformer model to use.
|
|
Must produce 384-dimensional embeddings.
|
|
Default: BAAI/bge-small-en-v1.5
|
|
"""
|
|
self.model_name = model_name
|
|
self._model = None
|
|
self._load_model()
|
|
|
|
def _load_model(self):
|
|
"""Lazy load and validate the SentenceTransformer model."""
|
|
if self._model is None:
|
|
try:
|
|
from sentence_transformers import SentenceTransformer
|
|
except ImportError:
|
|
raise ImportError(
|
|
"sentence-transformers is required for SentenceTransformersEmbeddings. "
|
|
"Install it with: pip install sentence-transformers"
|
|
)
|
|
|
|
logger.info(f"Loading embedding model: {self.model_name}...")
|
|
self._model = SentenceTransformer(self.model_name)
|
|
|
|
# Validate dimension matches database schema
|
|
model_dim = self._model.get_sentence_embedding_dimension()
|
|
if model_dim != EMBEDDING_DIMENSION:
|
|
raise ValueError(
|
|
f"Model {self.model_name} produces {model_dim}-dimensional embeddings, "
|
|
f"but database schema requires {EMBEDDING_DIMENSION} dimensions. "
|
|
f"Use a model that produces {EMBEDDING_DIMENSION}-dimensional embeddings."
|
|
)
|
|
|
|
logger.info(f"Model loaded (embedding dim: {model_dim})")
|
|
|
|
def encode(self, texts: List[str]) -> List[List[float]]:
|
|
"""
|
|
Generate 384-dimensional embeddings for a list of texts.
|
|
|
|
Args:
|
|
texts: List of text strings to encode
|
|
|
|
Returns:
|
|
List of 384-dimensional embedding vectors
|
|
"""
|
|
embeddings = self._model.encode(texts, convert_to_numpy=True, show_progress_bar=False)
|
|
return [emb.tolist() for emb in embeddings]
|