""" Step 4: Extract expert-curated reading lists from survey papers in unarXive 2024. This creates a proper evaluation dataset for paper recommendation. It replaces the weak citation-based eval (cited=2, co-cited=1, not cited=0) with an expert-curated reading list eval based on survey paper reference sections. WHY SURVEY PAPERS? Survey papers are gold-standard reading lists because: 1. The author spent months curating what to include and exclude 2. Citations are grouped by topic (section structure = topic clustering) 3. Citation ordering within sections encodes importance/chronology 4. Papers in Related Work vs Methods have different relationships 5. Papers NOT cited are expert-level exclusions (hard negatives) DATA SOURCE: ines-besrour/unarxive_2024 on HuggingFace - 2.28M full-text arXiv papers (1991-2024) - Each paper: metadata, sections with cite_spans, bib_entries with IDs - MIT licensed LABEL SCHEME: 4 = cited multiple times in Related Work (essential reading) 3 = cited in Related Work or mentioned heavily across paper (relevant) 2 = cited in Methods/Experiments (tool/resource/comparison) 1 = cited in Introduction only (background context) 0 = same category, NOT cited by survey (hard negative — expert chose to exclude) OUTPUT: eval_survey_reading_lists.parquet — for evaluation (surveys from 2023+) train_survey_reading_lists.parquet — for training better models (surveys pre-2023) eval_metadata.json — stats and documentation Columns: survey_arxiv_id, survey_title, cited_arxiv_id, section_name, section_type, section_index, paragraph_position, citation_count_in_section, total_mentions, label, is_hard_negative USAGE: # Full extraction (needs ~50GB disk, ~2-3 hours on CPU): python 04_extract_survey_reading_lists.py --max-surveys 1000 --min-citations 20 # Quick test (first 100 surveys): python 04_extract_survey_reading_lists.py --max-surveys 100 --min-citations 15 PREREQUISITES: pip install datasets pyarrow numpy tqdm huggingface_hub INTEGRATION: This script produces eval data that scripts/05_evaluate_on_surveys.py consumes. Together they replace the old eval approach in scripts/03_train_lightgbm.py. Author: ResearchIT ML Pipeline — Eval V2 """ from __future__ import annotations import argparse import json import os import re import sys import time from collections import defaultdict from pathlib import Path import numpy as np import pyarrow as pa import pyarrow.parquet as pq from tqdm import tqdm # ── Section Classification ─────────────────────────────────────────────────── RELATED_WORK_PATTERNS = [ r"related\s*work", r"literature\s*review", r"background", r"prior\s*work", r"previous\s*work", r"state\s*of\s*the\s*art", r"existing\s*(approaches|methods|work)", r"survey\s*of", ] METHODS_PATTERNS = [ r"method(ology|s)?", r"approach", r"proposed\s*(method|approach|framework)", r"our\s*(method|approach|framework)", r"model\s*(architecture|design)", r"framework", r"implementation", ] INTRO_PATTERNS = [r"introduction", r"overview", r"motivation"] EXPERIMENT_PATTERNS = [ r"experiment(s|al)?", r"evaluation", r"results", r"empirical", r"benchmark", r"comparison", ] def classify_section(section_name: str) -> str: """Classify a section into a semantic category.""" name_lower = section_name.lower().strip() for pattern in RELATED_WORK_PATTERNS: if re.search(pattern, name_lower): return "related_work" for pattern in METHODS_PATTERNS: if re.search(pattern, name_lower): return "methods" for pattern in INTRO_PATTERNS: if re.search(pattern, name_lower): return "introduction" for pattern in EXPERIMENT_PATTERNS: if re.search(pattern, name_lower): return "experiments" return "other" # ── Survey Detection ───────────────────────────────────────────────────────── SURVEY_TITLE_PATTERNS = [ r"\bsurvey\b", r"\breview\b", r"\boverview\b", r"\btutorial\b", r"\bcomprehensive\s+(survey|review|overview|study)\b", r"\bsystematic\s+(review|survey|study)\b", r"\bstate[\s-]of[\s-]the[\s-]art\b", r"\brecent\s+(advances|developments|progress)\b", ] def is_survey_paper(metadata: dict) -> bool: """Detect if a paper is a CS survey/review based on title and category.""" title = (metadata.get("title") or "").lower().replace("\n", " ") categories = (metadata.get("categories") or "") # Must be CS if not any(cat.strip().startswith("cs") for cat in categories.split()): return False for pattern in SURVEY_TITLE_PATTERNS: if re.search(pattern, title): return True return False # ── arXiv ID Extraction (multi-source) ─────────────────────────────────────── ARXIV_ID_REGEX = re.compile(r"(?:arXiv[:\s/]*)?(\d{4}\.\d{4,5}(?:v\d+)?)", re.IGNORECASE) ARXIV_DOI_REGEX = re.compile(r"10\.48550/arxiv\.(\d{4}\.\d{4,5})", re.IGNORECASE) ARXIV_OLD_FORMAT = re.compile(r"(?:arXiv[:\s/]*)?((?:astro-ph|cond-mat|cs|gr-qc|hep-ex|hep-lat|hep-ph|hep-th|math-ph|math|nlin|nucl-ex|nucl-th|physics|q-bio|q-fin|quant-ph|stat)/\d{7})", re.IGNORECASE) def extract_arxiv_id_from_bib_entry(bib_entry: dict) -> str | None: """ Extract arXiv ID from a bib entry using all available sources: 1. ids.arxiv_id field (direct, most reliable) 2. contained_arXiv_ids list 3. ids.doi field (when DOI is 10.48550/arxiv.XXXX.XXXXX) 4. bib_entry_raw text (regex search for arXiv:XXXX.XXXXX patterns) Returns normalized arXiv ID (without version suffix) or None. """ # Source 1: Direct arxiv_id field ids = bib_entry.get("ids", {}) arxiv_id = ids.get("arxiv_id", "") if arxiv_id and arxiv_id.strip(): return _normalize_arxiv_id(arxiv_id.strip()) # Source 2: contained_arXiv_ids contained = bib_entry.get("contained_arXiv_ids", []) if contained: return _normalize_arxiv_id(contained[0]) # Source 3: DOI containing arxiv doi = ids.get("doi", "") if doi: match = ARXIV_DOI_REGEX.search(doi) if match: return _normalize_arxiv_id(match.group(1)) # Source 4: Raw bib entry text raw = bib_entry.get("bib_entry_raw", "") if raw: # Try modern format first (YYMM.NNNNN) match = ARXIV_ID_REGEX.search(raw) if match: candidate = match.group(1) # Validate it looks like a real arXiv ID (not a random number) if re.match(r"\d{4}\.\d{4,5}", candidate): return _normalize_arxiv_id(candidate) # Try old format (category/NNNNNNN) match = ARXIV_OLD_FORMAT.search(raw) if match: return match.group(1) return None def _normalize_arxiv_id(arxiv_id: str) -> str: """Normalize arXiv ID: remove version suffix, lowercase category prefix.""" # Remove version: "2303.14957v2" → "2303.14957" normalized = re.sub(r"v\d+$", "", arxiv_id.strip()) # Remove prefix if present normalized = re.sub(r"^arXiv:", "", normalized, flags=re.IGNORECASE) return normalized # ── Citation Extraction ────────────────────────────────────────────────────── def extract_citations_from_paper(paper_data: dict) -> list[dict]: """ Extract all citations with section-level annotations. Uses multi-source arXiv ID extraction for higher resolution. """ sections = paper_data.get("sections", {}) bib_entries = paper_data.get("bib_entries", {}) # Build ref_id → arxiv_id mapping using all sources ref_to_arxiv: dict[str, str] = {} for ref_id, bib_entry in bib_entries.items(): arxiv_id = extract_arxiv_id_from_bib_entry(bib_entry) if arxiv_id: ref_to_arxiv[ref_id] = arxiv_id if not ref_to_arxiv: return [] # Track mentions across the paper total_mentions: dict[str, int] = defaultdict(int) section_mentions: dict[str, dict[str, int]] = defaultdict(lambda: defaultdict(int)) all_citations = [] for section_idx, (section_name, section_data) in enumerate(sections.items()): cite_spans = section_data.get("cite_spans", []) text = section_data.get("text", "") text_len = max(len(text), 1) for cite_span in cite_spans: ref_id = cite_span.get("ref_id", "") if ref_id not in ref_to_arxiv: continue cited_arxiv_id = ref_to_arxiv[ref_id] total_mentions[cited_arxiv_id] += 1 section_mentions[section_name][cited_arxiv_id] += 1 start_pos = cite_span.get("start", 0) paragraph_position = start_pos / text_len all_citations.append({ "cited_arxiv_id": cited_arxiv_id, "section_name": section_name, "section_type": classify_section(section_name), "section_index": section_idx, "paragraph_position": paragraph_position, }) # Enrich with counts for citation in all_citations: cited_id = citation["cited_arxiv_id"] section_name = citation["section_name"] citation["citation_count_in_section"] = section_mentions[section_name][cited_id] citation["total_mentions"] = total_mentions[cited_id] return all_citations def assign_relevance_label(citation: dict) -> int: """ Assign relevance tier based on section and citation frequency. 4 = essential (cited heavily in Related Work) 3 = relevant (appears in Related Work or mentioned heavily) 2 = tool/method (cited in Methods/Experiments) 1 = background (Introduction or passing mention elsewhere) """ section_type = citation["section_type"] count_in_section = citation["citation_count_in_section"] total_mentions = citation["total_mentions"] if section_type == "related_work": if total_mentions >= 3 or count_in_section >= 2: return 4 # Essential: heavily discussed in related work return 3 # Relevant: appears in related work if section_type in ("methods", "experiments"): return 2 # Tool/method: used in their approach or evaluation if section_type == "introduction": if total_mentions >= 2: return 3 # Important enough to mention multiple times return 1 # Background context only # "other" sections if total_mentions >= 3: return 3 # Heavily cited overall return 1 # Passing mention # ── Data Loading (handles multiple formats) ────────────────────────────────── def load_papers_from_jsonl(file_path: str, max_papers: int = None): """ Generator that yields papers from a JSONL file. Handles both single-paper-per-line and pretty-printed JSON. """ import gzip opener = gzip.open if file_path.endswith(".gz") else open with opener(file_path, "rt", encoding="utf-8", errors="replace") as f: count = 0 buffer = "" for line in f: line = line.strip() if not line: continue # Try as JSONL (one paper per line) try: paper = json.loads(line) yield paper count += 1 if max_papers and count >= max_papers: return continue except json.JSONDecodeError: pass # Accumulate for multi-line JSON buffer += line try: paper = json.loads(buffer) yield paper count += 1 buffer = "" if max_papers and count >= max_papers: return except json.JSONDecodeError: continue def load_papers_streaming_hf(max_papers: int = None): """ Load papers from HuggingFace using the datasets library (streaming). Falls back to file-based loading if streaming fails. """ try: from datasets import load_dataset ds = load_dataset("ines-besrour/unarxive_2024", split="train", streaming=True) count = 0 for paper in ds: yield paper count += 1 if max_papers and count >= max_papers: return except Exception as e: print(f"HF streaming failed: {e}") print("Please download the dataset manually and use --input-dir") sys.exit(1) # ── Main Pipeline ──────────────────────────────────────────────────────────── def process_papers( paper_source, max_surveys: int = 500, min_citations: int = 20, ) -> list[dict]: """ Process papers from any source and extract survey reading lists. """ surveys_found = [] papers_scanned = 0 for paper in tqdm(paper_source, desc="Scanning for surveys"): papers_scanned += 1 if len(surveys_found) >= max_surveys: break # Handle metadata (could be dict or JSON string) metadata = paper.get("metadata", {}) if isinstance(metadata, str): try: metadata = json.loads(metadata) except json.JSONDecodeError: continue paper_id = metadata.get("id") or paper.get("paper_id", "") if not paper_id: continue # Check if survey if not is_survey_paper(metadata): continue # Extract citations citations = extract_citations_from_paper(paper) arxiv_citations = [c for c in citations if c["cited_arxiv_id"]] unique_cited = set(c["cited_arxiv_id"] for c in arxiv_citations) if len(unique_cited) < min_citations: continue # Deduplicate: keep highest-relevance per cited paper best_per_paper: dict[str, dict] = {} for citation in arxiv_citations: cited_id = citation["cited_arxiv_id"] label = assign_relevance_label(citation) citation["label"] = label if cited_id not in best_per_paper or label > best_per_paper[cited_id]["label"]: best_per_paper[cited_id] = citation surveys_found.append({ "survey_arxiv_id": paper_id, "survey_title": (metadata.get("title") or "").replace("\n", " ").strip(), "survey_categories": metadata.get("categories", ""), "survey_update_date": metadata.get("update_date", ""), "citations": list(best_per_paper.values()), "total_unique_cited": len(unique_cited), "total_bib_entries": len(paper.get("bib_entries", {})), "arxiv_resolution_rate": len(unique_cited) / max(len(paper.get("bib_entries", {})), 1), }) if len(surveys_found) % 50 == 0: print(f" Found {len(surveys_found)} surveys after scanning {papers_scanned} papers") print(f"\nDone! Found {len(surveys_found)} surveys after scanning {papers_scanned} papers") return surveys_found def build_eval_dataset( surveys: list[dict], hard_negatives_per_query: int = 30, seed: int = 42, ) -> pa.Table: """ Build eval dataset with hard negatives from cross-survey mining. Hard negatives = papers cited by OTHER surveys in the same field but NOT cited by this survey. These represent expert-level exclusions: the survey author was aware of these papers and chose not to include them. """ rng = np.random.default_rng(seed) # Build category → cited papers mapping for hard negative mining category_papers: dict[str, set[str]] = defaultdict(set) for survey in surveys: cats = survey["survey_categories"].split() primary_cat = cats[0] if cats else "cs" for citation in survey["citations"]: category_papers[primary_cat].add(citation["cited_arxiv_id"]) rows = { "survey_arxiv_id": [], "survey_title": [], "cited_arxiv_id": [], "section_name": [], "section_type": [], "section_index": [], "paragraph_position": [], "citation_count_in_section": [], "total_mentions": [], "label": [], "is_hard_negative": [], } for survey in surveys: survey_id = survey["survey_arxiv_id"] survey_title = survey["survey_title"] cats = survey["survey_categories"].split() primary_cat = cats[0] if cats else "cs" # Add positive citations (the reading list) cited_set = set() for citation in survey["citations"]: cited_set.add(citation["cited_arxiv_id"]) rows["survey_arxiv_id"].append(survey_id) rows["survey_title"].append(survey_title) rows["cited_arxiv_id"].append(citation["cited_arxiv_id"]) rows["section_name"].append(citation["section_name"]) rows["section_type"].append(citation["section_type"]) rows["section_index"].append(citation["section_index"]) rows["paragraph_position"].append(float(citation["paragraph_position"])) rows["citation_count_in_section"].append(citation["citation_count_in_section"]) rows["total_mentions"].append(citation["total_mentions"]) rows["label"].append(citation["label"]) rows["is_hard_negative"].append(False) # Add hard negatives: same category, cited by OTHER surveys, not this one available_negatives = category_papers.get(primary_cat, set()) - cited_set - {survey_id} if available_negatives: neg_pool = list(available_negatives) n_neg = min(hard_negatives_per_query, len(neg_pool)) selected_negatives = rng.choice(neg_pool, size=n_neg, replace=False) for neg_id in selected_negatives: rows["survey_arxiv_id"].append(survey_id) rows["survey_title"].append(survey_title) rows["cited_arxiv_id"].append(str(neg_id)) rows["section_name"].append("") rows["section_type"].append("") rows["section_index"].append(-1) rows["paragraph_position"].append(-1.0) rows["citation_count_in_section"].append(0) rows["total_mentions"].append(0) rows["label"].append(0) rows["is_hard_negative"].append(True) table = pa.table({ "survey_arxiv_id": pa.array(rows["survey_arxiv_id"], type=pa.string()), "survey_title": pa.array(rows["survey_title"], type=pa.string()), "cited_arxiv_id": pa.array(rows["cited_arxiv_id"], type=pa.string()), "section_name": pa.array(rows["section_name"], type=pa.string()), "section_type": pa.array(rows["section_type"], type=pa.string()), "section_index": pa.array(rows["section_index"], type=pa.int32()), "paragraph_position": pa.array(rows["paragraph_position"], type=pa.float32()), "citation_count_in_section": pa.array(rows["citation_count_in_section"], type=pa.int32()), "total_mentions": pa.array(rows["total_mentions"], type=pa.int32()), "label": pa.array(rows["label"], type=pa.int32()), "is_hard_negative": pa.array(rows["is_hard_negative"], type=pa.bool_()), }) return table def time_split_surveys( surveys: list[dict], eval_cutoff_year: int = 2023, ) -> tuple[list[dict], list[dict]]: """ Time-split surveys into train/eval by publication year. Prevents temporal leakage: eval surveys cite newer papers. """ train = [] eval_set = [] for survey in surveys: date_str = survey.get("survey_update_date", "") try: year = int(date_str[:4]) except (ValueError, TypeError, IndexError): year = 2020 if year >= eval_cutoff_year: eval_set.append(survey) else: train.append(survey) return train, eval_set # ── CLI ────────────────────────────────────────────────────────────────────── def main(): parser = argparse.ArgumentParser( description="Extract survey paper reading lists from unarXive for evaluation" ) parser.add_argument("--input-dir", default=None, help="Directory with unarXive JSONL files (skip HF download)") parser.add_argument("--output-dir", default="./eval_v2_data") parser.add_argument("--max-surveys", type=int, default=500) parser.add_argument("--min-citations", type=int, default=20) parser.add_argument("--hard-negatives-per-query", type=int, default=30) parser.add_argument("--eval-cutoff-year", type=int, default=2023) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--push-to-hub", action="store_true", help="Push results to siddhm11/researchit-reranker-data") args = parser.parse_args() output_dir = Path(args.output_dir) output_dir.mkdir(parents=True, exist_ok=True) # Load papers if args.input_dir: # File-based loading (for HF job with downloaded data) input_path = Path(args.input_dir) files = sorted(input_path.glob("*.jsonl")) + sorted(input_path.glob("*.jsonl.gz")) if not files: print(f"ERROR: No JSONL files found in {input_path}") sys.exit(1) print(f"Found {len(files)} JSONL files in {input_path}") def paper_generator(): for fp in files: yield from load_papers_from_jsonl(str(fp)) all_surveys = process_papers(paper_generator(), args.max_surveys, args.min_citations) else: # HuggingFace streaming print("Streaming from ines-besrour/unarxive_2024...") paper_source = load_papers_streaming_hf() all_surveys = process_papers(paper_source, args.max_surveys, args.min_citations) if not all_surveys: print("ERROR: No survey papers found!") sys.exit(1) # Time split train_surveys, eval_surveys = time_split_surveys(all_surveys, args.eval_cutoff_year) print(f"\nTime split (cutoff={args.eval_cutoff_year}):") print(f" Train surveys: {len(train_surveys)}") print(f" Eval surveys: {len(eval_surveys)}") # Build and save datasets if eval_surveys: print(f"\nBuilding EVAL dataset ({len(eval_surveys)} surveys)...") eval_table = build_eval_dataset(eval_surveys, args.hard_negatives_per_query, args.seed) eval_file = output_dir / "eval_survey_reading_lists.parquet" pq.write_table(eval_table, str(eval_file), compression="snappy") print(f" Saved: {eval_file} ({len(eval_table)} rows)") if train_surveys: print(f"\nBuilding TRAIN dataset ({len(train_surveys)} surveys)...") train_table = build_eval_dataset(train_surveys, args.hard_negatives_per_query, args.seed + 1) train_file = output_dir / "train_survey_reading_lists.parquet" pq.write_table(train_table, str(train_file), compression="snappy") print(f" Saved: {train_file} ({len(train_table)} rows)") # Stats all_table = build_eval_dataset(all_surveys, args.hard_negatives_per_query, args.seed + 2) labels = all_table.column("label").to_pylist() label_dist = defaultdict(int) for l in labels: label_dist[l] += 1 n_surveys = len(set(all_table.column("survey_arxiv_id").to_pylist())) print(f"\n{'='*60}") print(f"EVAL FRAMEWORK SUMMARY") print(f"{'='*60}") print(f" Total surveys: {n_surveys}") print(f" Total rows: {len(labels)}") print(f" Avg papers/survey: {len(labels) / max(n_surveys, 1):.1f}") print(f"\n Label distribution:") print(f" 4 (essential): {label_dist[4]} ({100*label_dist[4]/max(len(labels),1):.1f}%)") print(f" 3 (related work): {label_dist[3]} ({100*label_dist[3]/max(len(labels),1):.1f}%)") print(f" 2 (methods/tools): {label_dist[2]} ({100*label_dist[2]/max(len(labels),1):.1f}%)") print(f" 1 (background): {label_dist[1]} ({100*label_dist[1]/max(len(labels),1):.1f}%)") print(f" 0 (hard negative): {label_dist[0]} ({100*label_dist[0]/max(len(labels),1):.1f}%)") # Resolution stats avg_resolution = np.mean([s["arxiv_resolution_rate"] for s in all_surveys]) print(f"\n arXiv ID resolution rate: {100*avg_resolution:.1f}%") # Save metadata meta = { "eval_version": "v2.0", "description": "Survey-curated reading lists for paper recommendation evaluation", "data_source": "ines-besrour/unarxive_2024", "eval_cutoff_year": args.eval_cutoff_year, "num_surveys_total": n_surveys, "num_surveys_train": len(train_surveys), "num_surveys_eval": len(eval_surveys), "total_rows": len(labels), "hard_negatives_per_query": args.hard_negatives_per_query, "min_citations": args.min_citations, "arxiv_resolution_rate": float(avg_resolution), "label_scheme": { "4": "Essential reading (cited heavily in Related Work)", "3": "Relevant (cited in Related Work or mentioned heavily)", "2": "Tool/method (cited in Methods/Experiments)", "1": "Background (cited in Introduction only)", "0": "Hard negative (same category, expert chose NOT to cite)", }, "label_distribution": dict(label_dist), "seed": args.seed, } meta_file = output_dir / "eval_metadata.json" with open(meta_file, "w") as f: json.dump(meta, f, indent=2) print(f"\n Metadata: {meta_file}") # Push to Hub if args.push_to_hub: print("\nPushing to siddhm11/researchit-reranker-data...") from huggingface_hub import HfApi api = HfApi() for fname in output_dir.iterdir(): api.upload_file( path_or_fileobj=str(fname), path_in_repo=f"eval_v2/{fname.name}", repo_id="siddhm11/researchit-reranker-data", repo_type="dataset", ) print(f" Uploaded: eval_v2/{fname.name}") print(f"\n✅ Done!") if __name__ == "__main__": main()