gopaljigaur/repro-cocoedit-bundle / scripts /cocoedit_gedit_subset_eval.py
gopaljigaur's picture
download
raw
6.92 kB
"""
Scaled reproduction of Claim 3 (Qwen-Image-Edit + CoCoEdit on GEdit-Bench-EN).
Runs the OFFICIAL inference recipe (test_gedit-bench.py) and OFFICIAL scorer
(geditbench/run_psnr_score.py) from the released repo/dataset back-to-back,
baseline vs. +CoCoEdit LoRA, on a random subset of GEdit-Bench-EN and prints
the PSNR/SSIM delta so it is directly comparable to the paper's claimed
+2.8 dB / +0.112 (Table 1, GEdit-Bench-EN, Qwen-Image-Edit row).
Intended to run on a Hugging Face GPU Job (1x A100-80GB is enough for the
Qwen-Image-Edit-2509 bf16 pipeline + LoRA). NOT executed in this session:
the Hub token used here lacks the `job.write` scope (see Claim 3 page for
the exact 403 and the canary command). Included in the reproduction bundle
as a ready-to-run artifact for whoever picks this up next with Jobs access.
hf jobs run --flavor a100-large -e HF_TOKEN \
ghcr.io/huggingface/diffusers-pytorch-cuda \
python cocoedit_gedit_subset_eval.py --n-samples 40
"""
import argparse
import json
import os
import random
import numpy as np
import torch
from huggingface_hub import hf_hub_download, snapshot_download
from PIL import Image
from skimage.metrics import structural_similarity as compare_ssim
SKIP_GROUPS = ["style_change", "tone_transfer", "subject-add"]
def read_mask(mask_path):
mask = Image.open(mask_path).convert("L")
mask = np.array(mask)
return (mask > 128).astype(np.float32)
def read_image(img_path):
img = Image.open(img_path).convert("RGB")
return np.array(img).astype(np.float32) / 255.0
def compute_psnr(img1, img2, mask):
if mask.sum() == 0:
return np.nan
mse = ((img1 - img2) ** 2 * mask[..., None]).sum() / (mask.sum() * 3)
return 100.0 if mse == 0 else 10 * np.log10(1.0 / mse)
def compute_ssim(img1, img2, mask):
total = 0.0
for c in range(3):
total += compare_ssim(
img1[..., c], img2[..., c], data_range=1.0, win_size=11,
gaussian_weights=True, use_sample_covariance=False, mask=mask,
)
return total / 3
def run_inference(pipeline, samples, bench_root, out_dir, steps, true_cfg_scale, guidance_scale):
os.makedirs(out_dir, exist_ok=True)
for key, item in samples.items():
ref_path = os.path.join(bench_root, item["ref_image_path"])
image = Image.open(ref_path).convert("RGB")
inputs = {
"image": image,
"prompt": item["caption"],
"generator": torch.manual_seed(0),
"true_cfg_scale": true_cfg_scale,
"guidance_scale": guidance_scale,
"negative_prompt": " ",
"num_inference_steps": steps,
}
with torch.inference_mode():
out = pipeline(**inputs).images[0]
save_path = os.path.join(out_dir, item["ref_image_path"])
os.makedirs(os.path.dirname(save_path), exist_ok=True)
out.save(save_path)
def score(samples, bench_root, out_dir):
psnrs, ssims = [], []
for key, item in samples.items():
if any(g in item["ref_image_path"] for g in SKIP_GROUPS):
continue
ref_img = read_image(os.path.join(bench_root, item["ref_image_path"]))
new_path = os.path.join(out_dir, item["ref_image_path"])
if not os.path.exists(new_path):
continue
new_img = read_image(new_path)
mask = read_mask(os.path.join(bench_root, item["mask_image_path"]))
if mask.shape != ref_img.shape[:2]:
mask = np.array(
Image.fromarray((mask * 255).astype(np.uint8)).resize(ref_img.shape[:2][::-1], Image.NEAREST)
) / 255.0
if new_img.shape != ref_img.shape:
new_img = np.array(
Image.fromarray((new_img * 255).astype(np.uint8)).resize(ref_img.shape[1::-1], Image.BILINEAR)
) / 255.0
psnrs.append(compute_psnr(new_img, ref_img, mask))
ssims.append(compute_ssim(new_img, ref_img, mask))
psnrs = [p for p in psnrs if not np.isnan(p)]
return float(np.mean(psnrs)), float(np.mean(ssims))
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--n-samples", type=int, default=40)
ap.add_argument("--steps", type=int, default=40)
ap.add_argument("--true-cfg-scale", type=float, default=4.0)
ap.add_argument("--guidance-scale", type=float, default=1.0)
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
from diffusers import QwenImageEditPlusPipeline
from safetensors.torch import load_file, save_file
bench_root = snapshot_download(
"wyh6666/GEditBench_ImgEditBench_with_mask", repo_type="dataset", allow_patterns=["geditbench/*"]
)
meta_path = os.path.join(bench_root, "geditbench", "metafile_geditbench.json")
with open(meta_path) as f:
full = json.load(f)
full = {k: v for k, v in full.items() if "/en/" in k and not any(g in k for g in SKIP_GROUPS)}
random.seed(args.seed)
keys = random.sample(list(full.keys()), min(args.n_samples, len(full)))
subset = {k: full[k] for k in keys}
print(f"Sampled {len(subset)} / {len(full)} eligible GEdit-Bench-EN entries")
pipeline = QwenImageEditPlusPipeline.from_pretrained("Qwen/Qwen-Image-Edit-2509")
pipeline.to(torch.bfloat16).to("cuda")
pipeline.set_progress_bar_config(disable=None)
baseline_out = "results/baseline"
cocoedit_out = "results/cocoedit"
print("Running baseline (no LoRA)...")
run_inference(pipeline, subset, bench_root, baseline_out, args.steps, args.true_cfg_scale, args.guidance_scale)
print("Loading CoCoEdit LoRA...")
lora_dir = snapshot_download("wyh6666/CoCoEdit")
converted = os.path.join(lora_dir, "adapter_model_converted.safetensors")
if not os.path.exists(converted):
sd = load_file(os.path.join(lora_dir, "adapter_model.safetensors"))
sd = {k.replace("base_model.model", "transformer"): v for k, v in sd.items()}
save_file(sd, converted)
pipeline.load_lora_weights(lora_dir, weight_name="adapter_model_converted.safetensors", adapter_name="lora")
pipeline.set_adapters(["lora"], adapter_weights=[1])
print("Running +CoCoEdit...")
run_inference(pipeline, subset, bench_root, cocoedit_out, args.steps, args.true_cfg_scale, args.guidance_scale)
base_psnr, base_ssim = score(subset, bench_root, baseline_out)
coco_psnr, coco_ssim = score(subset, bench_root, cocoedit_out)
print(f"\n=== Results on {len(subset)}-sample subset of GEdit-Bench-EN ===")
print(f"Qwen-Image-Edit PSNR={base_psnr:.3f} SSIM={base_ssim:.4f}")
print(f"Qwen-Image-Edit+CoCoEdit PSNR={coco_psnr:.3f} SSIM={coco_ssim:.4f}")
print(f"Delta PSNR={coco_psnr - base_psnr:+.3f} dB SSIM={coco_ssim - base_ssim:+.4f}")
print("Paper (full 606-sample GEdit-Bench-EN): PSNR +2.8 dB, SSIM +0.112 (Table 1)")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
6.92 kB
·
Xet hash:
b232110efe40979f70c98c5384a4ca954208816caa6b5fa66772f4dd4bce1ebd

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.