#!/usr/bin/env python3
"""Score every upscaled image against its reference: PSNR, SSIM and LPIPS.

Metric definitions match PhotoAIBench's free comparison tool (engine/tools/image-quality-compare),
so anyone can check a number by loading the two images there:

  PSNR (RGB)  10 log10(255^2 / MSE), MSE over R, G and B of every pixel. Higher is better.
  PSNR (Y)    the same on BT.601 full-range luma Y = 0.299 R + 0.587 G + 0.114 B.
  SSIM (Y)    Wang et al. 2004 on Y: Gaussian window sigma 1.5 (11 x 11), K1 0.01, K2 0.03,
              L 255, population covariance, valid windows only (scikit-image
              structural_similarity with gaussian_weights=True, use_sample_covariance=False).
              Higher is better; 1 means identical.
  LPIPS       Zhang et al. 2018, `lpips` package 0.1.4, AlexNet backbone, version 0.1, on RGB
              scaled to [-1, 1], averaged over the image. Lower is better; 0 means identical.

No border pixels are shaved: what a tool does at the edges is part of its output.

Inputs: data/private/photo-bench/<id>/{reference,<method>}.png. Output:
data/photo/bench/results-baselines.json (per image, per method, plus per-method summaries).

    data/private/photo-bench-venv/bin/python data/photo/bench/score.py
"""
from __future__ import annotations

import json
import sys
from datetime import date
from pathlib import Path

import numpy as np
from PIL import Image
from skimage.metrics import structural_similarity

REPO = Path(__file__).resolve().parents[3]
sys.path.insert(0, str(REPO / "data" / "photo"))
from fetch_samples import load_samples  # noqa: E402

BENCH = REPO / "data" / "private" / "photo-bench"
OUT = REPO / "data" / "photo" / "bench" / "results-baselines.json"
METHODS = ["bicubic", "lanczos", "realesrnet-x4plus", "realesrgan-x4plus", "realesr-general-x4v3"]


def luma(rgb: np.ndarray) -> np.ndarray:
    return 0.299 * rgb[..., 0] + 0.587 * rgb[..., 1] + 0.114 * rgb[..., 2]


def psnr(mse: float) -> float:
    return float("inf") if mse == 0 else float(10 * np.log10(255.0 ** 2 / mse))


class Lpips:
    def __init__(self):
        import lpips
        import torch

        self.torch = torch
        self.net = lpips.LPIPS(net="alex", version="0.1", verbose=False).eval()

    def __call__(self, a: np.ndarray, b: np.ndarray) -> float:
        t = lambda x: self.torch.from_numpy(x.astype(np.float32) / 127.5 - 1.0).permute(2, 0, 1).unsqueeze(0)
        with self.torch.no_grad():
            return float(self.net(t(a), t(b)).item())


def score_pair(ref: np.ndarray, out: np.ndarray, lp: Lpips) -> dict:
    ref_f, out_f = ref.astype(np.float64), out.astype(np.float64)
    yr, yo = luma(ref_f), luma(out_f)
    return {
        "psnr_rgb": round(psnr(float(np.mean((ref_f - out_f) ** 2))), 4),
        "psnr_y": round(psnr(float(np.mean((yr - yo) ** 2))), 4),
        "ssim_y": round(float(structural_similarity(yr, yo, data_range=255, gaussian_weights=True, sigma=1.5,
                                                    use_sample_covariance=False)), 5),
        "lpips": round(lp(ref, out), 5),
    }


def main(argv=None) -> int:
    import argparse

    import lpips
    import skimage
    import torch

    ap = argparse.ArgumentParser()
    ap.add_argument("--vendors", action="store_true",
                    help="score commercial tools' outputs (vendor-<id>.*) into results-vendors.json instead")
    args = ap.parse_args(argv)
    global OUT, METHODS
    if args.vendors:
        OUT = OUT.with_name("results-vendors.json")
        METHODS = sorted({p.stem for p in BENCH.glob("*/vendor-*.*")})

    samples = {s["id"]: s for s in load_samples() if s["use"] == "upscale"}
    lp = Lpips()
    per_image = {}
    for sid, s in samples.items():
        d = BENCH / sid
        ref = np.asarray(Image.open(d / "reference.png").convert("RGB"))
        per_image[sid] = {"category": s["category"], "size": [ref.shape[1], ref.shape[0]], "scores": {}}
        for m in METHODS:
            f = next((p for p in d.iterdir() if p.stem == m), None)
            if f is None:
                continue
            out = np.asarray(Image.open(f).convert("RGB"))
            assert out.shape == ref.shape, (sid, m, out.shape, ref.shape)
            per_image[sid]["scores"][m] = score_pair(ref, out, lp)
            print(sid, m, per_image[sid]["scores"][m], flush=True)

    summary = {}
    for m in METHODS:
        rows = [v["scores"][m] for v in per_image.values() if m in v["scores"]]
        if not rows:
            continue
        summary[m] = {
            "images": len(rows),
            **{f"mean_{k}": round(float(np.mean([r[k] for r in rows])), 4) for k in ("psnr_rgb", "psnr_y", "ssim_y", "lpips")},
        }
    # How often each method is best on each metric (ties count for all tied methods).
    wins = {m: {"psnr_y": 0, "ssim_y": 0, "lpips": 0} for m in summary}
    for v in per_image.values():
        sc = v["scores"]
        if not sc:
            continue
        for k, better in (("psnr_y", max), ("ssim_y", max), ("lpips", min)):
            best = better(sc[m][k] for m in sc)
            for m in sc:
                if sc[m][k] == best:
                    wins[m][k] += 1
    for m in summary:
        summary[m]["best_on"] = wins[m]

    OUT.write_text(json.dumps({
        "name": ("PhotoAIBench upscaler test: commercial web tools (DRAFT until vendors have replied)" if args.vendors
                 else "PhotoAIBench upscaler baselines (local open-source methods only)"),
        "run_date": date.today().isoformat(),
        "scale": 4,
        "images": len(per_image),
        "metrics": {
            "psnr_rgb": "dB, higher is better; MSE over R, G, B",
            "psnr_y": "dB, higher is better; BT.601 full-range luma",
            "ssim_y": "Wang et al. 2004, Gaussian sigma 1.5 (11x11), K1 0.01, K2 0.03, L 255, on luma; higher is better",
            "lpips": "Zhang et al. 2018, AlexNet, v0.1; lower is better",
        },
        "versions": {"scikit-image": skimage.__version__, "lpips": getattr(lpips, "__version__", "0.1.4"), "torch": torch.__version__},
        "methods": METHODS,
        "summary": summary,
        "per_image": per_image,
    }, indent=1) + "\n", encoding="utf-8")
    print(json.dumps(summary, indent=1))
    return 0


if __name__ == "__main__":
    sys.exit(main())
