fleet-memory/hindsight-api/hindsight_api/engine/embeddings.py
2025-12-10 12:12:21 +01:00

303 lines
10 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)).
Configuration via environment variables:
- HINDSIGHT_API_EMBEDDINGS_PROVIDER: "local" (default) or "tei"
For local provider:
- HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL: Model name (default: BAAI/bge-small-en-v1.5)
For TEI provider:
- HINDSIGHT_API_EMBEDDINGS_TEI_URL: TEI server URL (required)
"""
from abc import ABC, abstractmethod
from typing import List, Optional
import logging
import os
import httpx
logger = logging.getLogger(__name__)
# Fixed embedding dimension required by database schema
EMBEDDING_DIMENSION = 384
# Default model for local embeddings
DEFAULT_EMBEDDINGS_MODEL = "BAAI/bge-small-en-v1.5"
class Embeddings(ABC):
"""
Abstract base class for embedding generation.
All implementations MUST generate 384-dimensional embeddings to match
the database schema.
"""
@property
@abstractmethod
def provider_name(self) -> str:
"""Return a human-readable name for this provider (e.g., 'local', 'tei')."""
pass
@abstractmethod
async def initialize(self) -> None:
"""
Initialize the embedding model asynchronously.
This should be called during startup to load/connect to the model
and avoid cold start latency on first encode() call.
"""
pass
@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 LocalSTEmbeddings(Embeddings):
"""
Local embeddings implementation using SentenceTransformers.
Call initialize() during startup to load the model and avoid cold starts.
Default model is BAAI/bge-small-en-v1.5 which produces 384-dimensional
embeddings matching the database schema.
"""
def __init__(self, model_name: Optional[str] = None):
"""
Initialize local 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 or DEFAULT_EMBEDDINGS_MODEL
self._model = None
@property
def provider_name(self) -> str:
return "local"
async def initialize(self) -> None:
"""Load the embedding model."""
if self._model is not None:
return
try:
from sentence_transformers import SentenceTransformer
except ImportError:
raise ImportError(
"sentence-transformers is required for LocalSTEmbeddings. "
"Install it with: pip install sentence-transformers"
)
logger.info(f"Embeddings: initializing local provider with model {self.model_name}")
# Disable lazy loading (meta tensors) which causes issues with newer transformers/accelerate
# Setting low_cpu_mem_usage=False and device_map=None ensures tensors are fully materialized
self._model = SentenceTransformer(
self.model_name,
model_kwargs={"low_cpu_mem_usage": False, "device_map": None},
)
# 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"Embeddings: local provider initialized (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
"""
if self._model is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
embeddings = self._model.encode(texts, convert_to_numpy=True, show_progress_bar=False)
return [emb.tolist() for emb in embeddings]
class RemoteTEIEmbeddings(Embeddings):
"""
Remote embeddings implementation using HuggingFace Text Embeddings Inference (TEI) HTTP API.
TEI provides a high-performance inference server for embedding models.
See: https://github.com/huggingface/text-embeddings-inference
The server should be running a model that produces 384-dimensional embeddings.
"""
def __init__(
self,
base_url: str,
timeout: float = 30.0,
batch_size: int = 32,
max_retries: int = 3,
retry_delay: float = 0.5,
):
"""
Initialize remote TEI embeddings client.
Args:
base_url: Base URL of the TEI server (e.g., "http://localhost:8080")
timeout: Request timeout in seconds (default: 30.0)
batch_size: Maximum batch size for embedding requests (default: 32)
max_retries: Maximum number of retries for failed requests (default: 3)
retry_delay: Initial delay between retries in seconds, doubles each retry (default: 0.5)
"""
self.base_url = base_url.rstrip("/")
self.timeout = timeout
self.batch_size = batch_size
self.max_retries = max_retries
self.retry_delay = retry_delay
self._client: Optional[httpx.Client] = None
self._model_id: Optional[str] = None
@property
def provider_name(self) -> str:
return "tei"
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
"""Make an HTTP request with automatic retries on transient errors."""
import time
last_error = None
delay = self.retry_delay
for attempt in range(self.max_retries + 1):
try:
if method == "GET":
response = self._client.get(url, **kwargs)
else:
response = self._client.post(url, **kwargs)
response.raise_for_status()
return response
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
if attempt < self.max_retries:
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
time.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
last_error = e
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
time.sleep(delay)
delay *= 2
else:
raise
raise last_error
async def initialize(self) -> None:
"""Initialize the HTTP client and verify server connectivity."""
if self._client is not None:
return
logger.info(f"Embeddings: initializing TEI provider at {self.base_url}")
self._client = httpx.Client(timeout=self.timeout)
# Verify server is reachable and get model info
try:
response = self._request_with_retry("GET", f"{self.base_url}/info")
info = response.json()
self._model_id = info.get("model_id", "unknown")
logger.info(f"Embeddings: TEI provider initialized (model: {self._model_id})")
except httpx.HTTPError as e:
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
def encode(self, texts: List[str]) -> List[List[float]]:
"""
Generate embeddings using the remote TEI server.
Args:
texts: List of text strings to encode
Returns:
List of embedding vectors
"""
if self._client is None:
raise RuntimeError("Embeddings not initialized. Call initialize() first.")
if not texts:
return []
all_embeddings = []
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i:i + self.batch_size]
try:
response = self._request_with_retry(
"POST",
f"{self.base_url}/embed",
json={"inputs": batch},
)
batch_embeddings = response.json()
all_embeddings.extend(batch_embeddings)
except httpx.HTTPError as e:
raise RuntimeError(f"TEI embedding request failed: {e}")
return all_embeddings
def create_embeddings_from_env() -> Embeddings:
"""
Create an Embeddings instance based on environment variables.
Environment variables:
- HINDSIGHT_API_EMBEDDINGS_PROVIDER: "local" (default) or "tei"
For local provider:
- HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL: Model name (default: BAAI/bge-small-en-v1.5)
For TEI provider:
- HINDSIGHT_API_EMBEDDINGS_TEI_URL: TEI server URL (required)
Returns:
Configured Embeddings instance
"""
provider = os.environ.get("HINDSIGHT_API_EMBEDDINGS_PROVIDER", "local").lower()
if provider == "tei":
url = os.environ.get("HINDSIGHT_API_EMBEDDINGS_TEI_URL")
if not url:
raise ValueError(
"HINDSIGHT_API_EMBEDDINGS_TEI_URL is required when HINDSIGHT_API_EMBEDDINGS_PROVIDER is 'tei'"
)
return RemoteTEIEmbeddings(base_url=url)
elif provider == "local":
model = os.environ.get("HINDSIGHT_API_EMBEDDINGS_LOCAL_MODEL")
model_name = model or DEFAULT_EMBEDDINGS_MODEL
return LocalSTEmbeddings(model_name=model_name)
else:
raise ValueError(
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'"
)