File size: 5,192 Bytes
4811c23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
#!/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()