File size: 6,223 Bytes
e465a2f | 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 | """
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, device: torch.device, batch: int = 64):
os.makedirs(out_dir, exist_ok=True)
weights = load_model(model_path)
weights = {name: value.to(device) for name, value in weights.items()}
feats = coord_features(NATIVE_RES).to(device)
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).to(device)
z = encode_prompt(weights, emb)
rgb = decode_pixels(weights, z, feats)
arr = (rgb.reshape(len(chunk), NATIVE_RES, NATIVE_RES, 3).cpu().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, device: torch.device, batch: int = 32) -> float:
from torchmetrics.image.fid import FrechetInceptionDistance
fid = FrechetInceptionDistance(feature=2048, normalize=False).to(device)
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.to(device), 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, device: torch.device, batch: int = 32):
from torchmetrics.multimodal import CLIPScore
metric = CLIPScore(model_name_or_path="openai/clip-vit-base-patch32").to(device)
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.to(device), 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("--device", default="auto", help="auto, cpu, cuda, or a PyTorch device string")
p.add_argument("--skip-gen", action="store_true")
p.add_argument("--skip-fid", action="store_true")
args = p.parse_args()
device = torch.device("cuda" if args.device == "auto" and torch.cuda.is_available()
else "cpu" if args.device == "auto" else args.device)
print(f"device: {device}")
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, device)
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, device)
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, device)
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()
|