Aryan Sethi
Claude Opus 5 (1M context)
Publish the dataset to the Hub and verify the round trip
529f79d Download src/hub/pull.py from Aryan006/cone-distance: direct link, hf CLI and curl.
- Browser
- Download file 2.99 kB
-
https://huggingface.co/Aryan006/cone-distance/resolve/main/src/hub/pull.py
- Command line
-
hf download hf://Aryan006/cone-distance/src/hub/pull.py
-
curl -L -o pull.py https://huggingface.co/Aryan006/cone-distance/resolve/main/src/hub/pull.py
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() | |