linxin02 commited on
Commit
771ae23
·
verified ·
1 Parent(s): 141e91e

Upload viz/preference/pref_score.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. viz/preference/pref_score.py +75 -0
viz/preference/pref_score.py ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os, glob, json, sys
2
+ import ImageReward
3
+
4
+ DPG = "/mnt/localssd/ELLA_dpg/dpg_bench/prompts"
5
+
6
+ def prompt_for(idx):
7
+ p = os.path.join(DPG, f"{idx}.txt")
8
+ with open(p) as f:
9
+ return f.read().strip()
10
+
11
+ def run_domain(name, nogan_dir, gan_dir, model):
12
+ # indices present in BOTH dirs as {idx}_0.png and having a DPG prompt
13
+ def idxs(d):
14
+ out = set()
15
+ for f in glob.glob(os.path.join(d, "*_0.png")):
16
+ stem = os.path.basename(f)[:-len("_0.png")]
17
+ out.add(stem)
18
+ return out
19
+ common = sorted(idxs(nogan_dir) & idxs(gan_dir),
20
+ key=lambda x: (0, int(x)) if x.isdigit() else (1, x))
21
+ rows = []
22
+ for idx in common:
23
+ if not os.path.exists(os.path.join(DPG, f"{idx}.txt")):
24
+ continue
25
+ prompt = prompt_for(idx)
26
+ ng_img = os.path.join(nogan_dir, f"{idx}_0.png")
27
+ g_img = os.path.join(gan_dir, f"{idx}_0.png")
28
+ s_ng = float(model.score(prompt, ng_img))
29
+ s_g = float(model.score(prompt, g_img))
30
+ rows.append({"idx": idx, "nogan": s_ng, "gan": s_g, "delta": s_g - s_ng})
31
+ return rows
32
+
33
+ def summarize(name, rows):
34
+ n = len(rows)
35
+ wins = sum(1 for r in rows if r["gan"] > r["nogan"])
36
+ ties = sum(1 for r in rows if r["gan"] == r["nogan"])
37
+ losses = n - wins - ties
38
+ mean_ng = sum(r["nogan"] for r in rows) / n
39
+ mean_g = sum(r["gan"] for r in rows) / n
40
+ mean_d = sum(r["delta"] for r in rows) / n
41
+ return {
42
+ "domain": name, "n": n, "wins": wins, "ties": ties, "losses": losses,
43
+ "win_rate": 100.0 * wins / n,
44
+ "win_rate_excl_ties": 100.0 * wins / (wins + losses) if (wins + losses) else None,
45
+ "mean_nogan": mean_ng, "mean_gan": mean_g, "mean_delta": mean_d,
46
+ }
47
+
48
+ def main():
49
+ model = ImageReward.load("ImageReward-v1.0")
50
+ results = {"model": "ImageReward-v1.0", "domains": {}}
51
+
52
+ pixel_rows = run_domain(
53
+ "pixel",
54
+ "/mnt/localssd/eval_work/spectrum/gen/baseline",
55
+ "/mnt/localssd/eval_work/spectrum/gen/gan_dinov2_t035",
56
+ model)
57
+ results["domains"]["pixel"] = {"summary": summarize("pixel", pixel_rows),
58
+ "rows": pixel_rows}
59
+ print("PIXEL done:", json.dumps(results["domains"]["pixel"]["summary"]), flush=True)
60
+
61
+ latent_rows = run_domain(
62
+ "latent",
63
+ "/mnt/localssd/eval_work_sana/sana_blip3o_sft/imgs_50000",
64
+ "/mnt/localssd/eval_work_sana/sana_blip3o_gan/imgs_50000",
65
+ model)
66
+ results["domains"]["latent"] = {"summary": summarize("latent", latent_rows),
67
+ "rows": latent_rows}
68
+ print("LATENT done:", json.dumps(results["domains"]["latent"]["summary"]), flush=True)
69
+
70
+ with open("/mnt/localssd/pref_results.json", "w") as f:
71
+ json.dump(results, f, indent=2)
72
+ print("WROTE /mnt/localssd/pref_results.json", flush=True)
73
+
74
+ if __name__ == "__main__":
75
+ main()