cone-distance / src /hub /pull.py
Aryan Sethi
Claude Opus 5 (1M context)
Publish the dataset to the Hub and verify the round trip
529f79d
Raw History Blame Contribute Delete
2.99 kB
"""Download the dataset from the Hugging Face Hub and make it trainable.
The inverse of push.py: fetch the repo, unpack the shards into the layout
Ultralytics expects, and rewrite data.yaml with an absolute path so the training
script works from any working directory.
Usage:
python -m src.hub.pull --repo-id <user>/cone-distance-v1 --dest data/hub
python -m src.train.train --data data/hub/data.yaml
"""
from __future__ import annotations
import argparse
import os
import tarfile
from pathlib import Path
import yaml
def extract_shards(dest: Path) -> None:
shard_dir = dest / "shards"
shards = sorted(shard_dir.glob("*.tar"))
if not shards:
raise SystemExit(f"no shards under {shard_dir}")
for shard in shards:
with tarfile.open(shard) as archive:
archive.extractall(dest)
print(f" extracted {shard.name}")
def rewrite_data_yaml(dest: Path) -> Path:
"""Point `path` at wherever the dataset actually landed."""
path = dest / "data.yaml"
config = yaml.safe_load(path.read_text())
config["path"] = str(dest.resolve())
with open(path, "w") as handle:
yaml.safe_dump(config, handle, sort_keys=False)
return path
def pull(repo_id: str, dest: Path, token: str | None, keep_shards: bool) -> None:
from huggingface_hub import snapshot_download
dest.mkdir(parents=True, exist_ok=True)
print(f"downloading {repo_id} -> {dest}")
snapshot_download(
repo_id=repo_id,
repo_type="dataset",
local_dir=str(dest),
token=token,
)
print("unpacking")
extract_shards(dest)
if not keep_shards:
for shard in (dest / "shards").glob("*.tar"):
shard.unlink()
print(" removed shards (pass --keep-shards to keep them)")
data_yaml = rewrite_data_yaml(dest)
for split in ("train", "val"):
images = list((dest / split / "images").iterdir())
print(f" {split}: {len(images)} images")
print(f"\ndata config: {data_yaml}")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--repo-id", required=True)
parser.add_argument("--dest", type=Path, default=Path("data/hub"))
parser.add_argument("--token", default=os.environ.get("HF_TOKEN"),
help="defaults to $HF_TOKEN, then the cached "
"huggingface-cli login. Prefer the environment: a "
"token passed on the command line is visible to "
"anyone who can run ps.")
parser.add_argument("--keep-shards", action="store_true",
help="keep the .tar files after extracting (doubles disk use)")
args = parser.parse_args()
pull(repo_id=args.repo_id, dest=args.dest, token=args.token,
keep_shards=args.keep_shards)
if __name__ == "__main__":
main()