Aurelius / core /embedding.py
murtaza-2007
Rebuild Aurelius as a multi-source graph intelligence engine
75ce203
Raw
History Blame Contribute Delete
5.37 kB
"""Aurelius core — text-embedding model singleton + per-search caches.
v2 concurrency fix: v1 kept a single module-global `_emb_cache`/`_dead_ends`
pair that reset_session_caches() rebound at the start of EVERY search, while
the server allowed 4 concurrent searches — concurrent runs wiped each
other's caches mid-flight and re-embedded the same titles over and over.
The globals are gone. Each search now owns an EmbeddingCache instance
(navigator-scoped), and long-lived stores can own their own instance with
whatever lifetime they need. Nothing here is shared mutable state.
"""
from __future__ import annotations
import asyncio
import os
import time
from pathlib import Path
from typing import Optional
# Once the model has been downloaded once, sentence-transformers /
# huggingface_hub still spend 1-3s on every startup doing a network
# round-trip to check for updates. Must be set before sentence_transformers
# is imported; only skips the check when a cached model already exists
# locally (first run still goes online to fetch it).
_hf_cache = Path.home() / ".cache" / "huggingface" / "hub"
if _hf_cache.exists() and any(_hf_cache.glob("models--*")):
os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
import numpy as np
from sentence_transformers import SentenceTransformer
from config import EMBED_MODEL_NAME, EMBED_DEVICE, EMBED_BATCH_SIZE
_EMBED_MODEL: Optional[SentenceTransformer] = None
# Single-flight lock around encode(): torch with num_threads=1 gains nothing
# from interleaved encodes, and serializing them keeps per-call latency
# predictable when several searches run at once.
_ENCODE_LOCK = asyncio.Lock()
async def load_model():
"""Startup handler: load the model once, in-process.
No silent fallback on failure — if the load fails this re-raises so the
server refuses connections rather than degrading silently.
"""
global _EMBED_MODEL
# Free-tier hosts give one throttled vCPU; torch's default thread pool
# spawns one thread per logical core it thinks exists, each with its own
# arena. Cap to 1: no extra cores to run them anyway.
import torch
torch.set_num_threads(1)
print(f"[Embed] Loading sentence-transformers model '{EMBED_MODEL_NAME}'...")
t0 = time.time()
_EMBED_MODEL = SentenceTransformer(EMBED_MODEL_NAME, device=EMBED_DEVICE)
print(f"[Embed] {EMBED_MODEL_NAME} ready "
f"(dim={_EMBED_MODEL.get_embedding_dimension()}, {time.time()-t0:.1f}s)")
def model_loaded() -> bool:
return _EMBED_MODEL is not None
def embedding_dim() -> int:
return _EMBED_MODEL.get_embedding_dimension() if _EMBED_MODEL else 0
def cosine_similarity(a, b) -> float:
if a is None or b is None:
return 0.0
a = np.asarray(a, dtype=np.float32)
b = np.asarray(b, dtype=np.float32)
if a.size == 0 or b.size == 0:
return 0.0
na = np.linalg.norm(a)
nb = np.linalg.norm(b)
if na == 0 or nb == 0:
return 0.0
return float(np.dot(a, b) / (na * nb))
class EmbeddingCache:
"""A key → vector cache with batched, executor-offloaded encoding.
Keys are caller-chosen (the navigator uses NodeRef.key()); `texts` is
what actually gets encoded — always the enriched "{title}. {context}"
form, never a bare title (the v1 semantic-drift lesson).
"""
def __init__(self):
self._cache: dict[str, np.ndarray] = {}
self.encode_calls = 0
def get(self, key: str) -> Optional[np.ndarray]:
return self._cache.get(key)
def put(self, key: str, emb: np.ndarray):
self._cache[key] = emb
def __contains__(self, key: str) -> bool:
return key in self._cache
async def embed(self, keys: list[str],
texts: list[str] | None = None) -> list[np.ndarray]:
"""Return embeddings for keys (parallel lists), encoding only the
uncached ones in a single batched, off-loop encode call."""
if not keys or _EMBED_MODEL is None:
return [np.array([]) for _ in keys]
if texts is None:
texts = keys
uncached_idx = [i for i, k in enumerate(keys) if k not in self._cache]
if uncached_idx:
uncached_texts = [texts[i] for i in uncached_idx]
loop = asyncio.get_running_loop()
async with _ENCODE_LOCK:
embs = await loop.run_in_executor(
None,
lambda: _EMBED_MODEL.encode(
uncached_texts, convert_to_numpy=True,
batch_size=EMBED_BATCH_SIZE, show_progress_bar=False,
),
)
self.encode_calls += 1
for orig_i, emb in zip(uncached_idx, embs):
self._cache[keys[orig_i]] = emb
return [self._cache.get(k, np.array([])) for k in keys]
def encode_texts_sync(texts: list[str]) -> np.ndarray:
"""Synchronous batch encode for offline ingestion pipelines (no event
loop, no cache). Raises if the model isn't loaded."""
if _EMBED_MODEL is None:
raise RuntimeError("Embedding model not loaded — call load_model() first")
return _EMBED_MODEL.encode(texts, convert_to_numpy=True,
batch_size=EMBED_BATCH_SIZE,
show_progress_bar=False)