Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 11 additions & 3 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,10 @@ RAG_LLM_THINKING=disabled

# --- Embeddings ---------------------------------------------------------
RAG_EMBED_MODEL=BAAI/bge-small-en-v1.5
# onnx_quantized (default, fast) | sentence_transformers (torch baseline)
RAG_EMBED_BACKEND=onnx_quantized
# onnx_optimized (default: fastembed ONNX Runtime, fp16 graph-optimised) | sentence_transformers (torch baseline)
RAG_EMBED_BACKEND=onnx_optimized
# Optional local model directory (offline / air-gapped); empty = download from the Hub.
RAG_EMBED_MODEL_PATH=
RAG_EMBED_BATCH_SIZE=64
RAG_EMBED_CACHE_SIZE=1024

Expand All @@ -38,9 +40,15 @@ RAG_TEXT_COLUMNS=
# --- Retrieval / agent --------------------------------------------------
RAG_TOP_K=5
RAG_CANDIDATE_K=20
RAG_MMR_LAMBDA=0.5
# 1.0 = MMR off (pure relevance). Lower values trade relevance for diversity;
# on SciFact 0.5 cost 5 points of nDCG@10 (see README, Retrieval quality).
RAG_MMR_LAMBDA=1.0
# Max sub-queries the planner may decompose a question into.
RAG_MAX_SUBQUERIES=3
# Hybrid search: fuse BM25 keyword search with vector search.
RAG_HYBRID_SEARCH=false
RAG_KEYWORD_WEIGHT=0.5
RAG_RRF_K=60

