"""
upscale_thumbnails.py - 4K AI Thumbnail Upscaler for Historical Series
Uses NVIDIA GeForce GTX 1660 Super (CUDA FP32) via PyTorch + Pillow
Pipeline:
  1. 1x_PureVision (Denoise & JPEG artifact restoration)
  2. 4xPurePhoto-span (Swift Parameter-free Attention 4x Super-Resolution)
  3. Fused Bicubic Resampling to 4K UHD (3840x2160)
"""

import os
import sys
import time
import glob
import torch
import torch.nn.functional as F
import spandrel
import numpy as np
from PIL import Image

if sys.platform.startswith("win"):
    try:
        sys.stdout.reconfigure(encoding="utf-8")
        sys.stderr.reconfigure(encoding="utf-8")
    except Exception:
        pass

INPUT_DIR = os.path.abspath(r"content\thumbnails")
OUTPUT_DIR = os.path.abspath(r"content\thumbnails_4k")
PV_MODEL = os.path.abspath(r"models\1x_PureVision.pth")
SR_MODEL = os.path.abspath(r"models\4xPurePhoto-span.pth")

os.makedirs(OUTPUT_DIR, exist_ok=True)

def main():
    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    gpu_name = torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU"
    
    print("=" * 65)
    print(" 4K AI THUMBNAIL UPSCALER (PureVision + SPAN)")
    print(f" Device: {gpu_name}")
    print(f" Input Directory:  {INPUT_DIR}")
    print(f" Output Directory: {OUTPUT_DIR}")
    print("=" * 65)

    # 1. Load Models into GPU
    print("\n[1/3] Loading AI Models into GTX 1660 Super VRAM...")
    t0 = time.time()
    m_pv = spandrel.ModelLoader().load_from_file(PV_MODEL).to(device).eval()
    m_sr = spandrel.ModelLoader().load_from_file(SR_MODEL).to(device).eval()
    print(f"✓ Models loaded successfully in {time.time() - t0:.2f}s")

    # 2. Find all thumbnails
    thumb_files = sorted(glob.glob(os.path.join(INPUT_DIR, "ep*_thumb.jpg")))
    if not thumb_files:
        print("[-] Không tìm thấy file thumbnail nào trong", INPUT_DIR)
        return

    print(f"\n[2/3] Tìm thấy {len(thumb_files)} thumbnail cần upscale lên 4K:")
    for f in thumb_files:
        print(f"  - {os.path.basename(f)}")

    print("\n[3/3] Tiến hành xử lý AI 4K (Denoise -> 4x SR -> 4K UHD Resample)...")
    results = []

    with torch.no_grad():
        for idx, in_path in enumerate(thumb_files, 1):
            fname = os.path.basename(in_path)
            t_start = time.time()
            
            # Read image via Pillow
            pil_img = Image.open(in_path).convert("RGB")
            orig_w, orig_h = pil_img.size
            
            # Convert to PyTorch float32 tensor (1, 3, H, W)
            np_arr = np.array(pil_img, dtype=np.float32) / 255.0
            t_rgb = torch.from_numpy(np_arr).permute(2, 0, 1).unsqueeze(0).to(device, non_blocking=True)

            # Step 1: Denoise with PureVision
            torch.cuda.synchronize()
            t_pv_start = time.time()
            clean_rgb = m_pv(t_rgb)
            torch.cuda.synchronize()
            pv_time = time.time() - t_pv_start

            # Step 2: 4x Super-Resolution with SPAN
            t_sr_start = time.time()
            sr_rgb = m_sr(clean_rgb)
            torch.cuda.synchronize()
            sr_time = time.time() - t_sr_start

            # Step 3: Resample to exact 4K UHD (3840x2160)
            target_w, target_h = 3840, 2160
            out_4k = F.interpolate(sr_rgb, size=(target_h, target_w), mode="bicubic", align_corners=False)
            
            # Convert back to uint8 RGB
            out_arr = (out_4k.squeeze(0).permute(1, 2, 0).clamp(0.0, 1.0).cpu().numpy() * 255.0).astype(np.uint8)
            out_pil = Image.fromarray(out_arr)

            # Save JPEG 4K (Quality 98) & PNG 4K
            base_name = os.path.splitext(fname)[0]
            out_jpg = os.path.join(OUTPUT_DIR, f"{base_name}_4k.jpg")
            out_png = os.path.join(OUTPUT_DIR, f"{base_name}_4k.png")
            
            out_pil.save(out_jpg, "JPEG", quality=98, subsampling=0)
            out_pil.save(out_png, "PNG")
            
            total_time = time.time() - t_start
            jpg_mb = os.path.getsize(out_jpg) / (1024 * 1024)
            png_mb = os.path.getsize(out_png) / (1024 * 1024)
            peak_vram = torch.cuda.max_memory_allocated() / (1024 * 1024)

            print(f"[{idx}/{len(thumb_files)}] ✓ {fname} ({orig_w}x{orig_h}) -> 4K ({target_w}x{target_h}) | Time: {total_time:.2f}s (PV: {pv_time*1000:.0f}ms, SPAN: {sr_time*1000:.0f}ms) | Size: {jpg_mb:.1f}MB JPG, {png_mb:.1f}MB PNG | VRAM: {peak_vram:.0f}MB")
            
            results.append({
                "file": fname,
                "input_res": f"{orig_w}x{orig_h}",
                "output_res": f"{target_w}x{target_h}",
                "time_sec": round(total_time, 2),
                "jpg_mb": round(jpg_mb, 2),
                "png_mb": round(png_mb, 2),
                "jpg_path": out_jpg,
                "png_path": out_png
            })

    print("\n" + "=" * 65)
    print(f"🎉 HOÀN TẤT UPSCALE 4K TOÀN BỘ {len(results)} THUMBNAIL TRÊN GTX 1660 SUPER!")
    print(f"Thư mục lưu trữ: {OUTPUT_DIR}")
    print("=" * 65)

if __name__ == "__main__":
    main()
