Upload viz/preference/pref_score2.py with huggingface_hub
Browse files
viz/preference/pref_score2.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os, glob, json, math
|
| 2 |
+
import torch
|
| 3 |
+
from PIL import Image
|
| 4 |
+
|
| 5 |
+
DPG = "/mnt/localssd/ELLA_dpg/dpg_bench/prompts"
|
| 6 |
+
DEVICE = "cuda"
|
| 7 |
+
|
| 8 |
+
DOMAINS = {
|
| 9 |
+
"pixel": ("/mnt/localssd/eval_work/pref_pixel/baseline",
|
| 10 |
+
"/mnt/localssd/eval_work/pref_pixel/gan_dinov2_text"),
|
| 11 |
+
"latent": ("/mnt/localssd/eval_work_sana/sana_blip3o_sft/imgs_50000",
|
| 12 |
+
"/mnt/localssd/eval_work_sana/sana_blip3o_gan/imgs_50000"),
|
| 13 |
+
}
|
| 14 |
+
|
| 15 |
+
def common_idxs(nogan, gan):
|
| 16 |
+
def idxs(d):
|
| 17 |
+
return {os.path.basename(f)[:-len("_0.png")] for f in glob.glob(os.path.join(d,"*_0.png"))}
|
| 18 |
+
c = idxs(nogan) & idxs(gan)
|
| 19 |
+
c = {i for i in c if os.path.exists(os.path.join(DPG, i + ".txt"))}
|
| 20 |
+
return sorted(c, key=lambda x: (0, int(x)) if x.isdigit() else (1, x))
|
| 21 |
+
|
| 22 |
+
def prompt_for(i):
|
| 23 |
+
return open(os.path.join(DPG, i + ".txt")).read().strip()
|
| 24 |
+
|
| 25 |
+
import ImageReward
|
| 26 |
+
ir = ImageReward.load("ImageReward-v1.0", device=DEVICE)
|
| 27 |
+
def ir_score(p, path):
|
| 28 |
+
return float(ir.score(p, path))
|
| 29 |
+
|
| 30 |
+
from transformers import AutoProcessor, AutoModel
|
| 31 |
+
pk_proc = AutoProcessor.from_pretrained('laion/CLIP-ViT-H-14-laion2B-s32B-b79K')
|
| 32 |
+
pk_model = AutoModel.from_pretrained('yuvalkirstain/PickScore_v1').eval().to(DEVICE)
|
| 33 |
+
@torch.no_grad()
|
| 34 |
+
def pk_score(p, path):
|
| 35 |
+
img = Image.open(path).convert("RGB")
|
| 36 |
+
ii = pk_proc(images=img, return_tensors="pt").to(DEVICE)
|
| 37 |
+
ti = pk_proc(text=p, padding=True, truncation=True, max_length=77, return_tensors="pt").to(DEVICE)
|
| 38 |
+
ie = pk_model.get_image_features(**ii); ie = ie / ie.norm(dim=-1, keepdim=True)
|
| 39 |
+
te = pk_model.get_text_features(**ti); te = te / te.norm(dim=-1, keepdim=True)
|
| 40 |
+
return (pk_model.logit_scale.exp() * (te @ ie.T)).item()
|
| 41 |
+
|
| 42 |
+
def binom_two_sided(k, n):
|
| 43 |
+
if n == 0:
|
| 44 |
+
return 1.0
|
| 45 |
+
kk = min(k, n - k); ln2n = n * math.log(2)
|
| 46 |
+
tail = sum(math.exp(math.lgamma(n+1) - math.lgamma(i+1) - math.lgamma(n-i+1) - ln2n) for i in range(kk+1))
|
| 47 |
+
return min(1.0, 2 * tail)
|
| 48 |
+
|
| 49 |
+
MODELS = [("ImageReward-v1.0", ir_score), ("PickScore_v1", pk_score)]
|
| 50 |
+
|
| 51 |
+
results = {"domains": {}}
|
| 52 |
+
for dom, (ng, g) in DOMAINS.items():
|
| 53 |
+
idxs = common_idxs(ng, g)
|
| 54 |
+
results["domains"][dom] = {"n_pairs": len(idxs), "models": {}}
|
| 55 |
+
print(dom, "n_pairs", len(idxs), flush=True)
|
| 56 |
+
for mname, fn in MODELS:
|
| 57 |
+
rows = []
|
| 58 |
+
for i in idxs:
|
| 59 |
+
p = prompt_for(i)
|
| 60 |
+
s_ng = fn(p, os.path.join(ng, i + "_0.png"))
|
| 61 |
+
s_g = fn(p, os.path.join(g, i + "_0.png"))
|
| 62 |
+
rows.append((i, s_ng, s_g))
|
| 63 |
+
n = len(rows)
|
| 64 |
+
wins = sum(1 for _, a, b in rows if b > a)
|
| 65 |
+
ties = sum(1 for _, a, b in rows if b == a)
|
| 66 |
+
losses = n - wins - ties
|
| 67 |
+
mn = sum(a for _, a, _ in rows) / n
|
| 68 |
+
mg = sum(b for _, _, b in rows) / n
|
| 69 |
+
md = sum(b - a for _, a, b in rows) / n
|
| 70 |
+
eff = wins + losses
|
| 71 |
+
summ = {"n": n, "wins": wins, "ties": ties, "losses": losses,
|
| 72 |
+
"win_rate": 100 * wins / n,
|
| 73 |
+
"win_rate_excl_ties": 100 * wins / eff if eff else None,
|
| 74 |
+
"mean_nogan": mn, "mean_gan": mg, "mean_delta": md,
|
| 75 |
+
"binom_p_two_sided": binom_two_sided(wins, eff)}
|
| 76 |
+
results["domains"][dom]["models"][mname] = {
|
| 77 |
+
"summary": summ,
|
| 78 |
+
"rows": [{"idx": i, "nogan": a, "gan": b, "delta": b - a} for i, a, b in rows]}
|
| 79 |
+
print(dom, mname, json.dumps(summ), flush=True)
|
| 80 |
+
|
| 81 |
+
json.dump(results, open("/mnt/localssd/pref_results2.json", "w"), indent=2)
|
| 82 |
+
print("WROTE /mnt/localssd/pref_results2.json", flush=True)
|