CSIGv3_train_script / src /eval_val.py
XenderYang's picture
CSIGv3 AdcSR train scripts + A100 runbook
4811c23 verified
Raw
History Blame Contribute Delete
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()