| |
| """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 <val lq dir> --official_gt <val gt dir> |
| """ |
| 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()
|
|
|