PixelModel-v3 / eval /run_eval.py
TobiasLogic's picture
PixelModel v3: SIREN+FiLM CPPN, 919K params, beats v1 FID
63a1291 verified
Raw
History Blame Contribute Delete
5.73 kB
"""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()