Download src/eval_val.py from XenderYang/CSIGv3_train_script: direct link, hf CLI and curl.
- Browser
- Download file 5.19 kB
-
https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/eval_val.py
- Command line
-
hf download hf://XenderYang/CSIGv3_train_script/src/eval_val.py
-
curl -L -o eval_val.py https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/eval_val.py
5.19 kB
| #!/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 <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() | |