File size: 5,733 Bytes
63a1291
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
"""FID + CLIP Score for PixelModel v3 - identical protocol to v1.

Real set: MS-COCO val2014 pairs (coco_eval.npz from fetch_coco_subset.py),
256x256 center-crop, n=5000. Metrics via torchmetrics:
  FID  -> torchmetrics.image.fid.FrechetInceptionDistance
  CLIP -> torchmetrics.multimodal.CLIPScore, openai/clip-vit-base-patch32

Two modes, scored by the *same* code so numbers are comparable:

  # score v3 directly from model.png
  python eval/run_eval.py --arch v3 --work ../pm-work \
      --png model.png --config config.json --vocab vocab.json --n 5000

  # regression check: score any OTHER model (v1, v2) from pre-rendered PNGs
  # (files named 00000.png, 00001.png, ... aligned to eval order)
  python eval/run_eval.py --arch precomputed --work ../pm-work \
      --images-dir v1_render/ --n 5000

This is what makes the README's v1-vs-v3 comparison apples-to-apples: both go
through this script, on the same reals, with the same torchmetrics versions.
"""

from __future__ import annotations

import argparse
import glob
import json
import os
import sys

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 load_config, load_model_png, load_vocab, encode_caption, make_coord_grid


def resize_u8(arr_u8, size):
    return np.asarray(Image.fromarray(arr_u8, "RGB").resize((size, size), Image.BICUBIC),
                      dtype=np.uint8)


def render_v3(args, captions, device):
    cfg = load_config(args.config)
    model = load_model_png(args.png, cfg, map_location=device)
    vocab = load_vocab(args.vocab)
    coords = make_coord_grid(args.render_res, args.render_res, device=device,
                             dtype=torch.float32).unsqueeze(0)
    outs = []
    for i in range(0, len(captions), args.batch):
        batch_caps = captions[i:i + args.batch]
        toks = np.stack([encode_caption(c, vocab, cfg.max_tokens) for c in batch_caps])
        toks = torch.from_numpy(toks).long().to(device)
        c = coords.expand(len(batch_caps), -1, -1)
        with torch.no_grad():
            rgb = model(toks, c)
        rgb = (rgb.clamp(0, 1).reshape(len(batch_caps), args.render_res, args.render_res, 3)
               .cpu().numpy() * 255.0).round().astype(np.uint8)
        for im in rgb:
            outs.append(resize_u8(im, args.fid_size))
        if i % (args.batch * 20) == 0:
            print(f"[eval] rendered {i+len(batch_caps)}/{len(captions)}")
    return np.stack(outs)


def load_precomputed(images_dir, n, fid_size):
    files = sorted(glob.glob(os.path.join(images_dir, "*.png")))[:n]
    if not files:
        raise FileNotFoundError(f"no PNGs in {images_dir}")
    outs = [resize_u8(np.asarray(Image.open(f).convert("RGB"), dtype=np.uint8), fid_size)
            for f in files]
    return np.stack(outs)


def to_nchw_u8(arr):
    return torch.from_numpy(arr).permute(0, 3, 1, 2).contiguous()


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--arch", choices=["v3", "precomputed"], default="v3")
    ap.add_argument("--work", default="../pm-work", help="dir with coco_eval.npz")
    ap.add_argument("--png", default="model.png")
    ap.add_argument("--config", default="config.json")
    ap.add_argument("--vocab", default="vocab.json")
    ap.add_argument("--images-dir", default=None, help="precomputed renders")
    ap.add_argument("--n", type=int, default=5000)
    ap.add_argument("--render-res", type=int, default=128)
    ap.add_argument("--fid-size", type=int, default=256, help="match v1 (256 crop reals)")
    ap.add_argument("--batch", type=int, default=32)
    ap.add_argument("--clip-model", default="openai/clip-vit-base-patch32")
    ap.add_argument("--device", default="cuda")
    ap.add_argument("--out", default="eval_results.json")
    args = ap.parse_args()

    device = args.device if torch.cuda.is_available() else "cpu"
    from torchmetrics.image.fid import FrechetInceptionDistance
    from torchmetrics.multimodal.clip_score import CLIPScore

    ev = np.load(os.path.join(args.work, "coco_eval.npz"), allow_pickle=True)
    reals = ev["images"][:args.n]
    captions = [str(c) for c in ev["captions"][:args.n]]
    n = min(len(reals), len(captions), args.n)
    reals, captions = reals[:n], captions[:n]
    reals = np.stack([resize_u8(im, args.fid_size) for im in reals])
    print(f"[eval] arch={args.arch}  n={n}  fid_size={args.fid_size}  device={device}")

    if args.arch == "v3":
        fakes = render_v3(args, captions, device)
    else:
        fakes = load_precomputed(args.images_dir, n, args.fid_size)
        if len(fakes) != n:
            raise ValueError(f"precomputed count {len(fakes)} != {n}")

    fid = FrechetInceptionDistance(feature=2048, normalize=False).to(device)
    for i in range(0, n, args.batch):
        fid.update(to_nchw_u8(reals[i:i + args.batch]).to(device), real=True)
        fid.update(to_nchw_u8(fakes[i:i + args.batch]).to(device), real=False)
    fid_val = float(fid.compute())

    clip = CLIPScore(model_name_or_path=args.clip_model).to(device)
    for i in range(0, n, args.batch):
        imgs = to_nchw_u8(fakes[i:i + args.batch]).to(device)
        clip.update(imgs, captions[i:i + args.batch])
    clip_val = float(clip.compute())

    results = {
        "arch": args.arch, "n": n, "fid": round(fid_val, 2),
        "clip_score": round(clip_val, 2), "render_res": args.render_res,
        "fid_size": args.fid_size, "clip_model": args.clip_model,
    }
    with open(args.out, "w") as f:
        json.dump(results, f, indent=2)
    print(f"[eval] FID = {fid_val:.2f}   CLIP Score = {clip_val:.2f}   (n={n})")
    print(f"[eval] wrote {args.out}")


if __name__ == "__main__":
    main()