import os
import sys
import time
import torch
import spandrel
import numpy as np
from PIL import Image

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

torch.set_num_threads(56)

print("=== 1. Nap Model 1x_PureVision & 4x_SAFMN_PureScale ===")
loader = spandrel.ModelLoader()
t_load = time.time()
pv_model = loader.load_from_file(r"C:\Projects\Historical\models\1x_PureVision.pth").model.eval()
ps_model = loader.load_from_file(r"C:\Projects\Historical\models\4x_SAFMN_PureScale.pth").model.eval()
print(f"Nap xong 2 model trong {time.time()-t_load:.2f}s")

input_path = r"C:\Projects\Historical\downloads\nomos_test_in.png"
img = Image.open(input_path).convert("RGB")
print(f"Anh dau vao: {input_path} (Kich thuoc: {img.size})")

arr = np.array(img).astype(np.float32) / 255.0
tensor = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0)

# Buoc 1: Tien xu ly khu nhieu voi 1x_PureVision
print("\n[Buoc 1/3] Dang chay 1x_PureVision (Khu muoi nen H.264 & tai tao texture)...")
t1 = time.time()
with torch.no_grad():
    clean_tensor = pv_model(tensor)
t1_cost = time.time() - t1
print(f"Hoan tat Buoc 1 trong {t1_cost:.2f}s (Shape: {clean_tensor.shape})")

# Buoc 2: Phong dai 4x voi 4x_SAFMN_PureScale
print("\n[Buoc 2/3] Dang chay 4x_SAFMN_PureScale (Phong dai 4x len 5K sieu nhe)...")
t2 = time.time()
with torch.no_grad():
    out_5k = ps_model(clean_tensor)
t2_cost = time.time() - t2
print(f"Hoan tat Buoc 2 trong {t2_cost:.2f}s (Shape: {out_5k.shape})")

# Luu anh 5K
out_arr_5k = (out_5k.squeeze(0).permute(1, 2, 0).clamp(0, 1).numpy() * 255.0).astype(np.uint8)
img_5k = Image.fromarray(out_arr_5k)
out_5k_path = r"C:\Projects\Historical\downloads\purescale_5k_test.png"
img_5k.save(out_5k_path)
print(f"Da luu anh 5K: {out_5k_path} (Size: {img_5k.size}, Dung luong: {os.path.getsize(out_5k_path)/(1024*1024):.2f} MB)")

# Buoc 3: Ep ve chuan 4K bang Lanczos
print("\n[Buoc 3/3] Dang ep ve chuan 4K (3840x2160) bang thuat toan Lanczos...")
t3 = time.time()
img_4k = img_5k.resize((3840, 2160), Image.Resampling.LANCZOS)
out_4k_path = r"C:\Projects\Historical\downloads\purescale_4k_test.png"
img_4k.save(out_4k_path)
t3_cost = time.time() - t3
print(f"Da luu anh 4K: {out_4k_path} (Size: {img_4k.size}, Dung luong: {os.path.getsize(out_4k_path)/(1024*1024):.2f} MB)")

total_time = t1_cost + t2_cost + t3_cost
print(f"\n======================================================")
print(f"TONG KET PIPELINE PUREVISION + PURESCALE:")
print(f"- Buoc 1 (1x_PureVision):      {t1_cost:.2f}s")
print(f"- Buoc 2 (4x_SAFMN_PureScale): {t2_cost:.2f}s")
print(f"- Buoc 3 (Lanczos downscale):  {t3_cost:.2f}s")
print(f"TONG THOI GIAN 1 FRAME:        {total_time:.2f}s")
print(f"======================================================")
