Thicket/thicket/embedder.py

80 lines
2.9 KiB
Python

"""Local embedding engine — FastEmbed wrapper.
FastEmbed runs ONNX models fully locally (no API keys, no network
after the first model download). The catalog maps each model to its
vector dimension so the Qdrant collection is created with matching
geometry. BGE models want a short instruction prefix on *query* side
only — passages are embedded bare.
"""
from __future__ import annotations
# Curated catalog for technical corpora — code/config-heavy
# libraries want jina-code (English+code, 8k context); bge-base is the
# prose-strong alternative; bge-small for light setups.
# model name -> embedding dimension
EMBEDDING_MODELS: dict[str, int] = {
"jinaai/jina-embeddings-v2-base-code": 768,
"BAAI/bge-base-en-v1.5": 768,
"BAAI/bge-small-en-v1.5": 384,
}
DEFAULT_EMBED_MODEL = "jinaai/jina-embeddings-v2-base-code"
# BGE retrieval instruction — prefix for QUERIES only, never passages.
BGE_QUERY_PREFIX = "Represent this sentence for searching relevant passages: "
class EmbedderUnavailable(Exception):
"""Raised when fastembed is missing or the model cannot load."""
def model_dim(model_name: str) -> int:
"""Dimension for a catalog model (0 if unknown — Qdrant will tell us)."""
return EMBEDDING_MODELS.get(model_name, 0)
class EmbeddingEngine:
"""Lazy-loading FastEmbed engine. Construct anywhere; call load()
from the worker thread (first call may download the model)."""
def __init__(self, model_name: str = DEFAULT_EMBED_MODEL):
self.model_name = model_name
self._model = None
@property
def dim(self) -> int:
return EMBEDDING_MODELS.get(self.model_name, 0)
def load(self) -> None:
"""Import fastembed and load the model. Safe to call twice."""
if self._model is not None:
return
try:
from fastembed import TextEmbedding
except ImportError as e:
raise EmbedderUnavailable(
"fastembed not installed — run: pip install 'thicket[ingest]'"
) from e
try:
self._model = TextEmbedding(model_name=self.model_name)
except Exception as e:
raise EmbedderUnavailable(
f"failed to load embedding model '{self.model_name}': {e}"
) from e
def embed(self, texts: list[str]) -> list[list[float]]:
"""Embed passages (no instruction prefix). Order preserved.
Values are normalized to plain Python floats: FastEmbed yields
numpy scalars, which some targets' validators reject."""
self.load()
return [[float(x) for x in vec] for vec in self._model.embed(texts)]
def embed_query(self, text: str) -> list[float]:
"""Embed a retrieval query with the BGE instruction prefix
when the selected model is from the BGE family."""
self.load()
query = BGE_QUERY_PREFIX + text if "bge" in self.model_name.lower() else text
return [float(x) for x in next(iter(self._model.embed([query])))]