"""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 /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()