diff / scripts /make_other_methods_grid.py
siddharthdhara17's picture
Add comparison figures and filled evaluation tables
5777c7c verified
Raw
History Blame Contribute Delete
7.79 kB
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()