forked from t348575/kvcache-experiments
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathscbench_download.py
More file actions
89 lines (70 loc) · 3.06 KB
/
Copy pathscbench_download.py
File metadata and controls
89 lines (70 loc) · 3.06 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
"""Download SCBench config parquet files from HuggingFace into a local dir.
SCBench ships one parquet per config (subtask) under <config>/test-*.parquet.
This mirrors how the Bailian traces live under dataset/bailian/: after running
this, dataset/scbench/<config>.parquet holds each subtask, ready for
scbench_stats.py and the replay harness.
Examples:
# All configs into the default dataset/scbench/
python scripts/scbench_download.py
# Just two configs
python scripts/scbench_download.py --configs scbench_kv scbench_repoqa_and_kv
"""
from __future__ import annotations
import argparse
import os
import shutil
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from common.dataset import SCBENCH_CONFIGS, SCBENCH_REPO # noqa: E402
def download_config(api, config, out_dir, token):
"""Fetch every test parquet shard for one config, concatenate names to disk.
Most configs are a single shard (test-00000-of-00001.parquet); the glob keeps
multi-shard configs working. Shard 0 is written as <config>.parquet and any
further shards as <config>.partNN.parquet so a later reader can glob them.
"""
from huggingface_hub import hf_hub_download
repo_files = api.list_repo_files(SCBENCH_REPO, repo_type="dataset", token=token)
shards = sorted(
f for f in repo_files
if f.startswith(f"{config}/") and f.endswith(".parquet")
)
if not shards:
raise SystemExit(f"No parquet files found for config {config!r} in {SCBENCH_REPO}")
written = []
for idx, remote in enumerate(shards):
local = hf_hub_download(
repo_id=SCBENCH_REPO, repo_type="dataset", filename=remote, token=token
)
if idx == 0:
dst = out_dir / f"{config}.parquet"
else:
dst = out_dir / f"{config}.part{idx:02d}.parquet"
shutil.copy(local, dst)
written.append(dst)
print(f" {remote} -> {dst} ({os.path.getsize(dst) // 1024} KB)")
return written
def build_parser():
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--out", default="dataset/scbench",
help="Output directory (default: dataset/scbench).")
p.add_argument("--configs", nargs="+", default=SCBENCH_CONFIGS,
choices=SCBENCH_CONFIGS,
help="Configs to download (default: all).")
p.add_argument("--token", default=os.getenv("HF_TOKEN"),
help="HF token (defaults to $HF_TOKEN; not required for this public dataset).")
return p
def main():
args = build_parser().parse_args()
out_dir = Path(args.out)
out_dir.mkdir(parents=True, exist_ok=True)
from huggingface_hub import HfApi
api = HfApi()
print(f"Downloading {len(args.configs)} SCBench config(s) from {SCBENCH_REPO} -> {out_dir}")
for config in args.configs:
print(f"[{config}]")
download_config(api, config, out_dir, args.token)
print("Done.")
if __name__ == "__main__":
main()