File size: 4,367 Bytes
abb3324
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""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:

    <dataroot>/annotations/instances_train2017.json
    <dataroot>/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()