| """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 |
|
|
| |
| |
| |
| |
| |
| _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 |
|
|
| |
| |
| |
| _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 |
| |
| |
| |
| 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) |
|
|