Add SWIFT coding agent probe experiment scripts
This commit is contained in:
Executable
+84
@@ -0,0 +1,84 @@
|
||||
#!/usr/bin/env python3
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gzip
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
|
||||
|
||||
DEFAULT_REPO_NAME = "ti_coding_agent_training_probe_20260624"
|
||||
|
||||
|
||||
def resolve_dataset_id(raw: str, token: str | None, endpoint: str | None) -> str:
|
||||
if "/" in raw:
|
||||
return raw
|
||||
if not token:
|
||||
raise SystemExit(
|
||||
"HF_DATASET_REPO_ID must be owner/name when HF_TOKEN is not set. "
|
||||
f"Got unqualified repo name: {raw}"
|
||||
)
|
||||
owner = HfApi(token=token, endpoint=endpoint).whoami()["name"]
|
||||
return f"{owner}/{raw}"
|
||||
|
||||
|
||||
def gunzip_if_needed(src: Path, dst: Path) -> None:
|
||||
if dst.exists() and dst.stat().st_size > 0:
|
||||
return
|
||||
dst.parent.mkdir(parents=True, exist_ok=True)
|
||||
with gzip.open(src, "rb") as fin, dst.open("wb") as fout:
|
||||
shutil.copyfileobj(fin, fout, length=1024 * 1024)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--dataset-id", default=os.environ.get("HF_DATASET_REPO_ID", DEFAULT_REPO_NAME))
|
||||
parser.add_argument("--raw-dir", default="data/raw/training_probe")
|
||||
parser.add_argument("--out-dir", default="data/processed/training_probe")
|
||||
args = parser.parse_args()
|
||||
|
||||
token = os.environ.get("HF_TOKEN")
|
||||
endpoint = os.environ.get("HF_ENDPOINT")
|
||||
dataset_id = resolve_dataset_id(args.dataset_id, token, endpoint)
|
||||
|
||||
raw_dir = Path(args.raw_dir)
|
||||
out_dir = Path(args.out_dir)
|
||||
raw_dir.mkdir(parents=True, exist_ok=True)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
local_path = snapshot_download(
|
||||
repo_id=dataset_id,
|
||||
repo_type="dataset",
|
||||
local_dir=raw_dir,
|
||||
token=token,
|
||||
endpoint=endpoint,
|
||||
allow_patterns=[
|
||||
"README.md",
|
||||
"metadata.json",
|
||||
"train.parquet",
|
||||
"validation.parquet",
|
||||
"train.jsonl.gz",
|
||||
"validation.jsonl.gz",
|
||||
],
|
||||
)
|
||||
local = Path(local_path)
|
||||
for name in ("train", "validation"):
|
||||
gz = local / f"{name}.jsonl.gz"
|
||||
if gz.exists():
|
||||
gunzip_if_needed(gz, out_dir / f"{name}.jsonl")
|
||||
parquet = local / f"{name}.parquet"
|
||||
if parquet.exists():
|
||||
target = out_dir / f"{name}.parquet"
|
||||
if not target.exists():
|
||||
target.symlink_to(parquet.resolve())
|
||||
print(f"DATASET_ID={dataset_id}")
|
||||
print(f"TRAIN_JSONL={out_dir / 'train.jsonl'}")
|
||||
print(f"VALIDATION_JSONL={out_dir / 'validation.jsonl'}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
Reference in New Issue
Block a user