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