foto / fetch_authors.py
htohfa's picture
Upload 4 files
c51da3c verified
Raw
History Blame Contribute Delete
4.06 kB
"""Fetch author lists for every paper in the caption slice.
ArxivCap does not ship author names, so the author filter needs a one-time
metadata pass keyed on the arXiv IDs already in the slice. Uses the Semantic
Scholar batch endpoint (500 papers per request, ~110 requests for the astro
slice). Results append to authors.jsonl as they arrive, so an interrupted run
resumes, and the final parquet is rebuilt from that file.
export S2_API_KEY=... # optional but much faster
python fetch_authors.py
Output: authors.parquet with columns (arxiv_id, authors, author_ids).
"""
import argparse
import json
import os
import time
from pathlib import Path
import duckdb
import pyarrow as pa
import pyarrow.parquet as pq
import requests
BATCH_URL = "https://api.semanticscholar.org/graph/v1/paper/batch"
BATCH_SIZE = 500
def load_ids(slice_path: str) -> list[str]:
con = duckdb.connect()
rows = con.execute(
f"SELECT DISTINCT arxiv_id FROM read_parquet('{slice_path}') WHERE arxiv_id IS NOT NULL"
).fetchall()
return [r[0] for r in rows]
def fetch_batch(ids: list[str], api_key: str) -> list:
headers = {"x-api-key": api_key} if api_key else {}
payload = {"ids": [f"ARXIV:{i}" for i in ids]}
for attempt in range(6):
try:
r = requests.post(BATCH_URL, params={"fields": "authors"},
json=payload, headers=headers, timeout=60)
if r.status_code == 429:
wait = 10 * (attempt + 1)
print(f" rate limited, sleeping {wait}s")
time.sleep(wait)
continue
r.raise_for_status()
return r.json()
except requests.exceptions.RequestException as e:
if attempt == 5:
raise
time.sleep(5 * (attempt + 1))
return []
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--slice", default="astro_captions.parquet")
ap.add_argument("--out", default="authors.parquet")
ap.add_argument("--cache", default="authors.jsonl")
ap.add_argument("--sleep", type=float, default=1.0,
help="seconds between requests; raise if you see 429s")
args = ap.parse_args()
api_key = os.environ.get("S2_API_KEY", "")
if not api_key:
print("No S2_API_KEY set. This will work but is slower and rate limited.")
cache_path = Path(args.cache)
done = set()
if cache_path.exists():
for line in cache_path.open():
done.add(json.loads(line)["arxiv_id"])
ids = load_ids(args.slice)
todo = [i for i in ids if i not in done]
print(f"{len(ids):,} papers in slice, {len(done):,} already fetched, {len(todo):,} to go")
with cache_path.open("a") as f:
for start in range(0, len(todo), BATCH_SIZE):
batch = todo[start:start + BATCH_SIZE]
results = fetch_batch(batch, api_key)
for arxiv_id, entry in zip(batch, results):
authors, author_ids = [], []
if entry and entry.get("authors"):
for a in entry["authors"]:
if a.get("name"):
authors.append(a["name"])
author_ids.append(a.get("authorId") or "")
f.write(json.dumps({"arxiv_id": arxiv_id, "authors": authors,
"author_ids": author_ids}) + "\n")
f.flush()
print(f" {min(start + BATCH_SIZE, len(todo)):,}/{len(todo):,}")
time.sleep(args.sleep)
rows = [json.loads(l) for l in cache_path.open()]
found = sum(1 for r in rows if r["authors"])
pq.write_table(pa.table({
"arxiv_id": [r["arxiv_id"] for r in rows],
"authors": [r["authors"] for r in rows],
"author_ids": [r["author_ids"] for r in rows],
}), args.out)
print(f"\n{len(rows):,} records written to {args.out}")
print(f"{found:,} with at least one author ({100*found/max(len(rows),1):.1f}%)")
if __name__ == "__main__":
main()