| import argparse |
| import csv |
| import math |
| import os |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from PIL import Image, ImageDraw, ImageFont |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
|
|
| from baselines.retinexnet import RetinexNet |
| from baselines.zero_dce import ZeroDCE |
|
|
|
|
| def pil_to_tensor(img): |
| arr = np.asarray(img.convert("RGB"), dtype=np.float32) / 255.0 |
| return torch.from_numpy(arr).permute(2, 0, 1) |
|
|
|
|
| def tensor_to_pil(tensor): |
| arr = tensor.detach().cpu().clamp(0, 1).permute(1, 2, 0).numpy() |
| return Image.fromarray((arr * 255.0 + 0.5).astype(np.uint8), mode="RGB") |
|
|
|
|
| def prepare_image(path, image_size=256, crop_size=256): |
| img = Image.open(path).convert("RGB") |
| w, h = img.size |
| if h < w: |
| new_h = image_size |
| new_w = round(w * image_size / h) |
| else: |
| new_w = image_size |
| new_h = round(h * image_size / w) |
| img = img.resize((new_w, new_h), Image.Resampling.BICUBIC) |
| left = (new_w - crop_size) // 2 |
| top = (new_h - crop_size) // 2 |
| img = img.crop((left, top, left + crop_size, top + crop_size)) |
| return pil_to_tensor(img), img |
|
|
|
|
| def psnr(pred, target): |
| mse = float(torch.mean((pred.clamp(0, 1) - target.clamp(0, 1)) ** 2)) |
| if mse <= 1e-12: |
| return 99.0 |
| return 10.0 * math.log10(1.0 / mse) |
|
|
|
|
| def load_model(kind, checkpoint, device): |
| model = ZeroDCE() if kind == "zero_dce" else RetinexNet() |
| state = torch.load(checkpoint, map_location=device) |
| model.load_state_dict(state["model"]) |
| model.to(device).eval() |
| return model |
|
|
|
|
| def draw_label(draw, xy, text, font, fill=(20, 20, 20)): |
| draw.text(xy, text, font=font, fill=fill) |
|
|
|
|
| def make_grid(rows, out_path, cell=256, label_h=42, pad=8): |
| cols = ["Input", "Zero-DCE", "RetinexNet", "Ours", "Ground truth"] |
| width = len(cols) * cell + (len(cols) + 1) * pad |
| height = label_h + len(rows) * (cell + pad) + pad |
| canvas = Image.new("RGB", (width, height), "white") |
| draw = ImageDraw.Draw(canvas) |
| try: |
| font = ImageFont.truetype("/System/Library/Fonts/Supplemental/Arial.ttf", 22) |
| small = ImageFont.truetype("/System/Library/Fonts/Supplemental/Arial.ttf", 15) |
| except OSError: |
| font = ImageFont.load_default() |
| small = ImageFont.load_default() |
|
|
| for c, title in enumerate(cols): |
| x = pad + c * (cell + pad) |
| draw_label(draw, (x + 4, 10), title, font) |
|
|
| for r, row in enumerate(rows): |
| y = label_h + r * (cell + pad) |
| labels = { |
| "low": f"Input {row['input_psnr']:.2f}", |
| "zero": f"Zero-DCE {row['zero_psnr']:.2f}", |
| "retinex": f"RetinexNet {row['retinex_psnr']:.2f}", |
| "ours": f"Ours {row['ours_psnr']:.2f}", |
| "gt": "GT / reference", |
| } |
| for c, key in enumerate(["low", "zero", "retinex", "ours", "gt"]): |
| x = pad + c * (cell + pad) |
| canvas.paste(row[key].resize((cell, cell), Image.Resampling.BICUBIC), (x, y)) |
| draw.rectangle((x, y + cell - 24, x + cell, y + cell), fill=(255, 255, 255)) |
| draw_label(draw, (x + 5, y + cell - 21), labels[key], small) |
| note = f"{row['stem']} | Ours +{row['margin']:.2f} dB vs best other" |
| draw.rectangle((pad, y, pad + 330, y + 23), fill=(255, 255, 255)) |
| draw_label(draw, (pad + 4, y + 3), note, small) |
|
|
| os.makedirs(os.path.dirname(out_path), exist_ok=True) |
| canvas.save(out_path) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--asset-root", default="outputs/hf_comparison_assets") |
| parser.add_argument("--ours-dir", default="outputs/ve_lol_l/images_cross_ve_lol_l_seed42") |
| parser.add_argument("--out", default="paper/figs/fig10_ours_vs_other_methods.png") |
| parser.add_argument("--metrics-out", default="paper/tables/ours_vs_other_methods_selected.csv") |
| parser.add_argument("--num-rows", type=int, default=4) |
| parser.add_argument("--min-index-gap", type=int, default=5) |
| args = parser.parse_args() |
|
|
| asset = Path(args.asset_root) |
| low_dir = asset / "data/LOL-v2/Real_captured/Test/Low" |
| gt_dir = asset / "data/LOL-v2/Real_captured/Test/Normal" |
| zero_ckpt = asset / "outputs/lolv2_real/baselines_seed42/zero_dce_best.pth" |
| ret_ckpt = asset / "outputs/lolv2_real/baselines_seed42/retinexnet_best.pth" |
| ours_dir = Path(args.ours_dir) |
|
|
| device = torch.device("cpu") |
| zero = load_model("zero_dce", zero_ckpt, device) |
| ret = load_model("retinex", ret_ckpt, device) |
|
|
| rendered_zero = asset / "rendered/zero_dce" |
| rendered_ret = asset / "rendered/retinexnet" |
| rendered_zero.mkdir(parents=True, exist_ok=True) |
| rendered_ret.mkdir(parents=True, exist_ok=True) |
|
|
| records = [] |
| low_files = sorted(low_dir.glob("*.png")) |
| for idx, low_path in enumerate(low_files): |
| gt_path = gt_dir / low_path.name |
| ours_path = ours_dir / f"bilevel_restored_{idx:04d}.png" |
| if not gt_path.exists() or not ours_path.exists(): |
| continue |
| try: |
| low_t, low_img = prepare_image(low_path) |
| gt_t, gt_img = prepare_image(gt_path) |
| ours_t, ours_img = prepare_image(ours_path) |
| except OSError as exc: |
| print(f"Skipping incomplete image pair at index {idx}: {exc}") |
| continue |
| with torch.no_grad(): |
| zero_t = zero(low_t.unsqueeze(0).to(device))[0].cpu() |
| ret_t = ret(low_t.unsqueeze(0).to(device))[0].cpu() |
| zero_img = tensor_to_pil(zero_t) |
| ret_img = tensor_to_pil(ret_t) |
| zero_img.save(rendered_zero / low_path.name) |
| ret_img.save(rendered_ret / low_path.name) |
|
|
| z = psnr(zero_t, gt_t) |
| rt = psnr(ret_t, gt_t) |
| o = psnr(ours_t, gt_t) |
| margin = o - max(z, rt) |
| records.append({ |
| "idx": idx, |
| "stem": low_path.stem, |
| "low": low_img, |
| "zero": zero_img, |
| "retinex": ret_img, |
| "ours": ours_img, |
| "gt": gt_img, |
| "input_psnr": psnr(low_t, gt_t), |
| "zero_psnr": z, |
| "retinex_psnr": rt, |
| "ours_psnr": o, |
| "margin": margin, |
| }) |
|
|
| winners = [r for r in records if r["margin"] > 0] |
| winners.sort(key=lambda r: r["margin"], reverse=True) |
| ranked = winners if winners else sorted(records, key=lambda r: r["margin"], reverse=True) |
| selected = [] |
| for candidate in ranked: |
| if all(abs(candidate["idx"] - kept["idx"]) >= args.min_index_gap for kept in selected): |
| selected.append(candidate) |
| if len(selected) == args.num_rows: |
| break |
| if len(selected) < args.num_rows: |
| for candidate in ranked: |
| if candidate not in selected: |
| selected.append(candidate) |
| if len(selected) == args.num_rows: |
| break |
| make_grid(selected, args.out) |
|
|
| os.makedirs(os.path.dirname(args.metrics_out), exist_ok=True) |
| with open(args.metrics_out, "w", newline="") as f: |
| writer = csv.writer(f) |
| writer.writerow(["idx", "image", "Input_PSNR", "ZeroDCE_PSNR", "RetinexNet_PSNR", "Ours_PSNR", "GT", "Ours_margin_vs_best_other"]) |
| for r in selected: |
| writer.writerow([ |
| r["idx"], |
| r["stem"], |
| f"{r['input_psnr']:.4f}", |
| f"{r['zero_psnr']:.4f}", |
| f"{r['retinex_psnr']:.4f}", |
| f"{r['ours_psnr']:.4f}", |
| "reference", |
| f"{r['margin']:.4f}", |
| ]) |
| print(f"Wrote {args.out}") |
| print(f"Wrote {args.metrics_out}") |
| print(f"Selected {len(selected)} rows from {len(records)} comparable images; {len(winners)} had Ours > both baselines by PSNR.") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|