import os
import time
import subprocess
import cv2
import torch
import tensorrt as trt

ENGINE_PATH = "/workspace/srvgg_x4.engine"
INPUT_DIR = "/workspace/inputs"
in_video = os.path.join(INPUT_DIR, "ep2.mp4")

runtime = trt.Runtime(trt.Logger(trt.Logger.WARNING))
with open(ENGINE_PATH, "rb") as f:
    engine = runtime.deserialize_cuda_engine(f.read())
context = engine.create_execution_context()

stream = torch.cuda.Stream()
d_input = torch.empty(1, 3, 720, 1280, device="cuda", dtype=torch.float32).contiguous()
d_output = torch.empty(1, 3, 2880, 5120, device="cuda", dtype=torch.float32).contiguous()
context.set_tensor_address("input", d_input.data_ptr())
context.set_tensor_address("output", d_output.data_ptr())

cap = cv2.VideoCapture(in_video)

t_read = 0
t_prep = 0
t_trt = 0
t_interp = 0
t_cpu = 0

N = 20
for _ in range(N):
    t0 = time.time()
    ret, frame = cap.read()
    t_read += time.time() - t0

    t0 = time.time()
    t_frame = torch.from_numpy(frame).to("cuda", non_blocking=True)
    t_rgb = t_frame[:, :, [2, 1, 0]].permute(2, 0, 1).unsqueeze(0).float().div_(255.0)
    d_input.copy_(t_rgb)
    torch.cuda.synchronize()
    t_prep += time.time() - t0

    t0 = time.time()
    with torch.cuda.stream(stream):
        context.execute_async_v3(stream.cuda_stream)
    stream.synchronize()
    t_trt += time.time() - t0

    t0 = time.time()
    out_4k = torch.nn.functional.interpolate(d_output, size=(2160, 3840), mode="bicubic", align_corners=False)
    out_uint8 = out_4k.squeeze(0).clamp_(0.0, 1.0).mul_(255.0).to(torch.uint8)
    out_bgr = out_uint8[[2, 1, 0], :, :].permute(1, 2, 0)
    torch.cuda.synchronize()
    t_interp += time.time() - t0

    t0 = time.time()
    arr = out_bgr.cpu().numpy().tobytes()
    t_cpu += time.time() - t0

print(f"Per-frame breakdown over {N} frames:")
print(f"  OpenCV read:    {t_read/N*1000:.1f} ms")
print(f"  GPU prep:       {t_prep/N*1000:.1f} ms")
print(f"  TensorRT:       {t_trt/N*1000:.1f} ms")
print(f"  GPU interp 4K:  {t_interp/N*1000:.1f} ms")
print(f"  GPU->CPU numpy: {t_cpu/N*1000:.1f} ms")
total_ms = (t_read + t_prep + t_trt + t_interp + t_cpu) / N * 1000
print(f"Total pipeline without FFmpeg: {total_ms:.1f} ms -> {1000/total_ms:.1f} FPS")
