PixelModel-v5 / eval_v5.py
TobiasLogic's picture
PixelModel v5: same architecture, 36x the data
a012ee3 verified
Raw
History Blame Contribute Delete
4.89 kB
import argparse, json, os, sys, numpy as np, torch
from diffusers import AutoencoderKL
from transformers import CLIPTextModel, CLIPTokenizer, CLIPModel
sys.path.insert(0, '/root')
from dit import DiT
SCALE = 0.18215
def _img_feats(self, pixel_values=None, **kw):
return self.visual_projection(self.vision_model(pixel_values=pixel_values).pooler_output)
def _txt_feats(self, input_ids=None, attention_mask=None, **kw):
return self.text_projection(self.text_model(input_ids=input_ids,
attention_mask=attention_mask).pooler_output)
CLIPModel.get_image_features = _img_feats
CLIPModel.get_text_features = _txt_feats
@torch.no_grad()
def sample(model, seq, pool, ns, npool, steps, cfg, dev):
B = seq.shape[0]
x = torch.randn(B, 4, 32, 32, device=dev)
a, b = ns.expand(B, -1, -1), npool.expand(B, -1)
dt = 1.0 / steps
for i in range(steps):
t = torch.full((B,), i * dt, device=dev)
with torch.autocast('cuda', dtype=torch.bfloat16):
vc = model(x, t, seq, pool)
vu = model(x, t, a, b)
x = x + (vu + cfg * (vc - vu)).float() * dt
return x
@torch.no_grad()
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--ckpt', default='/root/runs/pm5/best.pt')
ap.add_argument('--data', default='/root/pm5eval/eval_256.npz')
ap.add_argument('--n', type=int, default=5000)
ap.add_argument('--batch', type=int, default=100)
ap.add_argument('--steps', type=int, default=50)
ap.add_argument('--cfgs', default='2,3,4,5,6')
ap.add_argument('--out', default='/root/pm5_eval.json')
args = ap.parse_args()
dev = 'cuda'
ck = torch.load(args.ckpt, map_location=dev)
model = DiT(dim=384, depth=12, heads=6).to(dev).eval()
model.load_state_dict(ck['ema'])
print(f"[eval] ckpt step {ck.get('step')} val {ck.get('val')}", flush=True)
vae = AutoencoderKL.from_pretrained('stabilityai/sd-vae-ft-mse').to(dev).half().eval()
tok = CLIPTokenizer.from_pretrained('openai/clip-vit-base-patch32')
txt = CLIPTextModel.from_pretrained('openai/clip-vit-base-patch32').to(dev).half().eval()
e = tok([''], padding='max_length', max_length=40, truncation=True, return_tensors='pt').to(dev)
o = txt(**e)
ns, npool = o.last_hidden_state.float(), o.pooler_output.float()
d = np.load(args.data, allow_pickle=True)
real = d['images'][:args.n]
caps = [str(x) for x in d['captions'][:args.n]]
n = len(caps)
print(f'[eval] {n} real images and captions', flush=True)
from torchmetrics.image.fid import FrechetInceptionDistance
from torchmetrics.multimodal.clip_score import CLIPScore
ref = CLIPScore(model_name_or_path='openai/clip-vit-base-patch32').to(dev)
for i in range(0, n, args.batch):
rb = torch.from_numpy(real[i:i + args.batch]).permute(0, 3, 1, 2).to(dev)
ref.update(rb, caps[i:i + args.batch])
real_clip = float(ref.compute().item())
print(f'[eval] REAL images CLIP score {real_clip:.2f} (metric sanity check / ceiling)', flush=True)
del ref
torch.cuda.empty_cache()
results = []
for cfg in [float(x) for x in args.cfgs.split(',')]:
fid = FrechetInceptionDistance(feature=2048, normalize=True).to(dev)
clip = CLIPScore(model_name_or_path='openai/clip-vit-base-patch32').to(dev)
for i in range(0, n, args.batch):
rb = torch.from_numpy(real[i:i + args.batch].astype(np.float32) / 255.0).permute(0, 3, 1, 2).to(dev)
fid.update(rb, real=True)
for i in range(0, n, args.batch):
cb = caps[i:i + args.batch]
t = tok(cb, padding='max_length', max_length=40, truncation=True, return_tensors='pt').to(dev)
oo = txt(**t)
z = sample(model, oo.last_hidden_state.float(), oo.pooler_output.float(),
ns, npool, args.steps, cfg, dev)
img = vae.decode((z / SCALE).half()).sample.float()
img = (img.clamp(-1, 1) + 1) / 2
fid.update(img, real=False)
clip.update((img * 255).to(torch.uint8), cb)
f = float(fid.compute().item()); c = float(clip.compute().item())
results.append({'cfg': cfg, 'fid': round(f, 2), 'clip_score': round(c, 2),
'n': n, 'steps': args.steps})
print(f'[eval] cfg {cfg}: FID {f:.2f} CLIP {c:.2f}', flush=True)
del fid, clip
torch.cuda.empty_cache()
best = min(results, key=lambda r: r['fid'])
json.dump({'ckpt_step': ck.get('step'), 'val_loss': ck.get('val'),
'real_image_clip_score': round(real_clip, 2),
'results': results, 'best_fid': best}, open(args.out, 'w'), indent=2)
print(f"[eval] BEST FID {best['fid']} at cfg {best['cfg']} (CLIP {best['clip_score']})", flush=True)
print('EVALDONE', flush=True)
if __name__ == '__main__':
main()