lexrag / app /retrieval.py
RV302001's picture
Initial commit: LexRAG hybrid legal RAG
cc0207f
Raw
History Blame Contribute Delete
3.77 kB
"""Hybrid retrieval: pgvector cosine + Postgres full-text, fused via RRF.
The bge model is loaded ONCE at import (module global), never per request.
"""
from __future__ import annotations
from sentence_transformers import SentenceTransformer
from sqlalchemy import text
from app.db import SessionLocal
MODEL_NAME = "BAAI/bge-small-en-v1.5"
TOP_K = 5
CANDIDATES = 20 # per-arm candidate pool before RRF
RRF_K = 60 # RRF damping constant
POOL = 25 # fused rows fetched before diversification
MAX_PER_CASE = 2 # cap chunks shown from any single judgment
# Loaded once at import time. Reused for every query.
_model = SentenceTransformer(MODEL_NAME)
# Single-CTE RRF: rank vector arm + keyword arm, fuse by reciprocal rank.
_HYBRID_SQL = text(
"""
WITH vec AS (
SELECT id, row_number() OVER (ORDER BY embedding <=> CAST(:qvec AS vector)) AS r
FROM chunks
ORDER BY embedding <=> CAST(:qvec AS vector)
LIMIT :cand
),
kw AS (
SELECT id, row_number() OVER (
ORDER BY ts_rank_cd(ts, plainto_tsquery('english', :q)) DESC) AS r
FROM chunks
WHERE ts @@ plainto_tsquery('english', :q)
LIMIT :cand
),
fused AS (
SELECT COALESCE(vec.id, kw.id) AS id,
(COALESCE(1.0/(:rrfk + vec.r), 0)
+ COALESCE(1.0/(:rrfk + kw.r), 0)) AS rrf
FROM vec
FULL OUTER JOIN kw USING (id)
ORDER BY rrf DESC
LIMIT :topk
)
SELECT c.id, c.case_name, c.content, fused.rrf
FROM fused
JOIN chunks c USING (id)
ORDER BY fused.rrf DESC;
"""
)
def embed_query(query: str) -> list[float]:
"""Embed one query string. Normalized to match ingest (cosine)."""
vec = _model.encode([query], normalize_embeddings=True, convert_to_numpy=True)[0]
return vec.tolist()
def _diversify(rows, top_k: int, max_per_case: int) -> list:
"""Cap chunks per case so the user sees several judgments, not one.
Two-pass over RRF-ordered rows: first take up to `max_per_case` per case,
then if short, backfill with leftover chunks (still RRF order).
"""
kept, overflow = [], []
seen: dict[str, int] = {}
for r in rows:
if seen.get(r.case_name, 0) < max_per_case:
kept.append(r)
seen[r.case_name] = seen.get(r.case_name, 0) + 1
else:
overflow.append(r)
if len(kept) >= top_k:
break
if len(kept) < top_k:
kept.extend(overflow[: top_k - len(kept)])
return kept[:top_k]
def hybrid_search(query: str, top_k: int = TOP_K) -> list[dict]:
"""Hybrid SQL -> diversify by case -> top-k [{case_name, content, score}]."""
qvec = embed_query(query)
db = SessionLocal()
try:
rows = db.execute(
_HYBRID_SQL,
{
"qvec": str(qvec), # pgvector accepts the "[...]" string form
"q": query,
"cand": CANDIDATES,
"rrfk": RRF_K,
"topk": POOL, # fetch a wider pool; trim after diversifying
},
).all()
finally:
db.close()
rows = _diversify(rows, top_k, MAX_PER_CASE)
return [
{"case_name": r.case_name, "content": r.content, "score": float(r.rrf)}
for r in rows
]
if __name__ == "__main__":
# Self-test: an exact statute term that pure-vector often misses but
# full-text nails. Confirms the keyword arm contributes to fusion.
q = "Section 302"
print(f"Query: {q!r}\n")
for i, hit in enumerate(hybrid_search(q), 1):
present = "Section 302" in hit["content"]
print(f"{i}. rrf={hit['score']:.5f} [exact-term:{present}] {hit['case_name']}")
print(f" {hit['content'][:140]!r}\n")