File size: 14,833 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 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 | """Robustness eval: 5 query styles per figure instead of one paraphrase.
For each sampled figure, one Haiku call produces five queries in different
registers, from terse to vague to notation-flipped. Each style is evaluated
separately, so the output shows how retrieval degrades as queries get more
human: shorter, vaguer, differently notated.
python eval_multi.py --indexes pilot_title_caption pilot_title_rewritten --n 300
Reuses helpers from eval_retrieval.py. Queries are cached in
multi_eval_queries.jsonl, so reruns only redo the cheap ranking math.
"""
import argparse
import json
import os
import random
import re
from pathlib import Path
import duckdb
import numpy as np
from eval_retrieval import embed_queries, evaluate, load_index
QUERY_MODEL = "claude-haiku-4-5-20251001"
STYLES = ["terse", "casual", "vague", "detailed", "notation"]
MULTI_PROMPT = """Here is a figure from an astrophysics paper:
Title: {title}
Caption: {caption}
Five different researchers are trying to find this figure in a search tool.
Write the query each would type. Do not reuse distinctive multi-word phrases
from the caption.
1. "terse": 4-8 keywords, no sentence structure, the way people actually type into search boxes.
2. "casual": one short natural sentence.
3. "vague": the researcher only half-remembers it. Just the general topic and roughly what kind of plot it was. It's fine to be imprecise or slightly wrong.
4. "detailed": a precise, complete description of the science and what is plotted.
5. "notation": like casual, but write any symbols or jargon in a DIFFERENT convention than the caption uses (e.g. sigma_8 vs S8 vs "amplitude of matter fluctuations", spelled-out names vs acronyms).
JSON only:
{{"terse": "...", "casual": "...", "vague": "...", "detailed": "...", "notation": "..."}}"""
def parse_json_obj(text: str) -> dict:
text = re.sub(r"```(json)?", "", text)
return json.loads(text[text.index("{"):text.rindex("}") + 1])
def generate_multi(samples: list[dict], cache_path: Path) -> dict:
from anthropic import Anthropic
client = None
cache = {}
if cache_path.exists():
for line in cache_path.open():
row = json.loads(line)
cache[row["key"]] = row["queries"]
with cache_path.open("a") as f:
for i, s in enumerate(samples):
key = f"{s['arxiv_id']}:{s['fig_idx']}"
if key in cache:
continue
if client is None:
client = Anthropic(api_key=os.environ["ANTHROPIC_API_KEY"])
prompt = MULTI_PROMPT.format(title=s["title"], caption=s["caption"][:1200])
try:
resp = client.messages.create(
model=QUERY_MODEL, max_tokens=500,
messages=[{"role": "user", "content": prompt}],
)
queries = parse_json_obj(resp.content[0].text)
if not all(st in queries and queries[st] for st in STYLES):
raise ValueError("missing styles")
except Exception as e:
print(f" query gen failed for {key}: {e}")
continue
cache[key] = queries
f.write(json.dumps({"key": key, "queries": queries}) + "\n")
if (i + 1) % 50 == 0:
print(f" {i+1}/{len(samples)} figures queried")
return cache
def generate_expansions(queries, style, n, cache_path):
"""LLM alternative phrasings, cached by (style, query)."""
import os
from anthropic import Anthropic
from query_expansion import expand_query
cache = {}
if cache_path.exists():
for line in cache_path.open():
row = json.loads(line)
cache[row["key"]] = row["variants"]
client = None
with cache_path.open("a") as f:
for i, q in enumerate(queries):
key = f"{style}:{q[:120]}"
if key in cache:
continue
if client is None:
client = Anthropic(api_key=os.environ["ANTHROPIC_API_KEY"])
variants = expand_query(q, client, QUERY_MODEL, n=n)
cache[key] = variants
f.write(json.dumps({"key": key, "variants": variants}) + "\n")
if (i + 1) % 50 == 0:
print(f" expanded {i+1}/{len(queries)}")
return [cache.get(f"{style}:{q[:120]}", []) for q in queries]
def generate_author_suggestions(queries, style, cache_path):
import os
from anthropic import Anthropic
from query_expansion import suggest_authors
cache = {}
if cache_path.exists():
for line in cache_path.open():
row = json.loads(line)
cache[row["key"]] = row["authors"]
client = None
with cache_path.open("a") as f:
for q in queries:
key = f"{style}:{q[:120]}"
if key in cache:
continue
if client is None:
client = Anthropic(api_key=os.environ["ANTHROPIC_API_KEY"])
names = suggest_authors(q, client, QUERY_MODEL)
cache[key] = names
f.write(json.dumps({"key": key, "authors": names}) + "\n")
return [cache.get(f"{style}:{q[:120]}", []) for q in queries]
def evaluate_retrieval(fi, queries, targets, k_list, variants_list=None,
filter_lists=None, boost_lists=None, depth=200):
"""Rank of each target using the FigureIndex search path."""
ranks = []
paper_ranks = []
for qi, q in enumerate(queries):
hits = fi.search_rows(
q, k=depth,
variants=variants_list[qi] if variants_list else None,
filter_authors=filter_lists[qi] if filter_lists else None,
boost_authors=boost_lists[qi] if boost_lists else None,
depth=depth,
)
tgt = targets[qi]
rank = None
p_rank = None
for r, (aid, fidx) in enumerate(hits):
if p_rank is None and aid == tgt[0]:
p_rank = r
if aid == tgt[0] and fidx == tgt[1]:
rank = r
break
ranks.append(rank)
paper_ranks.append(p_rank)
return ranks, paper_ranks
def evaluate_with_rerank(index, meta, info, queries, targets, reranker,
doc_lookup, depth):
Q = embed_queries(queries, info)
_, ids = index.search(Q, depth)
base_ranks = []
cand_lists = []
for qi, (tgt_id, tgt_fig) in enumerate(targets):
cand = [int(i) for i in ids[qi] if i >= 0]
cand_lists.append(cand)
rank = None
for r, idx in enumerate(cand):
m = meta[idx]
if m["arxiv_id"] == tgt_id and m["fig_idx"] == tgt_fig:
rank = r
break
base_ranks.append(rank)
if reranker is None:
return base_ranks, None
pairs = []
for qi, cand in enumerate(cand_lists):
for idx in cand:
m = meta[idx]
doc = doc_lookup.get((m["arxiv_id"], m["fig_idx"]), m.get("caption", ""))
pairs.append([queries[qi], doc])
scores = reranker.score(pairs)
rr_ranks = []
pos = 0
for qi, cand in enumerate(cand_lists):
s = scores[pos:pos + len(cand)]
pos += len(cand)
order = np.argsort(-s)
tgt_id, tgt_fig = targets[qi]
rank = None
for r, oi in enumerate(order):
m = meta[cand[oi]]
if m["arxiv_id"] == tgt_id and m["fig_idx"] == tgt_fig:
rank = r
break
rr_ranks.append(rank)
return base_ranks, rr_ranks
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--slice", default="astro_captions.parquet")
ap.add_argument("--index-dir", default="indexes")
ap.add_argument("--indexes", nargs="+",
default=["pilot_title_caption", "pilot_title_rewritten"])
ap.add_argument("--n", type=int, default=300)
ap.add_argument("--seed", type=int, default=11)
ap.add_argument("--expand", action="store_true",
help="fuse the query with LLM-generated alternative phrasings")
ap.add_argument("--expand-n", type=int, default=5)
ap.add_argument("--authors", default="authors.parquet",
help="author metadata parquet (from fetch_authors.py)")
ap.add_argument("--author-filter", action="store_true",
help="oracle author filter: restrict to papers sharing an "
"author with the target, simulating a user who "
"remembers one author name")
ap.add_argument("--author-boost", action="store_true",
help="soft boost using LLM-suggested authors (measured, not recommended)")
ap.add_argument("--rerank", action="store_true",
help="also rerank deep candidates with a local cross-encoder")
ap.add_argument("--rerank-model", default="BAAI/bge-reranker-base")
ap.add_argument("--rerank-depth", type=int, default=200)
ap.add_argument("--reuse-queries", action="store_true",
help="evaluate the figure set already in multi_eval_queries.jsonl "
"(for scale comparisons against a different index)")
args = ap.parse_args()
base = Path(args.index_dir)
first_index, first_meta, first_info = load_index(base, args.indexes[0])
con = duckdb.connect()
papers = {r[0]: r for r in con.execute(
f"SELECT arxiv_id, title, captions FROM read_parquet('{args.slice}')"
).fetchall()}
if args.reuse_queries:
picks = []
for line in Path("multi_eval_queries.jsonl").open():
arxiv_id, fig_idx = json.loads(line)["key"].rsplit(":", 1)
picks.append((arxiv_id, int(fig_idx)))
in_index = {(m["arxiv_id"], m["fig_idx"]) for m in first_meta}
missing = [p for p in picks if p not in in_index]
if missing:
print(f" WARNING: {len(missing)} cached figures not in this index, skipping them")
picks = [p for p in picks if p in in_index]
else:
rng = random.Random(args.seed)
indexed = [(m["arxiv_id"], m["fig_idx"]) for m in first_meta if m["fig_idx"] >= 1]
picks = rng.sample(indexed, min(args.n, len(indexed)))
samples = []
for arxiv_id, fig_idx in picks:
p = papers.get(arxiv_id)
if not p or fig_idx > len(p[2]):
continue
samples.append({"arxiv_id": arxiv_id, "fig_idx": fig_idx,
"title": p[1], "caption": p[2][fig_idx - 1]})
print(f"{len(samples)} eval figures, 5 query styles each")
cache = generate_multi(samples, Path("multi_eval_queries.jsonl"))
usable = [s for s in samples if f"{s['arxiv_id']}:{s['fig_idx']}" in cache]
print(f"{len(usable)} figures with full query sets\n")
reranker = None
doc_lookup = {}
if args.rerank:
from embedders import LocalReranker
from build_caption_index import clean_latex
print(f"Loading reranker {args.rerank_model}...")
reranker = LocalReranker(args.rerank_model)
for arxiv_id, (aid, title, captions) in papers.items():
for fi, cap in enumerate(captions):
doc_lookup[(arxiv_id, fi + 1)] = f"{title} | {clean_latex(cap)}"
depth = args.rerank_depth
use_search_path = args.expand or args.author_filter or args.author_boost
fig_index = None
author_lookup = {}
if use_search_path:
from retrieval import FigureIndex
fig_index = FigureIndex(base / args.indexes[0], authors_path=args.authors)
if not fig_index.authors_by_paper and (args.author_filter or args.author_boost):
raise SystemExit(f"No author metadata found at {args.authors}. "
f"Run fetch_authors.py first.")
author_lookup = fig_index.authors_by_paper
results = {}
for name in args.indexes:
index, meta, info = load_index(base, name)
for style in STYLES:
queries = [cache[f"{s['arxiv_id']}:{s['fig_idx']}"][style] for s in usable]
targets = [(s["arxiv_id"], s["fig_idx"]) for s in usable]
ranks, rr_ranks = evaluate_with_rerank(index, meta, info, queries, targets,
reranker, doc_lookup, depth)
n = len(ranks)
rec = lambda rs, k: sum(1 for r in rs if r is not None and r < k) / n
row = [rec(ranks, 1), rec(ranks, 5), rec(ranks, 20),
rec(ranks, 50), rec(ranks, depth)]
if rr_ranks is not None:
row += [rec(rr_ranks, 1), rec(rr_ranks, 5), rec(rr_ranks, 20)]
if use_search_path and name == args.indexes[0]:
variants_list = None
if args.expand:
print(f" [{style}] expanding queries...")
variants_list = generate_expansions(
queries, style, args.expand_n, Path("expansion_cache.jsonl"))
filter_lists = None
if args.author_filter:
filter_lists = [author_lookup.get(t[0], []) for t in targets]
boost_lists = None
if args.author_boost:
print(f" [{style}] suggesting authors...")
boost_lists = generate_author_suggestions(
queries, style, Path("author_suggestion_cache.jsonl"))
ex_ranks, _ = evaluate_retrieval(
fig_index, queries, targets, None,
variants_list=variants_list, filter_lists=filter_lists,
boost_lists=boost_lists, depth=depth)
row += [rec(ex_ranks, 1), rec(ex_ranks, 5), rec(ex_ranks, 20)]
results[(name, style)] = row
cols = ["R@1", "R@5", "R@20", "R@50", f"R@{depth}"]
if args.rerank:
cols += ["rr@1", "rr@5", "rr@20"]
if use_search_path:
tag = "ex" if args.expand else ("af" if args.author_filter else "ab")
cols += [f"{tag}@1", f"{tag}@5", f"{tag}@20"]
header = f"{'index':24s} {'style':10s} " + " ".join(f"{c:>7s}" for c in cols)
lines = [header, "-" * len(header)]
for name in args.indexes:
for style in STYLES:
vals = results[(name, style)]
lines.append(f"{name:24s} {style:10s} " + " ".join(f"{v:7.3f}" for v in vals))
ncols = min(len(results[(name, s)]) for s in STYLES)
avg = [sum(results[(name, s)][i] for s in STYLES) / len(STYLES) for i in range(ncols)]
lines.append(f"{name:24s} {'MEAN':10s} " + " ".join(f"{v:7.3f}" for v in avg))
lines.append("")
out = "\n".join(lines)
print(out)
Path("multi_eval_report.txt").write_text(out)
print("Report written to multi_eval_report.txt")
if __name__ == "__main__":
main()
|