G Davies
Deploy JW Search Docker Space
722bda8 verified
Raw
History Blame Contribute Delete
4.98 kB
"""
Text embedding wrapper for scene descriptions.
Uses BAAI/bge-m3 via sentence-transformers. bge-m3 is multilingual (JW.org
content exists in many languages) and produces 1024-dim dense vectors that
slot into the same sqlite-vec storage pattern the rest of the project uses.
"""
from __future__ import annotations
import math
from dataclasses import dataclass
from typing import Sequence
from utils import log_message
DEFAULT_MODEL_ID = "BAAI/bge-m3"
DEFAULT_BATCH_SIZE = 16
# bge-m3 has an 8192-token context. ~3.5 chars/token average for our
# descriptions, so 24000 chars is a safe per-chunk budget that leaves
# headroom for tokenization edge cases.
SUMMARY_CHUNK_CHAR_BUDGET = 24000
@dataclass
class EmbeddingResult:
"""Embedding for one scene description."""
vector: list[float]
model_id: str
dimension: int
class TextEmbedder:
"""Lazy-loaded sentence-transformers wrapper."""
def __init__(
self,
model_id: str = DEFAULT_MODEL_ID,
*,
batch_size: int = DEFAULT_BATCH_SIZE,
) -> None:
self.model_id = model_id
self.batch_size = batch_size
self._model = None
def _ensure_loaded(self) -> None:
if self._model is not None:
return
try:
from sentence_transformers import SentenceTransformer
except ImportError as exc:
raise RuntimeError(
"sentence-transformers is required for the text embedder."
) from exc
log_message(f"Loading text embedder: {self.model_id}")
self._model = SentenceTransformer(self.model_id)
log_message(f"Text embedder ready: {self.model_id}")
def embed_many(self, texts: Sequence[str]) -> list[EmbeddingResult]:
"""Embed a list of texts. Returns one ``EmbeddingResult`` per input."""
self._ensure_loaded()
if not texts:
return []
vectors = self._model.encode(
list(texts),
batch_size=self.batch_size,
normalize_embeddings=True,
convert_to_numpy=True,
show_progress_bar=False,
)
return [
EmbeddingResult(
vector=vector.tolist(),
model_id=self.model_id,
dimension=int(vector.shape[0]),
)
for vector in vectors
]
def chunk_pieces_for_embedding(
pieces: Sequence[str],
*,
char_budget: int = SUMMARY_CHUNK_CHAR_BUDGET,
) -> list[str]:
"""Group text pieces into newline-joined chunks under a character budget.
Used to keep summary embeddings from silently truncating when the joined
text exceeds the embedder's context window. Each output chunk is the
newline-joined text of a contiguous slice of ``pieces`` whose total
length stays under ``char_budget``.
"""
chunks: list[str] = []
current: list[str] = []
current_len = 0
for piece in pieces:
piece_len = len(piece) + 1 # +1 for the newline that will join them
if current and current_len + piece_len > char_budget:
chunks.append("\n".join(current))
current = [piece]
current_len = piece_len
else:
current.append(piece)
current_len += piece_len
if current:
chunks.append("\n".join(current))
return chunks
def mean_pool_normalized(vectors: Sequence[Sequence[float]]) -> list[float]:
"""Mean across vectors then re-normalize to unit length.
Used to roll multiple chunk embeddings into one representative vector
that still lives on the unit sphere (so cosine similarity remains
well-defined against single-chunk embeddings).
"""
if not vectors:
return []
dim = len(vectors[0])
pooled = [0.0] * dim
for vec in vectors:
for i in range(dim):
pooled[i] += vec[i]
n = len(vectors)
pooled = [x / n for x in pooled]
norm = math.sqrt(sum(x * x for x in pooled))
if norm == 0:
return pooled
return [x / norm for x in pooled]
def embed_pieces_pooled(
pieces: Sequence[str],
embedder: TextEmbedder,
) -> tuple[list[float], int, str]:
"""Chunk ``pieces`` to fit the embedder, embed each chunk, mean-pool the result.
Returns ``(vector, dimension, model_id)``. For an empty input returns
``([], 0, embedder.model_id)`` so callers can detect "nothing to embed"
without checking length on the vector themselves.
"""
if not pieces:
return [], 0, embedder.model_id
chunks = chunk_pieces_for_embedding(pieces)
chunk_embeddings = embedder.embed_many(chunks)
if not chunk_embeddings:
return [], 0, embedder.model_id
if len(chunk_embeddings) == 1:
emb = chunk_embeddings[0]
return emb.vector, emb.dimension, emb.model_id
pooled = mean_pool_normalized([e.vector for e in chunk_embeddings])
return pooled, chunk_embeddings[0].dimension, chunk_embeddings[0].model_id