# --- Service ------------------------------------------------------------
RAG_HOST=0.0.0.0
Expand Down
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -21,3 +21,6 @@ bench/results/*.json
.pytest_cache/
.ruff_cache/
.DS_Store

# Evaluation data (downloaded by eval/download_scifact.py); results are committed
eval/data/
10 changes: 8 additions & 2 deletions Makefile
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
.PHONY: help install install-baseline dev test lint run ingest bench docker-build docker-up docker-down clean
.PHONY: help install install-baseline dev test lint run ingest bench eval-data eval docker-build docker-up docker-down clean

PYTHON ?= python3
VENV ?= .venv
Expand All @@ -22,7 +22,7 @@ test: ## Run the test suite
$(BIN)/pytest

lint: ## Lint with ruff
$(BIN)/ruff check app bench scripts tests
$(BIN)/ruff check app bench eval scripts tests

run: ## Start the API on :8000 with reload
$(BIN)/uvicorn app.main:app --reload --host 0.0.0.0 --port 8000
Expand All @@ -39,6 +39,12 @@ demo-docker: ## Same demo, against the container (corpus mounted at /corpus)
bench: ## Run the latency benchmark against $(CORPUS)
$(BIN)/python -m bench.benchmark --corpus $(CORPUS)

eval-data: ## Download the BEIR SciFact benchmark (~4.6 MB) into eval/data/
$(BIN)/python -m eval.download_scifact

eval: ## Retrieval-quality evaluation (nDCG@10, recall, MRR) on SciFact
$(BIN)/python -m eval.retrieval_eval

docker-build: ## Build the container image
docker compose build

Expand Down
192 changes: 151 additions & 41 deletions README.md

Large diffs are not rendered by default.

15 changes: 12 additions & 3 deletions app/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,8 @@ class Settings(BaseSettings):

# --- Embeddings ---
embed_model: str = "BAAI/bge-small-en-v1.5"
embed_backend: str = "onnx_quantized"
embed_backend: str = "onnx_optimized"
embed_model_path: str = "" # local model directory; empty = download from the Hub
embed_batch_size: int = 64
embed_cache_size: int = 1024

Expand All @@ -45,8 +46,16 @@ class Settings(BaseSettings):
# --- Retrieval ---
top_k: int = 5
candidate_k: int = 20
mmr_lambda: float = 0.5
# 1.0 = MMR off (pure relevance). On SciFact, MMR at 0.5 cost 4.9 points of
# nDCG@10 (eval/results/scifact.json), so diversification is opt-in.
mmr_lambda: float = 1.0
max_subqueries: int = 3
# Fuse BM25 keyword search with vector search. Off by default: on SciFact it raised
# recall@100 but not the top-10 ranking. Turn it on for corpora full of exact
# identifiers (part numbers, error codes, gene names) that embeddings blur.
hybrid_search: bool = False
keyword_weight: float = 0.5 # BM25's weight in the rank fusion (vector list = 1.0)
rrf_k: int = 60

# --- Service ---
host: str = "0.0.0.0"
Expand All @@ -56,7 +65,7 @@ class Settings(BaseSettings):
@field_validator("embed_backend")
@classmethod
def _valid_backend(cls, v: str) -> str:
allowed = {"onnx_quantized", "sentence_transformers"}
allowed = {"onnx_optimized", "onnx_quantized", "sentence_transformers"}
if v not in allowed:
raise ValueError(f"embed_backend must be one of {sorted(allowed)}, got {v!r}")
return v
Expand Down
5 changes: 5 additions & 0 deletions app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ def build_state(settings: Settings) -> AppState:
model_name=settings.embed_model,
batch_size=settings.embed_batch_size,
cache_size=settings.embed_cache_size,
model_path=settings.embed_model_path or None,
)
# Load the model and initialise the compute graph before serving, so the
# first real request doesn't pay for it.
Expand All @@ -70,12 +71,16 @@ def build_state(settings: Settings) -> AppState:
hnsw_m=settings.hnsw_m,
hnsw_ef_construction=settings.hnsw_ef_construction,
hnsw_ef_search=settings.hnsw_ef_search,
lexical=settings.hybrid_search,
)
retriever = Retriever(
store=store,
top_k=settings.top_k,
candidate_k=settings.candidate_k,
mmr_lambda=settings.mmr_lambda,
hybrid=settings.hybrid_search,
rrf_k=settings.rrf_k,
keyword_weight=settings.keyword_weight,
)
agent = RetrievalAgent(retriever=retriever, max_subqueries=settings.max_subqueries)
answerer = build_answerer(
Expand Down
51 changes: 40 additions & 11 deletions app/retrieval/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,16 @@
The baseline. Runs the model through PyTorch in fp32. This is the
conventional way to serve a Hugging Face embedding model.

``onnx_quantized``
The optimisation. Runs the *same* model architecture through ONNX Runtime
with int8-quantised weights (via ``fastembed``). Same vector space, same
retrieval quality in practice, materially less CPU work per call.
``onnx_optimized``
The optimisation. Runs the same model through ONNX Runtime via ``fastembed``.
For ``BAAI/bge-small-en-v1.5`` fastembed ships a graph-optimised export
(fused attention, GELU and layer-norm kernels) with fp16 weights. It is *not*
int8-quantised, despite the repository name ``Qdrant/bge-small-en-v1.5-onnx-Q``:
its ``ort_config.json`` has an empty quantization block and every weight tensor
is FLOAT16. Its vectors match the fp32 model's to a cosine similarity of
0.99999, and ``eval/retrieval_eval.py`` shows identical retrieval quality.

``onnx_quantized`` is accepted as a deprecated alias for backwards compatibility.

A small LRU cache sits in front of query embedding, because repeated and
near-repeated queries are the norm in a served system and re-encoding them is
Expand Down Expand Up @@ -56,17 +62,28 @@ def warmup(self) -> None:
self.embed_query("warmup")


class OnnxQuantizedEmbedder(Embedder):
"""int8-quantised ONNX Runtime backend (fastembed). The served default."""
class OnnxEmbedder(Embedder):
"""Graph-optimised ONNX Runtime backend (fastembed). The served default.

``model_path`` points fastembed at a local model directory instead of the
Hugging Face Hub, for offline or air-gapped deployments.
"""

name = "onnx_quantized"
name = "onnx_optimized"

def __init__(self, model_name: str, batch_size: int = 64, threads: int | None = None):
def __init__(
self,
model_name: str,
batch_size: int = 64,
threads: int | None = None,
model_path: str | None = None,
):
from fastembed import TextEmbedding

self.model_name = model_name
self.batch_size = batch_size
self._model = TextEmbedding(model_name=model_name, threads=threads)
kwargs = {"specific_model_path": model_path} if model_path else {}
self._model = TextEmbedding(model_name=model_name, threads=threads, **kwargs)
self._dimension: int | None = None

def embed_documents(self, texts: list[str]) -> list[list[float]]:
Expand Down Expand Up @@ -198,15 +215,27 @@ def cache_stats(self) -> dict[str, int | float]:
}


# Kept so existing imports keep working after the rename.
OnnxQuantizedEmbedder = OnnxEmbedder

ONNX_BACKENDS = {"onnx_optimized", "onnx_quantized"} # the second is a deprecated alias


def build_embedder(
backend: str,
model_name: str,
batch_size: int = 64,
cache_size: int = 1024,
model_path: str | None = None,
) -> Embedder:
"""Construct the configured backend, optionally wrapped in a query cache."""
if backend == "onnx_quantized":
embedder: Embedder = OnnxQuantizedEmbedder(model_name, batch_size=batch_size)
if backend in ONNX_BACKENDS:
if backend == "onnx_quantized":
logger.warning(
"RAG_EMBED_BACKEND=onnx_quantized is deprecated (the model is fp16 and "
"graph-optimised, not int8); use onnx_optimized"
)
embedder: Embedder = OnnxEmbedder(model_name, batch_size=batch_size, model_path=model_path)
elif backend == "sentence_transformers":
embedder = SentenceTransformerEmbedder(model_name, batch_size=batch_size)
else:
Expand Down
161 changes: 161 additions & 0 deletions app/retrieval/lexical.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
"""In-memory BM25 keyword index, used alongside vector search for hybrid retrieval.

Dense embeddings are good at paraphrase ("heart attack" ~ "myocardial infarction") but
weak on exact tokens a user types verbatim: gene names, drug names, error codes, IDs.
BM25 is the opposite. Fusing the two ranked lists with reciprocal rank fusion gets
most of the benefit of each; ``eval/retrieval_eval.py`` measures the difference.

The index lives in memory and is rebuilt from the vector store at startup. With
posting lists in plain dicts it handles tens of thousands of chunks comfortably;
beyond a few hundred thousand, a dedicated engine (Elasticsearch, OpenSearch, or
a database full-text index) is the better home for it.
"""

from __future__ import annotations

import math
import re
import threading
from collections import Counter, defaultdict
from typing import Any

_TOKEN = re.compile(r"[a-z0-9]+(?:[-'][a-z0-9]+)*")

# A short English stopword list. Removing these keeps posting lists small and stops
# very common words from contributing noise to the score.
STOPWORDS = frozenset(
{
"a", "an", "and", "are", "as", "at", "be", "but", "by", "for", "from", "has", "have", "he",
"her", "his", "i", "if", "in", "into", "is", "it", "its", "of", "on", "or", "our", "she",
"so", "than", "that", "the", "their", "them", "then", "there", "these", "they", "this",
"to", "was", "we", "were", "what", "when", "where", "which", "who", "whom", "why", "will",
"with",
"would", "you", "your", "do", "does", "did", "not", "no", "can", "could", "should", "may",
"might", "also", "been", "being", "about", "over", "under", "between", "after", "before",
}
)


# Plural folding, so "genes" matches "gene" and "studies" matches "study". This is the
# plural step of the Porter stemmer only: a full stemmer conflates more word forms but
# also more unrelated words. On the SciFact benchmark it adds 0.6 points of nDCG@10
# and 1.9 points of recall@100 over no stemming (eval/retrieval_eval.py).
def stem(token: str) -> str:
if len(token) > 4 and token.endswith("ies"):
return token[:-3] + "y"
if token.endswith("sses"):
return token[:-2]
if len(token) > 3 and token.endswith("s") and not token.endswith(("ss", "us", "is")):
return token[:-1]
return token


def tokenize(text: str) -> list[str]:
"""Lowercase alphanumeric tokens (hyphenated terms kept whole), no stopwords, plurals folded."""
return [stem(t) for t in _TOKEN.findall(text.lower()) if t not in STOPWORDS and len(t) > 1]


class BM25Index:
"""Okapi BM25 over chunk texts, keyed by chunk ID.

``k1`` controls term-frequency saturation and ``b`` document-length
normalisation. The defaults (0.9, 0.4) are the Anserini/Pyserini defaults
that BEIR-style evaluations commonly use.
"""

def __init__(self, k1: float = 0.9, b: float = 0.4):
self.k1 = k1
self.b = b
self._lock = threading.Lock()
self._postings: dict[str, dict[str, int]] = defaultdict(dict) # term -> {doc: tf}
self._doc_len: dict[str, int] = {}
self._doc_terms: dict[str, Counter] = {}
self._metadata: dict[str, dict[str, Any]] = {}
self._text: dict[str, str] = {}
self._total_len = 0

# ------------------------------------------------------------------
# Writes
# ------------------------------------------------------------------
def add(self, ids: list[str], texts: list[str], metadatas: list[dict[str, Any]] | None = None):
"""Add or replace documents. Re-adding an ID replaces it (upsert semantics)."""
metadatas = metadatas or [{} for _ in ids]
with self._lock:
for doc_id, text, metadata in zip(ids, texts, metadatas, strict=True):
if doc_id in self._doc_len:
self._remove(doc_id)
terms = Counter(tokenize(text))
for term, tf in terms.items():
self._postings[term][doc_id] = tf
length = sum(terms.values())
self._doc_len[doc_id] = length
self._doc_terms[doc_id] = terms
self._metadata[doc_id] = dict(metadata or {})
self._text[doc_id] = text
self._total_len += length

def _remove(self, doc_id: str) -> None:
for term in self._doc_terms.pop(doc_id, {}):
postings = self._postings.get(term)
if postings is not None:
postings.pop(doc_id, None)
if not postings:
del self._postings[term]
self._total_len -= self._doc_len.pop(doc_id, 0)
self._metadata.pop(doc_id, None)
self._text.pop(doc_id, None)

def clear(self) -> None:
with self._lock:
self._postings.clear()
self._doc_len.clear()
self._doc_terms.clear()
self._metadata.clear()
self._text.clear()
self._total_len = 0

# ------------------------------------------------------------------
# Reads
# ------------------------------------------------------------------
def __len__(self) -> int:
return len(self._doc_len)

def search(
self, query: str, k: int, where: dict[str, Any] | None = None
) -> list[tuple[str, float]]:
"""Return up to ``k`` ``(doc_id, score)`` pairs, best first.

``where`` supports the same simple equality filter the service uses for
vector search (for example ``{"source_tag": "my-dataset"}``).
"""
terms = tokenize(query)
n_docs = len(self._doc_len)
if not terms or n_docs == 0 or k <= 0:
return []
avg_len = self._total_len / n_docs

scores: dict[str, float] = defaultdict(float)
with self._lock:
for term in set(terms):
postings = self._postings.get(term)
if not postings:
continue
df = len(postings)
idf = math.log(1 + (n_docs - df + 0.5) / (df + 0.5))
for doc_id, tf in postings.items():
norm = self.k1 * (1 - self.b + self.b * self._doc_len[doc_id] / avg_len)
scores[doc_id] += idf * tf * (self.k1 + 1) / (tf + norm)

if where:
scores = {
d: s
for d, s in scores.items()
if all(
self._metadata.get(d, {}).get(key) == value for key, value in where.items()
)
}

return sorted(scores.items(), key=lambda kv: kv[1], reverse=True)[:k]

def document(self, doc_id: str) -> tuple[str, dict[str, Any]]:
return self._text.get(doc_id, ""), self._metadata.get(doc_id, {})
Loading
Loading