File size: 3,158 Bytes
6c311ad
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
import io
import json
import os
import random
import zipfile
from concurrent.futures import ThreadPoolExecutor

import numpy as np
import requests
from PIL import Image

ANN_URL = "http://images.cocodataset.org/annotations/annotations_trainval2014.zip"

def csr(img, size):
    img = img.convert("RGB")
    w, h = img.size
    s = min(w, h)
    l, t = (w - s) // 2, (h - s) // 2
    return np.asarray(img.crop((l, t, l + s, t + s)).resize((size, size), Image.BICUBIC), dtype=np.uint8)

def fetch_one(item, size):
    url, cap = item
    for _ in range(3):
        try:
            r = requests.get(url, timeout=15)
            if r.status_code == 200:
                return csr(Image.open(io.BytesIO(r.content)), size), cap
        except Exception:
            pass
    return None

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--work", default="/root/v6cache")
    ap.add_argument("--n", type=int, default=5000)
    ap.add_argument("--size", type=int, default=256)
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--workers", type=int, default=48)
    args = ap.parse_args()
    os.makedirs(args.work, exist_ok=True)

    out_path = os.path.join(args.work, "eval_256.npz")
    if os.path.exists(out_path):
        print(f"[eval-set] already exists at {out_path}", flush=True)
        return

    ann_path = os.path.join(args.work, "captions_val2014.json")
    if not os.path.exists(ann_path):
        print("[eval-set] downloading annotations", flush=True)
        z = os.path.join(args.work, "ann.zip")
        with requests.get(ANN_URL, stream=True, timeout=120) as r:
            with open(z, "wb") as f:
                for chunk in r.iter_content(1 << 20):
                    f.write(chunk)
        with zipfile.ZipFile(z) as zf:
            with zf.open("annotations/captions_val2014.json") as src, open(ann_path, "wb") as dst:
                dst.write(src.read())
        os.remove(z)

    ann = json.load(open(ann_path))
    url_by_id = {im["id"]: im["coco_url"] for im in ann["images"]}
    cap_by_id = {}
    for a in ann["annotations"]:
        cap_by_id.setdefault(a["image_id"], a["caption"])
    items = [(url_by_id[i], cap_by_id[i]) for i in cap_by_id if i in url_by_id]
    random.Random(args.seed).shuffle(items)
    print(f"[eval-set] {len(items)} val2014 pairs available, target {args.n}", flush=True)

    imgs, caps = [], []
    pool = ThreadPoolExecutor(max_workers=args.workers)
    idx, batch = 0, 64
    while len(imgs) < args.n and idx < len(items):
        chunk = items[idx:idx + batch]
        idx += batch
        results = [r for r in pool.map(lambda it: fetch_one(it, args.size), chunk) if r is not None]
        for a, c in results:
            imgs.append(a); caps.append(c)
        if idx % (batch * 20) == 0:
            print(f"[eval-set] {len(imgs)}/{args.n}", flush=True)

    imgs = np.stack(imgs[:args.n])
    caps = np.array(caps[:args.n], dtype=object)
    np.savez(out_path, images=imgs, captions=caps)
    print(f"[eval-set] DONE {imgs.shape} -> {out_path}", flush=True)

if __name__ == "__main__":
    main()