Add arxiv-complete retrieval layer (semantic + keyword)
This commit is contained in:
@@ -0,0 +1,136 @@
|
||||
"""
|
||||
Data access layer for the arxiv-complete metadata dataset.
|
||||
|
||||
The full metadata parquet is 1.64 GB / 3,148,796 rows (one per arXiv paper),
|
||||
stored as a single file on Hugging Face:
|
||||
|
||||
https://huggingface.co/api/datasets/secemp9/arxiv-complete/parquet/metadata/train/0.parquet
|
||||
|
||||
Loading the whole thing into RAM (~a few GB of Python objects) is feasible on
|
||||
a big machine but wasteful for a search prototype. Instead we stream rows from
|
||||
the remote parquet with polars' *scan_parquet* + lazy filtering, so the query
|
||||
path only materializes the handful of rows that actually match.
|
||||
|
||||
The same lazy-scanning approach works against a local copy of the parquet,
|
||||
which is what the build pipeline uses to precompute embeddings.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Iterable
|
||||
|
||||
import polars as pl
|
||||
|
||||
DEFAULT_PARQUET_URL = (
|
||||
"https://huggingface.co/api/datasets/secemp9/arxiv-complete/"
|
||||
"parquet/metadata/train/0.parquet"
|
||||
)
|
||||
|
||||
# Columns that are cheap to keep in memory and useful for display / filtering.
|
||||
DISPLAY_COLUMNS = [
|
||||
"paper_id",
|
||||
"title",
|
||||
"primary_category",
|
||||
"first_version_date",
|
||||
"arxiv_abs_url",
|
||||
"abstract",
|
||||
]
|
||||
|
||||
# Free-text terms that mark an abstract as a survey / overview / taxonomy paper.
|
||||
SURVEY_KEYWORDS = ("survey", "comprehensive", "overview", "taxonomy")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Paper:
|
||||
"""A single arXiv paper, as surfaced by the search layer."""
|
||||
|
||||
paper_id: str
|
||||
title: str
|
||||
primary_category: str
|
||||
year: int
|
||||
url: str
|
||||
abstract: str = ""
|
||||
score: float = 0.0
|
||||
|
||||
@staticmethod
|
||||
def from_row(row: pl.Row) -> "Paper":
|
||||
raw = dict(zip(DISPLAY_COLUMNS, row))
|
||||
fvd = raw["first_version_date"]
|
||||
year = (
|
||||
fvd.year
|
||||
if isinstance(fvd, datetime)
|
||||
else (int(fvd) if fvd is not None else 0)
|
||||
)
|
||||
return Paper(
|
||||
paper_id=str(raw["paper_id"]),
|
||||
title=str(raw["title"] or ""),
|
||||
primary_category=str(raw["primary_category"] or ""),
|
||||
year=year,
|
||||
url=str(raw["arxiv_abs_url"] or ""),
|
||||
abstract=str(raw["abstract"] or ""),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataStore:
|
||||
"""
|
||||
Lazy view over the arxiv-complete metadata parquet.
|
||||
|
||||
We deliberately do not load the full 3.15M rows into memory. ``scan_parquet``
|
||||
returns a lazy frame that polars only executes when we call ``collect()`` --
|
||||
and it pushes down the WHERE filters (year range, category, survey) so the
|
||||
download/shuffle work stays proportional to the result set.
|
||||
"""
|
||||
|
||||
source: str = field(default_factory=lambda: os.environ.get(
|
||||
"ARXIV_METADATA_PARQUET", DEFAULT_PARQUET_URL
|
||||
))
|
||||
|
||||
def _scan(self) -> pl.LazyFrame:
|
||||
return pl.scan_parquet(self.source)
|
||||
|
||||
def scan_filtered(
|
||||
self,
|
||||
*,
|
||||
year_min: int | None = None,
|
||||
year_max: int | None = None,
|
||||
primary_category: str | None = None,
|
||||
survey_only: bool = False,
|
||||
) -> pl.LazyFrame:
|
||||
"""
|
||||
Return a lazy frame of the full metadata filtered to the given criteria.
|
||||
|
||||
``survey_only`` keeps only rows whose abstract mentions a survey/overview
|
||||
keyword. All filters are AND-ed together and pushed into the scan.
|
||||
"""
|
||||
q = self._scan()
|
||||
|
||||
if year_min is not None:
|
||||
q = q.filter(pl.col("first_version_date").dt.year() >= year_min)
|
||||
if year_max is not None:
|
||||
q = q.filter(pl.col("first_version_date").dt.year() <= year_max)
|
||||
if primary_category:
|
||||
q = q.filter(pl.col("primary_category") == primary_category)
|
||||
if survey_only:
|
||||
needle = "|".join(SURVEY_KEYWORDS)
|
||||
q = q.filter(
|
||||
pl.col("abstract")
|
||||
.str.to_lowercase()
|
||||
.str.contains(needle, strict=False)
|
||||
)
|
||||
return q
|
||||
|
||||
def fetch(self, limit: int | None = None) -> list[Paper]:
|
||||
"""Execute the (already filtered) scan and return Paper objects."""
|
||||
df = self.scan_filtered().select(DISPLAY_COLUMNS)
|
||||
if limit is not None:
|
||||
df = df.limit(limit)
|
||||
rows = df.collect().to_dicts()
|
||||
return [Paper.from_row(r.values()) for r in rows]
|
||||
|
||||
def count(self) -> int:
|
||||
"""Total rows matching the current filter (cheap COUNT push-down)."""
|
||||
return int(self.scan_filtered().select(pl.len()).collect()["len"][0])
|
||||
Reference in New Issue
Block a user