mirror of
https://github.com/qdrant/landing_page.git
synced 2026-09-30 00:18:32 +02:00
Improve category filter and scale corpus to match notebook
This commit is contained in:
+19
-10
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user