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()
|