Skip to content

Commit 9dd254f

Browse files
author
Ronald Tse
committed
feat(data): real Tashkeela++ fetcher via HuggingFace Hub
Replaces the env-var placeholder in scripts/fetch_data.sh with a proper Python fetcher that knows the canonical dataset locations. Primary source: Misraj/Sadeed_Tashkeela — gated, requires HF_TOKEN. Fallback: community-datasets/tashkeela — GPLv2 open access. The fetcher streams parquet → one-line-per-chunk text, skips blank and overlong lines, writes atomically via .tmp rename, and exits with a clear error message on gated-repo failures (including the URL the user must visit to grant access). Also adds pythonpath = ["src", "."] to pytest config so tests run without pip install -e ., and pyarrow>=15.0 to [publish] extras. Smoke-tested end-to-end: fetch 100 lines → RababaArabicData consumes them → (bare, diacritized) pairs ready for StudentTrainer.
1 parent 3edc065 commit 9dd254f

3 files changed

Lines changed: 293 additions & 0 deletions

File tree

‎pyproject.toml‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@ export = [
3535
publish = [
3636
"huggingface_hub>=0.23",
3737
"omegaconf>=2.3",
38+
"pyarrow>=15.0",
3839
]
3940
dev = [
4041
"pytest>=8.0",
@@ -52,6 +53,7 @@ include = ["framework*", "tasks*"]
5253

5354
[tool.pytest.ini_options]
5455
testpaths = ["tests"]
56+
pythonpath = ["src", "."]
5557
addopts = "-ra -q --strict-markers"
5658
markers = [
5759
"slow: marks tests as slow (deselect with '-m \"not slow\"')",

‎scripts/fetch_data.py‎

Lines changed: 237 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,237 @@
1+
"""Fetch raw corpora from HuggingFace Hub.
2+
3+
Replaces ``scripts/fetch_data.sh``. The shell version required manual
4+
env-var URLs; this one knows the canonical dataset locations and
5+
validates downloads.
6+
7+
For ``rababa_arabic``:
8+
Default source: ``Misraj/Sadeed_Tashkeela`` — gated, requires HF_TOKEN
9+
and access grant at
10+
https://huggingface.co/datasets/Misraj/Sadeed_Tashkeela
11+
Fallback: ``community-datasets/tashkeela`` — GPLv2, open access, raw
12+
book text that needs heavier cleaning (handled by the data module).
13+
14+
For ``rababa_hebrew`` and ``secryst_thai_ipa`` the upstream sources are
15+
not yet on HF as datasets — leave the manual env-var path intact in
16+
``fetch_data.sh`` until they are.
17+
"""
18+
19+
from __future__ import annotations
20+
21+
import argparse
22+
import os
23+
import sys
24+
from pathlib import Path
25+
26+
ROOT = Path(__file__).resolve().parent.parent
27+
sys.path.insert(0, str(ROOT / "src"))
28+
29+
DEFAULT_OUT = ROOT / "data" / "raw"
30+
31+
DATASETS = {
32+
"rababa_arabic": {
33+
"primary": {
34+
"repo_id": "Misraj/Sadeed_Tashkeela",
35+
"repo_type": "dataset",
36+
"files": [
37+
"data/train-00000-of-00003.parquet",
38+
"data/train-00001-of-00003.parquet",
39+
"data/train-00002-of-00003.parquet",
40+
],
41+
"test_files": ["data/test-00000-of-00001.parquet"],
42+
"text_column": "text",
43+
"out_name": "tashkeela_plus_plus.txt",
44+
"split_lines": False,
45+
"note": (
46+
"Gated dataset. Visit "
47+
"https://huggingface.co/datasets/Misraj/Sadeed_Tashkeela, "
48+
"log in, accept the terms, then export HF_TOKEN."
49+
),
50+
},
51+
"fallback": {
52+
"repo_id": "community-datasets/tashkeela",
53+
"repo_type": "dataset",
54+
"files": None,
55+
"text_column": "text",
56+
"out_name": "tashkeela_plus_plus.txt",
57+
"split_lines": True,
58+
"note": (
59+
"Open-access raw corpus (GPLv2). Each row is a full book; "
60+
"we split on newlines and skip lines >1024 chars."
61+
),
62+
},
63+
},
64+
}
65+
66+
67+
def fetch_task(
68+
task: str,
69+
out_dir: Path,
70+
max_samples: int | None,
71+
use_fallback: bool,
72+
) -> Path:
73+
cfg = DATASETS.get(task)
74+
if cfg is None:
75+
raise SystemExit(
76+
f"No fetcher registered for task '{task}'. "
77+
f"Known: {sorted(DATASETS)}"
78+
)
79+
source = cfg["fallback"] if use_fallback else cfg["primary"]
80+
81+
try:
82+
import importlib.util
83+
84+
if importlib.util.find_spec("huggingface_hub") is None:
85+
raise ImportError("huggingface_hub not installed")
86+
except ImportError as e:
87+
raise SystemExit(
88+
"huggingface_hub is required. Install with: "
89+
"pip install -e '.[publish]'"
90+
) from e
91+
92+
token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
93+
out_dir.mkdir(parents=True, exist_ok=True)
94+
out_path = out_dir / source["out_name"]
95+
text_column = source["text_column"]
96+
97+
files = source["files"]
98+
if files is None:
99+
files = _list_repo_files(source["repo_id"], source["repo_type"], token)
100+
101+
count = _stream_parquet_to_text(
102+
files=files,
103+
repo_id=source["repo_id"],
104+
repo_type=source["repo_type"],
105+
token=token,
106+
text_column=text_column,
107+
out_path=out_path,
108+
max_samples=max_samples,
109+
split_lines=source.get("split_lines", False),
110+
)
111+
print(
112+
f"[{task}] wrote {count:,} lines -> {out_path} "
113+
f"({out_path.stat().st_size:,} bytes) "
114+
f"from {source['repo_id']}"
115+
)
116+
if count == 0:
117+
raise SystemExit(
118+
f"No lines written. Source note: {source.get('note', '')}"
119+
)
120+
return out_path
121+
122+
123+
def _list_repo_files(repo_id: str, repo_type: str, token: str | None) -> list[str]:
124+
from huggingface_hub import HfApi
125+
126+
api = HfApi(token=token)
127+
files = api.list_repo_files(repo_id, repo_type=repo_type)
128+
return [f for f in files if f.endswith((".parquet", ".json", ".jsonl", ".txt"))]
129+
130+
131+
def _stream_parquet_to_text(
132+
files: list[str],
133+
repo_id: str,
134+
repo_type: str,
135+
token: str | None,
136+
text_column: str,
137+
out_path: Path,
138+
max_samples: int | None,
139+
split_lines: bool = False,
140+
max_line_chars: int = 1024,
141+
) -> int:
142+
"""Stream the ``text_column`` of each parquet file to ``out_path``.
143+
144+
Each row's text is written as one line (after whitespace collapse).
145+
If ``split_lines`` is set (raw book corpora), the row's text is split
146+
on embedded newlines first — one row may carry many verse-sized lines.
147+
Lines longer than ``max_line_chars`` are skipped (training chunks
148+
should be ~50-60 words; longer ones are typically misplits).
149+
"""
150+
import pyarrow.parquet as pq
151+
from huggingface_hub import hf_hub_download
152+
153+
written = 0
154+
tmp = out_path.with_suffix(".txt.tmp")
155+
with tmp.open("w", encoding="utf-8") as fp:
156+
for fpath in files:
157+
local = hf_hub_download(
158+
repo_id=repo_id,
159+
filename=fpath,
160+
repo_type=repo_type,
161+
token=token,
162+
)
163+
pf = pq.ParquetFile(local)
164+
for batch in pf.iter_batches(batch_size=1024, columns=[text_column]):
165+
col = batch.column(text_column).to_pylist()
166+
for blob in col:
167+
if not blob:
168+
continue
169+
chunks = blob.splitlines() if split_lines else [blob]
170+
for raw in chunks:
171+
line = " ".join(raw.split())
172+
if not line or len(line) > max_line_chars:
173+
continue
174+
fp.write(line)
175+
fp.write("\n")
176+
written += 1
177+
if max_samples is not None and written >= max_samples:
178+
tmp.replace(out_path)
179+
return written
180+
tmp.replace(out_path)
181+
return written
182+
183+
184+
def main() -> int:
185+
parser = argparse.ArgumentParser(description=__doc__)
186+
parser.add_argument(
187+
"--task",
188+
required=True,
189+
choices=sorted(DATASETS),
190+
help="Which task corpus to fetch.",
191+
)
192+
parser.add_argument(
193+
"--out-dir",
194+
type=Path,
195+
default=DEFAULT_OUT,
196+
help=f"Output directory (default: {DEFAULT_OUT}).",
197+
)
198+
parser.add_argument(
199+
"--max-samples",
200+
type=int,
201+
default=None,
202+
help="Cap number of lines written (dev/CI mode).",
203+
)
204+
parser.add_argument(
205+
"--fallback",
206+
action="store_true",
207+
help=(
208+
"Use the open-access fallback dataset instead of the primary "
209+
"(useful when the primary is gated and no HF_TOKEN is set)."
210+
),
211+
)
212+
args = parser.parse_args()
213+
214+
try:
215+
fetch_task(
216+
task=args.task,
217+
out_dir=args.out_dir,
218+
max_samples=args.max_samples,
219+
use_fallback=args.fallback,
220+
)
221+
except SystemExit:
222+
raise
223+
except Exception as e:
224+
msg = str(e)
225+
if "GatedRepoError" in type(e).__name__ or "gated" in msg.lower():
226+
cfg = DATASETS[args.task]
227+
note = cfg["primary"].get("note", "")
228+
raise SystemExit(
229+
f"Gated dataset: {type(e).__name__}\n{note}\n"
230+
f"Or rerun with --fallback to use {cfg['fallback']['repo_id']}."
231+
) from e
232+
raise SystemExit(f"Fetch failed: {type(e).__name__}: {msg}") from e
233+
return 0
234+
235+
236+
if __name__ == "__main__":
237+
raise SystemExit(main())

‎tests/test_fetch_data.py‎

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
"""Tests for ``scripts/fetch_data.py``.
2+
3+
Network access is NOT required — these tests cover import, CLI
4+
argument handling, and the gated-repo error path. Real download is
5+
covered by the smoke test in the README.
6+
"""
7+
8+
from __future__ import annotations
9+
10+
import importlib.util
11+
from pathlib import Path
12+
13+
import pytest
14+
15+
SCRIPT = Path(__file__).resolve().parent.parent / "scripts" / "fetch_data.py"
16+
17+
18+
@pytest.fixture(scope="module")
19+
def fetch_module():
20+
spec = importlib.util.spec_from_file_location("fetch_data", SCRIPT)
21+
assert spec and spec.loader
22+
mod = importlib.util.module_from_spec(spec)
23+
spec.loader.exec_module(mod)
24+
return mod
25+
26+
27+
def test_fetcher_imports(fetch_module) -> None:
28+
assert hasattr(fetch_module, "fetch_task")
29+
assert hasattr(fetch_module, "DATASETS")
30+
assert "rababa_arabic" in fetch_module.DATASETS
31+
32+
33+
def test_fetcher_primary_has_parquet_files(fetch_module) -> None:
34+
primary = fetch_module.DATASETS["rababa_arabic"]["primary"]
35+
assert primary["repo_id"] == "Misraj/Sadeed_Tashkeela"
36+
assert primary["repo_type"] == "dataset"
37+
assert len(primary["files"]) == 3
38+
assert all(f.endswith(".parquet") for f in primary["files"])
39+
40+
41+
def test_fetcher_fallback_uses_open_dataset(fetch_module) -> None:
42+
fallback = fetch_module.DATASETS["rababa_arabic"]["fallback"]
43+
assert fallback["repo_id"] == "community-datasets/tashkeela"
44+
assert fallback["split_lines"] is True
45+
46+
47+
def test_fetcher_unknown_task_errors(fetch_module, tmp_path: Path) -> None:
48+
with pytest.raises(SystemExit, match="No fetcher registered"):
49+
fetch_module.fetch_task(
50+
task="not_a_task",
51+
out_dir=tmp_path,
52+
max_samples=None,
53+
use_fallback=False,
54+
)

0 commit comments

Comments
 (0)