| """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() |
|
|