File size: 7,406 Bytes
c51da3c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 | """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]
|