#!/usr/bin/env python """Val 评测 + 代理分: 在隔离 val(真实对/官方对) 上计算 8 指标 + proxy + 保真护栏。 用法: python src/eval_val.py --weights weight/s2/net_params_X.pkl --pairs_json data/manifest_val.json python src/eval_val.py --weights ... --official_lq --official_gt """ import argparse, json, os, sys from pathlib import Path REPO = Path(__file__).resolve().parents[1] sys.path.insert(0, str(REPO)); sys.path.insert(0, str(REPO / "src")); sys.path.insert(0, str(REPO / "official")) import numpy as np, torch from PIL import Image from inference_4k import build_model, infer_4k_single def load_metrics(device): import pyiqa return {k: pyiqa.create_metric(k, device=device) for k in ["psnr", "ssim", "lpips", "dists", "niqe", "maniqa", "musiq", "clipiqa"]} def resize_pair(sr, hr, max_side=1024): sr = sr.convert("RGB"); hr = hr.convert("RGB") for im in (sr, hr): w, h = im.size if max(w, h) > max_side: s = max_side / max(w, h) im.thumbnail((int(w * s), int(h * s)), Image.LANCZOS) return sr, hr def metric_dict(metrics, sr, hr, device): import torchvision.transforms as T def t(im): return T.ToTensor()(im).to(device).unsqueeze(0) a, b = t(sr), t(hr) out = {"psnr": metrics["psnr"](a, b).item(), "ssim": metrics["ssim"](a, b).item(), "lpips": metrics["lpips"](a, b).item(), "dists": metrics["dists"](a, b).item(), "niqe": metrics["niqe"](a).item(), "maniqa": metrics["maniqa"](a).item(), "musiq": metrics["musiq"](a).item(), "clipiqa": metrics["clipiqa"](a).item()} return out def proxy(m): return (0.25 * m["clipiqa"] + 0.2 * (m["musiq"] / 100.0) + 0.15 * m["maniqa"] + 0.15 * (1 - m["niqe"] / 10.0) + 0.25 * (1 - m["lpips"])) def main(): ap = argparse.ArgumentParser() ap.add_argument("--weights", required=True) ap.add_argument("--half_decoder", default="weight/pretrained/halfDecoder.ckpt") ap.add_argument("--model_id", default="models/stable-diffusion-2-1-base") ap.add_argument("--pairs_json", default="", help='[{lr,hr}] 或 manifest 结构') ap.add_argument("--official_lq", default="", help="同分辨率官方对(整图滑窗)") ap.add_argument("--official_gt", default="") ap.add_argument("--out", default="logs/eval_result.json") ap.add_argument("--max_side", type=int, default=1024) args = ap.parse_args() device = "cuda" if torch.cuda.is_available() else "cpu" net, tail = build_model(args.weights, args.half_decoder, args.model_id, device, bf16=True) metrics = load_metrics(device) rows = [] if args.pairs_json: with open(args.pairs_json, encoding="utf-8") as fh: data = json.load(fh) pairs = data.get("real_pairs", data if isinstance(data, list) else []) for i, p in enumerate(pairs): with torch.no_grad(): lr = Image.open(p["lr"]).convert("RGB") hr = Image.open(p["hr"]).convert("RGB") lr128 = lr.resize((max(1, lr.width // 4), max(1, lr.height // 4)), Image.LANCZOS) t = torch.from_numpy(np.asarray(lr128, dtype=np.float32).transpose(2, 0, 1) / 255.0 * 2 - 1)[None].to(device) with torch.autocast("cuda", dtype=torch.bfloat16): z = net(t); sr_arr = tail(z) sr = Image.fromarray(((sr_arr[0].float().cpu().numpy().transpose(1, 2, 0) + 1) / 2 * 255).clip(0, 255).astype(np.uint8)) sr, hr = resize_pair(sr, hr, args.max_side) m = metric_dict(metrics, sr, hr, device) m["name"] = os.path.basename(p["lr"]); m["proxy"] = proxy(m) rows.append(m) print(f" [{i+1}] {m['name']} proxy {m['proxy']:.4f}", flush=True) if args.official_lq and args.official_gt: lq_files = sorted(os.listdir(args.official_lq)) for f in lq_files: lq_p = os.path.join(args.official_lq, f) gt_p = os.path.join(args.official_gt, f.replace("_lq", "_gt").replace("lq.jpg", "gt.png")) if not os.path.exists(gt_p): gt_p = os.path.join(args.official_gt, f.replace("_lq.jpg", "_gt.jpg")) if not os.path.exists(gt_p): continue sr = infer_4k_single(lq_p, net, tail, device) gt = Image.open(gt_p).convert("RGB") sr, gt = resize_pair(sr, gt, args.max_side) m = metric_dict(metrics, sr, gt, device) m["name"] = f; m["proxy"] = proxy(m) rows.append(m) print(f" [official] {f} proxy {m['proxy']:.4f}", flush=True) if not rows: print("无评测样本"); return keys = ["psnr", "ssim", "lpips", "dists", "niqe", "maniqa", "musiq", "clipiqa", "proxy"] agg = {k: float(np.mean([r[k] for r in rows])) for k in keys} result = {"per_image": rows, "mean": agg} os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) with open(args.out, "w", encoding="utf-8") as fh: json.dump(result, fh, ensure_ascii=False, indent=1) print(json.dumps(agg, indent=1)) if __name__ == "__main__": main()