"""
purescale_video_upscaler.py - Next-Gen 3-Step Video Upscaler for Historical Series
Pipeline:
  Step 1: 1x_PureVision (Denoise & H.264 artifact restoration in FP32)
  Step 2: 4x Super-Resolution (Default: 4xPurePhoto-span.pth - Swift Parameter-free Attention)
  Step 3: Downscaling 5K -> 4K UHD (3840x2160)
Encoding:
  NVIDIA NVENC HEVC (CQ 19, Preset p4) with piped I/O (zero disk intermediate frames)
"""

import argparse
import os
import subprocess
import sys
import time
import json
import torch
import torch.nn.functional as F
import spandrel
import numpy as np

FFMPEG_PATH = r"C:\Youtube\ffmpeg.exe"
FFPROBE_PATH = r"C:\Youtube\ffprobe.exe"


def get_video_info(video_path: str):
    cmd = [
        FFPROBE_PATH,
        "-v", "error",
        "-select_streams", "v:0",
        "-show_entries", "stream=width,height,r_frame_rate,nb_frames,duration",
        "-of", "json",
        video_path
    ]
    res = subprocess.run(cmd, capture_output=True, text=True, check=True)
    data = json.loads(res.stdout)
    stream = data["streams"][0]
    
    width = int(stream["width"])
    height = int(stream["height"])
    r_fps = stream["r_frame_rate"].split("/")
    fps = float(r_fps[0]) / float(r_fps[1]) if len(r_fps) == 2 else float(r_fps[0])
    
    nb_frames = int(stream.get("nb_frames", 0))
    duration = float(stream.get("duration", 0.0))
    if nb_frames == 0 and duration > 0:
        nb_frames = int(duration * fps)
        
    return {"width": width, "height": height, "fps": fps, "nb_frames": nb_frames, "duration": duration}


def build_tile_tensors(h, w, tile_size=512, tile_pad=32, scale=4, device="cuda:0"):
    stride = tile_size - 2 * tile_pad
    y_starts = list(range(0, h, stride))
    x_starts = list(range(0, w, stride))
    
    tiles_meta = []
    for y0 in y_starts:
        for x0 in x_starts:
            y1 = min(y0 + tile_size, h)
            x1 = min(x0 + tile_size, w)
            y0_adj = max(0, y1 - tile_size)
            x0_adj = max(0, x1 - tile_size)
            
            th = y1 - y0_adj
            tw = x1 - x0_adj
            
            w_h = torch.ones(th * scale, device=device)
            w_w = torch.ones(tw * scale, device=device)
            pad_s = tile_pad * scale
            if y0_adj > 0:
                w_h[:pad_s] = torch.linspace(0, 1, pad_s, device=device)
            if y1 < h:
                w_h[-pad_s:] = torch.linspace(1, 0, pad_s, device=device)
            if x0_adj > 0:
                w_w[:pad_s] = torch.linspace(0, 1, pad_s, device=device)
            if x1 < w:
                w_w[-pad_s:] = torch.linspace(1, 0, pad_s, device=device)
                
            mask = (w_h[:, None] * w_w[None, :]).unsqueeze(0).unsqueeze(0)
            tiles_meta.append((y0_adj, y1, x0_adj, x1, mask))
            
    return tiles_meta


def tile_upscale(model, x, tiles_meta, scale=4):
    b, c, h, w = x.shape
    out_h, out_w = h * scale, w * scale
    output = torch.zeros((b, c, out_h, out_w), device=x.device, dtype=torch.float32)
    weights = torch.zeros((1, 1, out_h, out_w), device=x.device, dtype=torch.float32)

    for y0, y1, x0, x1, mask in tiles_meta:
        tile = x[:, :, y0:y1, x0:x1]
        tile_out = model(tile)
        output[:, :, y0*scale:y1*scale, x0*scale:x1*scale] += tile_out * mask
        weights[:, :, y0*scale:y1*scale, x0*scale:x1*scale] += mask

    output /= weights.clamp(min=1e-5)
    return output


