Improve category filter and scale corpus to match notebook

This commit is contained in:
Dylan Couzon
2026-05-06 16:54:52 -04:00
parent 81747bcbf2
commit 66264ac07c
@@ -28,27 +28,36 @@ Use Python <3.13. Not all dependencies support the newest Python versions yet.
## Dataset ## Dataset
You'll work with a sample of arXiv papers from the [`gfissore/arxiv-abstracts-2021`](https://huggingface.co/datasets/gfissore/arxiv-abstracts-2021) Hugging Face dataset, filtered to machine-learning and computer-science categories. Each paper has a title, an abstract, and category tags, which gives you four natural representations once the abstract is split into chunks: title, full abstract as a summary, abstract sentences as chunks, and categories as tags. You'll work with 20 000 arXiv papers from the [`gfissore/arxiv-abstracts-2021`](https://huggingface.co/datasets/gfissore/arxiv-abstracts-2021) Hugging Face dataset, filtered to ML/CS categories and to papers from 2018 onward — earlier ML papers predate most of the topics queries care about. Each paper has a title, an abstract, and category tags, which gives you four natural representations once the abstract is split into chunks: title, full abstract as a summary, abstract sentences as chunks, and categories as tags.
```python ```python
from datasets import load_dataset from datasets import load_dataset
ML_CATEGORIES = {"cs.LG", "cs.CV", "cs.CL", "cs.AI", "stat.ML"} ML_CATEGORIES = {"cs.LG", "cs.CV", "cs.CL", "cs.AI", "stat.ML"}
dataset = load_dataset( # Non-streaming so HF caches the parquet locally; first run downloads ~2.5 GB, re-runs are instant.
"gfissore/arxiv-abstracts-2021", split="train", streaming=True dataset = load_dataset("gfissore/arxiv-abstracts-2021", split="train")
)
papers = [] papers = []
for row in dataset: # IDs are roughly chronological; iterate from the end to land on 2021/2020/2019 papers first.
if len(papers) >= 2000: for i in range(len(dataset) - 1, -1, -1):
break # 2000 ML/CS papers is enough for this tutorial if len(papers) >= 20000:
break
row = dataset[i]
if not row["abstract"] or not row["title"]: if not row["abstract"] or not row["title"]:
continue continue
cats = list(row["categories"]) # categories arrive as space-joined strings (e.g. ["cs.LG cs.CV"]); split each entry.
cats = [tok for entry in row["categories"] for tok in entry.split()]
if not any(c in ML_CATEGORIES for c in cats): if not any(c in ML_CATEGORIES for c in cats):
continue # ML/CS papers only continue
# Year lives in the YYMM prefix of new-format arXiv IDs ("2104.01234" -> 2021).
arxiv_id = row["id"]
if "/" in arxiv_id or "." not in arxiv_id:
continue # skip pre-2007 IDs like "math/0506001"
if 2000 + int(arxiv_id[:2]) < 2018:
continue
papers.append({ papers.append({
"arxiv_id": row["id"], "arxiv_id": arxiv_id,
"title": row["title"].strip(), "title": row["title"].strip(),
"abstract": row["abstract"].strip(), "abstract": row["abstract"].strip(),
"categories": cats, "categories": cats,