Add arxiv-complete retrieval layer (semantic + keyword)
This commit is contained in:
@@ -0,0 +1,207 @@
|
||||
"""
|
||||
Query interface for the retrieval layer.
|
||||
|
||||
``query()`` is the single entry point the CLI (and any caller) uses. It:
|
||||
|
||||
1. embeds the natural-language query,
|
||||
2. runs an ANN search against the FAISS index,
|
||||
3. attaches title / year / category / URL / abstract snippet to each hit,
|
||||
4. applies the requested filters (year range, primary category, survey-only),
|
||||
5. returns a ranked list of ``Result`` objects.
|
||||
|
||||
If no index exists on disk yet (or building it failed), ``query()`` falls back
|
||||
to a **pure-keyword** search that streams the metadata and matches on title/
|
||||
abstract with a TF-style score. This guarantees the CLI always returns
|
||||
something, even before an embedding index has been built.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
|
||||
from data import DataStore, SURVEY_KEYWORDS
|
||||
|
||||
# Cap the working set for the keyword fallback so it stays fast on the 3.15M-row
|
||||
# metadata file (term-overlap scoring doesn't need the full corpus).
|
||||
_KEYWORD_WORKING_SET = 200_000
|
||||
|
||||
|
||||
@dataclass
|
||||
class Result:
|
||||
paper_id: str
|
||||
title: str
|
||||
year: int
|
||||
category: str
|
||||
url: str
|
||||
snippet: str
|
||||
score: float
|
||||
method: str # "semantic" | "keyword"
|
||||
|
||||
|
||||
@dataclass
|
||||
class QueryOptions:
|
||||
top_k: int = 10
|
||||
year_min: int | None = None
|
||||
year_max: int | None = None
|
||||
primary_category: str | None = None
|
||||
survey_only: bool = False
|
||||
embedder = None # set by caller; optional
|
||||
|
||||
|
||||
def _snippet(abstract: str, maxlen: int = 200) -> str:
|
||||
text = re.sub(r"\s+", " ", abstract or "").strip()
|
||||
if len(text) <= maxlen:
|
||||
return text
|
||||
return text[:maxlen].rsplit(" ", 1)[0] + " …"
|
||||
|
||||
|
||||
def _normalize(text: str) -> str:
|
||||
return re.sub(r"[^a-z0-9]+", " ", (text or "").lower()).strip()
|
||||
|
||||
|
||||
def _keyword_score(text: str, terms: list[str]) -> float:
|
||||
"""TF-style overlap score between a doc and the query terms."""
|
||||
words = set(text.split())
|
||||
score = sum(1 for t in terms if t in words)
|
||||
# small bonus for title matches
|
||||
return score
|
||||
|
||||
|
||||
def _keyword_search(
|
||||
query: str,
|
||||
store: DataStore,
|
||||
opts: QueryOptions,
|
||||
) -> list[Result]:
|
||||
"""
|
||||
Pure-keyword fallback: stream matching rows, score by term overlap, filter,
|
||||
and return the top-k. Used when no ANN index is available.
|
||||
"""
|
||||
terms = [t for t in _normalize(query).split() if len(t) > 1]
|
||||
if not terms:
|
||||
return []
|
||||
|
||||
# Pull a working set of rows (push the cheap filters into the scan).
|
||||
scan = store.scan_filtered(
|
||||
year_min=opts.year_min,
|
||||
year_max=opts.year_max,
|
||||
primary_category=opts.primary_category,
|
||||
survey_only=opts.survey_only,
|
||||
)
|
||||
rows = (
|
||||
scan.select(
|
||||
["paper_id", "title", "primary_category", "first_version_date", "arxiv_abs_url", "abstract"]
|
||||
)
|
||||
.limit(_KEYWORD_WORKING_SET)
|
||||
.collect()
|
||||
.to_dicts()
|
||||
)
|
||||
if not rows:
|
||||
return []
|
||||
|
||||
scored: list[Result] = []
|
||||
for r in rows:
|
||||
rec = dict(r)
|
||||
title = str(rec["title"] or "")
|
||||
abstract = str(rec["abstract"] or "")
|
||||
body = f"{title} {abstract}"
|
||||
score = _keyword_score(body.lower(), terms)
|
||||
if score <= 0:
|
||||
continue
|
||||
fvd = rec["first_version_date"]
|
||||
year = fvd.year if isinstance(fvd, __import__("datetime").datetime) else (int(fvd) if fvd is not None else 0)
|
||||
scored.append(
|
||||
Result(
|
||||
paper_id=str(rec["paper_id"]),
|
||||
title=title,
|
||||
year=year,
|
||||
category=str(rec["primary_category"] or ""),
|
||||
url=str(rec["arxiv_abs_url"] or ""),
|
||||
snippet=_snippet(abstract),
|
||||
score=score,
|
||||
method="keyword",
|
||||
)
|
||||
)
|
||||
scored.sort(key=lambda x: x.score, reverse=True)
|
||||
return scored[: opts.top_k]
|
||||
|
||||
|
||||
def query(
|
||||
query_text: str,
|
||||
store: DataStore,
|
||||
ann=None,
|
||||
embedder=None,
|
||||
opts: QueryOptions | None = None,
|
||||
) -> list[Result]:
|
||||
"""
|
||||
Run a search. Tries the ANN index first; falls back to keyword search.
|
||||
|
||||
Returns a list of ranked ``Result`` objects.
|
||||
"""
|
||||
opts = opts or QueryOptions()
|
||||
query_text = query_text or ""
|
||||
|
||||
# --- semantic path ---
|
||||
if ann is not None and embedder is not None:
|
||||
try:
|
||||
return _semantic_query(query_text, store, ann, embedder, opts)
|
||||
except Exception as exc: # pragma: no cover - defensive
|
||||
print(f"[query] ANN search failed ({exc}); falling back to keyword search.")
|
||||
|
||||
# --- keyword fallback ---
|
||||
return _keyword_search(query_text, store, opts)
|
||||
|
||||
|
||||
def _semantic_query(
|
||||
query_text: str,
|
||||
store: DataStore,
|
||||
ann,
|
||||
embedder,
|
||||
opts: QueryOptions,
|
||||
) -> list[Result]:
|
||||
import faiss
|
||||
|
||||
qvec = np.ascontiguousarray(embedder.embed_one(query_text), dtype=np.float32)
|
||||
k = min(opts.top_k * 8, len(ann), ann.index.ntotal)
|
||||
k = max(k, 1)
|
||||
dists, idx = ann.index.search(qvec.reshape(1, -1), k)
|
||||
dists = dists[0]
|
||||
idx = idx[0]
|
||||
|
||||
# Gather candidate hits and apply filters.
|
||||
candidates: list[Result] = []
|
||||
for d, i in zip(dists.tolist(), idx.tolist()):
|
||||
if i < 0 or i >= len(ann._ids):
|
||||
continue
|
||||
pid = ann._ids[i]
|
||||
year = ann._years[i]
|
||||
cat = ann._cats[i]
|
||||
if opts.year_min is not None and year < opts.year_min:
|
||||
continue
|
||||
if opts.year_max is not None and year > opts.year_max:
|
||||
continue
|
||||
if opts.primary_category and cat != opts.primary_category:
|
||||
continue
|
||||
if opts.survey_only:
|
||||
if not any(kw in ann._abstracts[i].lower() for kw in SURVEY_KEYWORDS):
|
||||
continue
|
||||
# FAISS L2 distance -> similarity in [0,1]; closer => higher.
|
||||
sim = 1.0 / (1.0 + float(d))
|
||||
candidates.append(
|
||||
Result(
|
||||
paper_id=pid,
|
||||
title=ann._titles[i],
|
||||
year=year,
|
||||
category=cat,
|
||||
url=ann._urls[i],
|
||||
snippet=_snippet(ann._abstracts[i]),
|
||||
score=sim,
|
||||
method="semantic",
|
||||
)
|
||||
)
|
||||
|
||||
candidates.sort(key=lambda x: x.score, reverse=True)
|
||||
return candidates[: opts.top_k]
|
||||
Reference in New Issue
Block a user