def upscale_video(
    input_video: str,
    output_video: str,
    pv_model_path: str = "models/1x_PureVision.pth",
    sr_model_path: str = "models/4xPurePhoto-span.pth",
    start_sec: float = 0.0,
    duration: float = None,
    tile_size: int = 512,
    tile_pad: int = 32,
    force_tile: bool = False,
    cq: int = 19,
):
    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    print(f"==================================================")
    print(f" PureVision + SPAN 4K AI Video Upscaler")
    print(f" Device: {torch.cuda.get_device_name(0)}")
    print(f"==================================================")

    info = get_video_info(input_video)
    in_w, in_h = info["width"], info["height"]
    fps = info["fps"]
    total_video_frames = info["nb_frames"]
    
    if duration is not None and duration > 0:
        frames_to_process = int(duration * fps)
    else:
        frames_to_process = total_video_frames - int(start_sec * fps)

    target_4k_w = 3840
    target_4k_h = 2160

    print(f"Input: {input_video} ({in_w}x{in_h} @ {fps:.2f} fps)")
    print(f"Target: {output_video} ({target_4k_w}x{target_4k_h} 4K UHD)")
    print(f"Processing range: {start_sec:.2f}s to {start_sec + (duration or info['duration']):.2f}s ({frames_to_process} frames)")

    # 1. Load Models
    print(f"\n[1/3] Loading models...")
    print(f"  - Step 1 (Denoise): {pv_model_path}")
    print(f"  - Step 2 (Super-Res): {sr_model_path}")
    t_load = time.time()
    m_pv = spandrel.ModelLoader().load_from_file(pv_model_path).to(device).eval()
    m_sr = spandrel.ModelLoader().load_from_file(sr_model_path).to(device).eval()
    raw_sr = m_sr.model
    arch_name = getattr(m_sr.architecture, "name", str(m_sr.architecture))
    print(f"Models loaded in {time.time() - t_load:.2f}s (SR Architecture: {arch_name})")

    # SPAN and Compact run direct with ~3.6GB VRAM peak, SAFMN uses tiles
    use_tile = force_tile or ("SAFMN" in arch_name and not force_tile)
    if use_tile:
        print(f"Tiling: Enabled ({tile_size}x{tile_size}, pad: {tile_pad}px)")
        tiles_meta = build_tile_tensors(in_h, in_w, tile_size=tile_size, tile_pad=tile_pad, scale=4, device=device)
        print(f"Pre-computed {len(tiles_meta)} tiles per frame.")
    else:
        print(f"Tiling: Disabled (Direct Full-Frame Mode, ~3.6GB VRAM on GTX 1660 Super)")
        tiles_meta = None

    # 2. Setup FFmpeg Pipes
    if duration is not None and duration <= 0:
        duration = None
    print("\n[2/3] Initializing zero-disk FFmpeg video pipeline...")
    reader_cmd = [
        FFMPEG_PATH,
        "-ss", str(start_sec),
    ]
    if duration is not None:
        reader_cmd.extend(["-t", str(duration)])
    reader_cmd.extend([
        "-i", input_video,
        "-f", "rawvideo",
        "-pix_fmt", "rgb24",
        "-vcodec", "rawvideo",
        "-"
    ])

    writer_cmd = [
        FFMPEG_PATH, "-y",
        "-f", "rawvideo",
        "-vcodec", "rawvideo",
        "-s", f"{target_4k_w}x{target_4k_h}",
        "-pix_fmt", "rgb24",
        "-r", str(fps),
        "-i", "-",
        "-ss", str(start_sec),
    ]
    if duration is not None:
        writer_cmd.extend(["-t", str(duration)])
    writer_cmd.extend([
        "-i", input_video,
        "-map", "0:v:0",
        "-map", "1:a:0?",
        "-c:v", "hevc_nvenc",
        "-preset", "p4",
        "-cq", str(cq),
        "-bf", "0",
        "-avoid_negative_ts", "make_zero",
        "-pix_fmt", "yuv420p",
        "-c:a", "copy",
        "-shortest",
        output_video
    ])

    reader_proc = subprocess.Popen(reader_cmd, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, bufsize=10**7)
    writer_proc = subprocess.Popen(writer_cmd, stdin=subprocess.PIPE, stderr=subprocess.DEVNULL, bufsize=10**7)

    frame_bytes = in_w * in_h * 3
    processed_count = 0
    start_time = time.time()

    pv_time_total = 0.0
    sr_time_total = 0.0
    down_time_total = 0.0

    print("\n[3/3] Commencing 4K Upscale Rendering Loop...")
    try:
        while True:
            raw_frame = reader_proc.stdout.read(frame_bytes)
            if not raw_frame or len(raw_frame) < frame_bytes:
                break

            processed_count += 1
            t_frame_start = time.time()

            # Ingest to GPU tensor (use .copy() to ensure array is writable)
            frame_np = np.frombuffer(raw_frame, dtype=np.uint8).reshape((in_h, in_w, 3)).copy()
            t_in = torch.from_numpy(frame_np.transpose(2, 0, 1)).float().div(255.0).unsqueeze(0).to(device)

            # Step 1: Denoise with 1x_PureVision (FP32)
            t_pv0 = time.time()
            with torch.inference_mode():
                t_denoised = m_pv(t_in).clamp(0.0, 1.0)
            pv_time_total += (time.time() - t_pv0)

            # Step 2: 4x Super-Resolution to 5K with SPAN / SR model
            t_sr0 = time.time()
            with torch.inference_mode():
                if use_tile:
                    t_5k = tile_upscale(raw_sr, t_denoised, tiles_meta, scale=4).clamp(0.0, 1.0)
                else:
                    t_5k = raw_sr(t_denoised).clamp(0.0, 1.0)
            sr_time_total += (time.time() - t_sr0)

            # Step 3: Downscale 5K -> 4K UHD (Bicubic)
            t_down0 = time.time()
            with torch.inference_mode():
                t_4k = F.interpolate(t_5k, size=(target_4k_h, target_4k_w), mode="bicubic", align_corners=False).clamp(0.0, 1.0)
            down_time_total += (time.time() - t_down0)

            # Step 4: Write frame to NVENC encoder pipe
            out_bytes = (t_4k.squeeze(0).permute(1, 2, 0).mul(255.0).byte().cpu().numpy()).tobytes()
            writer_proc.stdin.write(out_bytes)

            # Stats & Progress
            elapsed = time.time() - start_time
            current_fps = processed_count / elapsed
            eta_sec = (frames_to_process - processed_count) / max(current_fps, 1e-4) if frames_to_process > processed_count else 0
            vram_mb = torch.cuda.max_memory_allocated() / (1024 ** 2)

            if eta_sec >= 3600:
                eta_str = f"{int(eta_sec // 3600)}h {int((eta_sec % 3600) // 60):02d}m"
            else:
                eta_str = f"{eta_sec / 60:.1f}m"

            pct = (processed_count / max(frames_to_process, 1)) * 100
            msg = (
                f"Frame [{processed_count}/{frames_to_process}] ({pct:.1f}%) "
                f"| Speed: {current_fps:.2f} FPS "
                f"| Elapsed: {int(elapsed // 60)}m {int(elapsed % 60):02d}s "
                f"| ETA: {eta_str} "
                f"| VRAM: {vram_mb:.0f} MB"
            )

            if processed_count % 50 == 0 or processed_count == 1:
                print(f"\r{msg}", flush=True)
            else:
                print(f"\r{msg}", end="", flush=True)

    except KeyboardInterrupt:
        print("\n\nUser interrupted process.")
    finally:
        reader_proc.stdout.close()
        reader_proc.terminate()
        writer_proc.stdin.close()
        writer_proc.wait()

    total_time = time.time() - start_time
    avg_fps = processed_count / total_time if total_time > 0 else 0
    full_ep_eta_hours = (total_video_frames / avg_fps) / 3600 if avg_fps > 0 else 0

    print(f"\n\n==================================================")
    print(f" Upscale Benchmark & Quality Report")
    print(f"==================================================")
    print(f"Total Frames Processed: {processed_count} frames")
    print(f"Total Elapsed Time: {total_time:.2f}s ({total_time/60:.2f} mins)")
    print(f"Average Upscale Speed: {avg_fps:.3f} FPS ({1.0/avg_fps:.2f}s per frame)")
    print(f"Peak VRAM Consumed: {torch.cuda.max_memory_allocated() / (1024**2):.1f} MB / 6144 MB")
    print(f"Pipeline Stage Breakdown per frame:")
    if processed_count > 0:
        print(f"  - Step 1 (1x_PureVision Denoise): {pv_time_total/processed_count*1000:.1f} ms")
        print(f"  - Step 2 (4x Super-Resolution):   {sr_time_total/processed_count*1000:.1f} ms")
        print(f"  - Step 3 (5K -> 4K Downsample):   {down_time_total/processed_count*1000:.1f} ms")
    print(f"Projected Full Episode Render Time ({total_video_frames} frames): {full_ep_eta_hours:.2f} hours")
    print(f"Output File: {output_video}")
    print(f"==================================================")


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="PureVision + SPAN Next-Gen AI 4K Video Upscaler")
    parser.add_argument("--input", default="downloads/ep2.mp4", help="Path to input video")
    parser.add_argument("--output", default="downloads/ep2_span_4k_test.mp4", help="Path to output 4K video")
    parser.add_argument("--pv_model", default="models/1x_PureVision.pth", help="Path to Step 1 PureVision model")
    parser.add_argument("--sr_model", default="models/4xPurePhoto-span.pth", help="Path to Step 2 4x SR model (default: SPAN)")
    parser.add_argument("--start_sec", type=float, default=0.0, help="Start time in seconds")
    parser.add_argument("--duration", type=float, default=None, help="Duration in seconds (e.g. 5.0 for benchmark, omit for full video)")
    parser.add_argument("--tile_size", type=int, default=512, help="Tile size if tiling is forced")
    parser.add_argument("--tile_pad", type=int, default=32, help="Overlap padding for tiling")
    parser.add_argument("--force_tile", action="store_true", help="Force tiling on SPAN (default is disabled for max speed)")
    parser.add_argument("--cq", type=int, default=19, help="NVENC CQ quality factor")

    args = parser.parse_args()
    upscale_video(
        input_video=args.input,
        output_video=args.output,
        pv_model_path=args.pv_model,
        sr_model_path=args.sr_model,
        start_sec=args.start_sec,
        duration=args.duration,
        tile_size=args.tile_size,
        tile_pad=args.tile_pad,
        force_tile=args.force_tile,
        cq=args.cq,
    )
