|
| 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()) |
0 commit comments