File size: 5,574 Bytes
c500926
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
"""
run_eval.py - benchmark eval: COCO FID + CLIP Score via torchmetrics.

Protocol (matches the Tiny-T2I leaderboard requirements):
  - FID: torchmetrics.image.fid.FrechetInceptionDistance (InceptionV3,
    2048-dim pool3 features). Real set: n COCO val2014 images (256x256
    center-crop) from sayakpaul/coco-30-val-2014, rows 0..n-1 of the
    stream — disjoint by image hash from the training set (see
    fetch_coco_subset.py). Generated set: model output at native 64x64
    for those same n captions.
  - CLIP Score: torchmetrics.multimodal.CLIPScore with
    openai/clip-vit-base-patch32 (the default), generated image vs the
    caption that produced it.

Usage (after fetch_coco_subset.py has populated --work):
  python eval/run_eval.py --work ../pm-work --model model.png --n 5000
"""

import argparse
import json
import os
import sys
import time

import numpy as np
import torch
from PIL import Image

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from model import NATIVE_RES, coord_features, decode_pixels, encode_prompt, load_model, prompts_to_embeddings  # noqa: E402


def generate(model_path: str, captions, out_dir: str, batch: int = 64):
    os.makedirs(out_dir, exist_ok=True)
    weights = load_model(model_path)
    feats = coord_features(NATIVE_RES)
    t0 = time.time()
    for start in range(0, len(captions), batch):
        chunk = captions[start:start + batch]
        with torch.no_grad():
            emb = prompts_to_embeddings(chunk)
            z = encode_prompt(weights, emb)
            rgb = decode_pixels(weights, z, feats)
        arr = (rgb.reshape(len(chunk), NATIVE_RES, NATIVE_RES, 3).numpy()
               * 255).clip(0, 255).astype(np.uint8)
        for j in range(len(chunk)):
            Image.fromarray(arr[j], mode="RGB").save(
                os.path.join(out_dir, f"gen_{start + j:05d}.png"))
        if (start // batch) % 20 == 0:
            print(f"  gen {start + len(chunk)}/{len(captions)}  "
                  f"({time.time() - t0:.0f}s)", flush=True)
    print(f"  generated {len(captions)} images @ {NATIVE_RES}x{NATIVE_RES} "
          f"in {time.time() - t0:.0f}s", flush=True)


def load_batch(paths):
    imgs = [np.array(Image.open(p).convert("RGB"), dtype=np.uint8) for p in paths]
    return torch.from_numpy(np.stack(imgs)).permute(0, 3, 1, 2)   # (B,3,H,W) uint8


def compute_fid(real_dir: str, gen_dir: str, n: int, batch: int = 32) -> float:
    from torchmetrics.image.fid import FrechetInceptionDistance
    fid = FrechetInceptionDistance(feature=2048, normalize=False)
    t0 = time.time()
    for label, dir_, real in (("real", real_dir, True), ("gen", gen_dir, False)):
        files = sorted(os.listdir(dir_))[:n]
        for start in range(0, len(files), batch):
            imgs = load_batch([os.path.join(dir_, f) for f in files[start:start + batch]])
            fid.update(imgs, real=real)
            if (start // batch) % 25 == 0:
                print(f"  fid/{label}: {start + imgs.shape[0]}/{len(files)}  "
                      f"({time.time() - t0:.0f}s)", flush=True)
    return float(fid.compute())


def compute_clip_score(gen_dir: str, captions, batch: int = 32):
    from torchmetrics.multimodal import CLIPScore
    metric = CLIPScore(model_name_or_path="openai/clip-vit-base-patch32")
    files = sorted(os.listdir(gen_dir))[:len(captions)]
    t0 = time.time()
    for start in range(0, len(files), batch):
        imgs = load_batch([os.path.join(gen_dir, f) for f in files[start:start + batch]])
        metric.update(imgs, captions[start:start + imgs.shape[0]])
        if (start // batch) % 25 == 0:
            print(f"  clip: {start + imgs.shape[0]}/{len(files)}  "
                  f"({time.time() - t0:.0f}s)", flush=True)
    return float(metric.compute())


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--work",  required=True, help="dir from fetch_coco_subset.py")
    p.add_argument("--model", default="model.png")
    p.add_argument("--n",     type=int, default=5000)
    p.add_argument("--skip-gen", action="store_true")
    p.add_argument("--skip-fid", action="store_true")
    args = p.parse_args()

    with open(os.path.join(args.work, "eval_captions.json"), encoding="utf-8") as f:
        captions = json.load(f)[:args.n]
    real_dir = os.path.join(args.work, "eval_real")
    gen_dir = os.path.join(args.work, "eval_gen")

    if not args.skip_gen:
        print(f"[1/3] generating {len(captions)} images from '{args.model}'...")
        generate(args.model, captions, gen_dir)
    fid = None
    if not args.skip_fid:
        print("[2/3] FID (torchmetrics.image.fid, InceptionV3 2048)...")
        fid = compute_fid(real_dir, gen_dir, args.n)
        print(f"FID = {fid:.4f}", flush=True)
    print("[3/3] CLIP Score (torchmetrics, openai/clip-vit-base-patch32)...")
    clip = compute_clip_score(gen_dir, captions)
    print(f"CLIP Score = {clip:.4f}  (cosine {clip / 100:.4f})")

    print(f"\nRESULTS  n={args.n}  native_res={NATIVE_RES}x{NATIVE_RES}")
    if fid is not None:
        print(f"  FID        = {fid:.2f}")
    print(f"  CLIP Score = {clip:.2f}")
    out_path = os.path.join(args.work, "eval_results.json")
    results = {"n": args.n, "native_resolution": f"{NATIVE_RES}x{NATIVE_RES}",
               "fid": fid, "clip_score": clip}
    if fid is None and os.path.exists(out_path):
        old = json.load(open(out_path))
        results["fid"] = old.get("fid")
    with open(out_path, "w") as f:
        json.dump(results, f, indent=2)


if __name__ == "__main__":
    main()