"""Download only the COCO images that contain a stop sign. COCO train2017 is 19 GB. About 1,700 of its images contain a stop sign, and those are the only ones we want, so this reads the annotation file first and then fetches that subset directly -- a few hundred MB instead of nineteen GB. Lays the result out exactly as the real COCO distribution does, so src.data.extract_coco runs against it unchanged: /annotations/instances_train2017.json /train2017/*.jpg Usage: python scripts/fetch_coco_stopsigns.py --dataroot data/raw/coco """ from __future__ import annotations import argparse import json import urllib.request import zipfile from concurrent.futures import ThreadPoolExecutor from pathlib import Path ANNOTATIONS_URL = "http://images.cocodataset.org/annotations/annotations_trainval2017.zip" IMAGE_URL = "http://images.cocodataset.org/{split}/{file_name}" CATEGORY = "stop sign" def download(url: str, dest: Path) -> None: dest.parent.mkdir(parents=True, exist_ok=True) temporary = dest.with_suffix(dest.suffix + ".part") with urllib.request.urlopen(url, timeout=120) as response, open(temporary, "wb") as handle: while chunk := response.read(1 << 20): handle.write(chunk) temporary.rename(dest) def fetch_annotations(dataroot: Path, split: str) -> Path: path = dataroot / "annotations" / f"instances_{split}.json" if path.exists(): print(f"annotations already present: {path}") return path archive = dataroot / "annotations_trainval2017.zip" if not archive.exists(): print(f"downloading {ANNOTATIONS_URL} (~241 MB)") download(ANNOTATIONS_URL, archive) print(f"extracting instances_{split}.json") with zipfile.ZipFile(archive) as zf: zf.extract(f"annotations/instances_{split}.json", dataroot) archive.unlink() return path def fetch_images(annotations_path: Path, dataroot: Path, split: str, workers: int, limit: int | None) -> None: print("reading annotations (large file, takes a moment)") with open(annotations_path) as handle: coco = json.load(handle) category_ids = {c["id"] for c in coco["categories"] if c["name"] == CATEGORY} if not category_ids: raise SystemExit(f"no category named {CATEGORY!r}") wanted = {a["image_id"] for a in coco["annotations"] if a["category_id"] in category_ids} images = [i for i in coco["images"] if i["id"] in wanted] images.sort(key=lambda i: i["id"]) if limit is not None: images = images[:limit] print(f"{len(images)} images contain a {CATEGORY}") target_dir = dataroot / split target_dir.mkdir(parents=True, exist_ok=True) todo = [i for i in images if not (target_dir / i["file_name"]).exists()] print(f"{len(images) - len(todo)} already on disk, fetching {len(todo)}") def fetch(image: dict) -> None: download(IMAGE_URL.format(split=split, file_name=image["file_name"]), target_dir / image["file_name"]) failures = [] with ThreadPoolExecutor(max_workers=workers) as pool: for index, (image, future) in enumerate( [(i, pool.submit(fetch, i)) for i in todo], start=1 ): try: future.result() except Exception as error: failures.append((image["file_name"], error)) if index % 200 == 0: print(f" {index}/{len(todo)}") total = sum(p.stat().st_size for p in target_dir.glob("*.jpg")) print(f"\n{len(list(target_dir.glob('*.jpg')))} images, {total / 1e6:.0f} MB in {target_dir}") if failures: print(f"!! {len(failures)} downloads failed; re-run to retry") def main() -> None: parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) parser.add_argument("--dataroot", type=Path, default=Path("data/raw/coco")) parser.add_argument("--split", default="train2017") parser.add_argument("--workers", type=int, default=16) parser.add_argument("--limit", type=int, default=None) args = parser.parse_args() annotations = fetch_annotations(args.dataroot, args.split) fetch_images(annotations, args.dataroot, args.split, args.workers, args.limit) if __name__ == "__main__": main()