"""Figure retrieval over the caption index. Supports three things beyond plain nearest-neighbour search: multi-query fusion - search with several phrasings of the same request and combine the ranked lists with reciprocal rank fusion author filter - restrict the search to papers by named authors, applied inside FAISS as an exact ID constraint, not a post-hoc trim, so the top k is the top k among that author's figures rather than whatever survived a global search author boost - softly promote papers by named authors without excluding anything else (off unless you pass boost_authors) Author names are matched loosely: "McQuinn", "M. McQuinn" and "Matthew McQuinn" all match the same person. A name that matches nobody in the corpus contributes nothing, so a wrong or invented name degrades to a no-op rather than corrupting the result. """ import json import re import unicodedata from collections import defaultdict from pathlib import Path import numpy as np import pyarrow.parquet as pq RRF_K = 60 BOOST_WEIGHT = 0.25 def normalize_name(name: str) -> str: """Lowercase, strip accents and punctuation, and put the family name last. Handles both "Given Family" (as Semantic Scholar returns) and "Family, Given" (as people often type), so the two normalize alike. """ name = unicodedata.normalize("NFKD", name) name = "".join(c for c in name if not unicodedata.combining(c)) name = name.lower() if "," in name: family, _, given = name.partition(",") name = f"{given} {family}" name = re.sub(r"['\u2019\-]", "", name) name = re.sub(r"[^a-z0-9 ]", " ", name) return re.sub(r"\s+", " ", name).strip() def last_name(name: str) -> str: parts = normalize_name(name).split() return parts[-1] if parts else "" class FigureIndex: def __init__(self, index_dir, authors_path=None): import faiss index_dir = Path(index_dir) self.index = faiss.read_index(str(index_dir / "index.faiss")) self.meta = pq.read_table(index_dir / "meta.parquet").to_pylist() info_path = index_dir / "info.json" self.info = (json.load(info_path.open()) if info_path.exists() else {"backend": "openai", "model": "text-embedding-3-small", "dim": self.index.d}) self._embedder = None self.rows_by_paper = defaultdict(list) for row, m in enumerate(self.meta): self.rows_by_paper[m["arxiv_id"]].append(row) self.papers_by_lastname = defaultdict(set) self.authors_by_paper = {} if authors_path and Path(authors_path).exists(): table = pq.read_table(authors_path).to_pylist() for entry in table: names = entry["authors"] or [] if not names: continue self.authors_by_paper[entry["arxiv_id"]] = names for n in names: self.papers_by_lastname[last_name(n)].add(entry["arxiv_id"]) @property def embedder(self): if self._embedder is None: from embedders import embedder_from_info self._embedder = embedder_from_info(self.info) return self._embedder def papers_for_authors(self, queries: list[str]) -> set: """arXiv IDs whose author list matches any of the given names.""" matched = set() for q in queries: qn = normalize_name(q) if not qn: continue candidates = self.papers_by_lastname.get(last_name(q), set()) for arxiv_id in candidates: for full in self.authors_by_paper.get(arxiv_id, []): fn = normalize_name(full) if qn == fn or qn == last_name(full): matched.add(arxiv_id) break q_parts, f_parts = qn.split(), fn.split() if (len(q_parts) > 1 and q_parts[-1] == f_parts[-1] and q_parts[0][0] == f_parts[0][0]): matched.add(arxiv_id) break return matched def _row_selector(self, papers: set): import faiss rows = [] for arxiv_id in papers: rows.extend(self.rows_by_paper.get(arxiv_id, [])) if not rows: return None, 0 ids = np.array(sorted(rows), dtype="int64") return faiss.SearchParameters(sel=faiss.IDSelectorBatch(ids)), len(ids) def search(self, query: str, k: int = 20, variants=None, filter_authors=None, boost_authors=None, depth=None): """Return up to k figure matches, best first. variants: extra phrasings of the same query, fused with RRF filter_authors: hard restriction to papers by these authors boost_authors: soft promotion of papers by these authors depth: per-query retrieval depth before fusion (default 5k) """ texts = [query] + list(variants or []) depth = depth or max(k * 5, 100) params = None if filter_authors: papers = self.papers_for_authors(filter_authors) params, n_rows = self._row_selector(papers) if params is None: return [] depth = min(depth, n_rows) Q = self.embedder.embed(texts, is_query=True) if params is not None: sims, ids = self.index.search(Q, depth, params=params) else: sims, ids = self.index.search(Q, depth) scores = defaultdict(float) best_sim = {} for qi in range(len(texts)): for rank, row in enumerate(ids[qi]): if row < 0: continue row = int(row) scores[row] += 1.0 / (RRF_K + rank + 1) sim = float(sims[qi][rank]) if sim > best_sim.get(row, -1e9): best_sim[row] = sim if boost_authors: boosted = self.papers_for_authors(boost_authors) if boosted: bonus = BOOST_WEIGHT / RRF_K for row in list(scores): if self.meta[row]["arxiv_id"] in boosted: scores[row] += bonus ranked = sorted(scores.items(), key=lambda kv: -kv[1])[:k] out = [] for row, score in ranked: m = self.meta[row] out.append({ "arxiv_id": m["arxiv_id"], "fig_idx": m["fig_idx"], "caption": m.get("caption", ""), "authors": self.authors_by_paper.get(m["arxiv_id"], []), "fusion_score": score, "similarity": best_sim.get(row, 0.0), }) return out def search_rows(self, query: str, k: int, variants=None, filter_authors=None, boost_authors=None, depth=None): """Same as search() but returns raw meta row indices, for evaluation.""" hits = self.search(query, k=k, variants=variants, filter_authors=filter_authors, boost_authors=boost_authors, depth=depth) return [(h["arxiv_id"], h["fig_idx"]) for h in hits]