#!/usr/bin/env python3
"""Upscale every low-resolution test input 4x with the local, open-source baselines.

Methods (licenses and weight hashes are recorded in data/photo/tools.yaml):

  bicubic               Pillow BICUBIC, 4x. The floor: what any image editor does.
  lanczos               Pillow LANCZOS, 4x. A sharper classic filter.
  realesrnet-x4plus     Real-ESRGAN project, RealESRNet_x4plus.pth: RRDBNet trained with pixel loss
                        only (no GAN), on synthetic real-world degradations.
  realesrgan-x4plus     Real-ESRGAN project, RealESRGAN_x4plus.pth: the same network fine-tuned with
                        perceptual and GAN losses. The model most "AI upscalers" started from.
  realesr-general-x4v3  Real-ESRGAN project, realesr-general-x4v3.pth: a small compact network
                        (SRVGGNetCompact) for general scenes.

The networks run with PyTorch on the CPU, whole image at once (no tiles), float32, input scaled to
[0, 1]; output clipped to [0, 1], multiplied by 255 and rounded, as the project's own inference
script does. Models are loaded with spandrel, which reads the original .pth files.

Inputs: data/private/photo-bench/<id>/lr.png (from prepare.py).
Outputs: data/private/photo-bench/<id>/<method>.png, plus timings in
data/photo/bench/upscale-timings.json.

    data/private/photo-bench-venv/bin/python data/photo/bench/upscale_local.py [--only id,...] [--methods m,...]
"""
from __future__ import annotations

import argparse
import json
import platform
import sys
import time
from pathlib import Path

import numpy as np
from PIL import Image

REPO = Path(__file__).resolve().parents[3]
BENCH = REPO / "data" / "private" / "photo-bench"
MODELS = REPO / "data" / "private" / "photo-models"
TIMINGS = REPO / "data" / "photo" / "bench" / "upscale-timings.json"
SCALE = 4

NETWORKS = {
    "realesrnet-x4plus": "RealESRNet_x4plus.pth",
    "realesrgan-x4plus": "RealESRGAN_x4plus.pth",
    "realesr-general-x4v3": "realesr-general-x4v3.pth",
}
METHODS = ["bicubic", "lanczos", *NETWORKS]


def classic(lr: Image.Image, method: str) -> Image.Image:
    f = {"bicubic": Image.BICUBIC, "lanczos": Image.LANCZOS}[method]
    return lr.resize((lr.width * SCALE, lr.height * SCALE), f)


_loaded: dict = {}


def network(lr: Image.Image, method: str) -> Image.Image:
    import torch
    from spandrel import ModelLoader

    if method not in _loaded:
        desc = ModelLoader().load_from_file(str(MODELS / NETWORKS[method]))
        assert desc.scale == SCALE, f"{method}: scale {desc.scale}"
        desc.model.eval()
        _loaded[method] = desc
    model = _loaded[method].model
    x = torch.from_numpy(np.asarray(lr, dtype=np.float32) / 255.0).permute(2, 0, 1).unsqueeze(0)
    with torch.no_grad():
        y = model(x)
    arr = (y.squeeze(0).clamp(0, 1).permute(1, 2, 0).numpy() * 255.0).round().astype(np.uint8)
    return Image.fromarray(arr, "RGB")


def main(argv=None) -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--only", help="comma-separated sample ids")
    ap.add_argument("--methods", help=f"comma-separated subset of {','.join(METHODS)}")
    ap.add_argument("--force", action="store_true", help="redo outputs that already exist")
    args = ap.parse_args(argv)
    import torch

    torch.set_num_threads(8)
    ids = sorted(p.name for p in BENCH.iterdir() if (p / "lr.png").exists())
    if args.only:
        ids = [i for i in ids if i in set(args.only.split(","))]
    methods = args.methods.split(",") if args.methods else METHODS
    timings = json.loads(TIMINGS.read_text()) if TIMINGS.exists() else {"runs": {}}
    timings["machine"] = f"{platform.machine()} {platform.processor() or ''} {platform.system()} {platform.release()}".strip()
    timings["torch"] = torch.__version__
    timings["threads"] = torch.get_num_threads()
    for sid in ids:
        lr = Image.open(BENCH / sid / "lr.png").convert("RGB")
        for m in methods:
            out = BENCH / sid / f"{m}.png"
            if out.exists() and not args.force:
                continue
            t0 = time.perf_counter()
            img = classic(lr, m) if m in ("bicubic", "lanczos") else network(lr, m)
            dt = time.perf_counter() - t0
            assert img.size == (lr.width * SCALE, lr.height * SCALE)
            img.save(out)
            timings["runs"].setdefault(sid, {})[m] = round(dt, 2)
            print(f"{sid:30s} {m:22s} {dt:7.1f} s", flush=True)
            TIMINGS.write_text(json.dumps(timings, indent=1, sort_keys=True) + "\n")
    return 0


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