atharvak30's picture
Upload 4 files
22d966c verified
Raw History Blame
9.06 kB
# app.py
# AI Video Enhancer 4K - Gradio app for Hugging Face Spaces
# Simplified version for better ZeroGPU compatibility
import os
import shutil
import subprocess
import tempfile
import time
from pathlib import Path
from typing import Tuple
import gradio as gr
import spaces
import torch
import numpy as np
from PIL import Image
import cv2
from huggingface_hub import hf_hub_download
# Config
TEMP_DIR = Path(tempfile.gettempdir()) / "hf_video_enhancer"
TEMP_DIR.mkdir(parents=True, exist_ok=True)
def run_cmd(cmd):
p = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
if p.returncode != 0:
raise RuntimeError(f"Command failed: {p.stderr.decode()}")
return p.stdout.decode()
def probe_video(video_path: str) -> Tuple[float, int, int, float]:
cmd = [
"ffprobe", "-v", "error",
"-select_streams", "v:0",
"-show_entries", "stream=width,height,duration,r_frame_rate",
"-of", "default=noprint_wrappers=1:nokey=0",
video_path
]
p = subprocess.run(cmd, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
out = p.stdout.decode()
width = height = 0
duration = 0.0
fps = 30.0
for line in out.splitlines():
if line.startswith("width="):
width = int(line.split("=")[1])
elif line.startswith("height="):
height = int(line.split("=")[1])
elif line.startswith("duration="):
try:
duration = float(line.split("=")[1])
except:
pass
elif line.startswith("r_frame_rate="):
try:
fps_str = line.split("=")[1]
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den)
else:
fps = float(fps_str)
except:
pass
return duration, width, height, fps
def extract_frames(video_path: str, frames_dir: Path):
frames_dir.mkdir(parents=True, exist_ok=True)
run_cmd([
"ffmpeg", "-y", "-i", video_path,
"-vsync", "0",
str(frames_dir / "%06d.png")
])
def reassemble_video(frames_dir: Path, audio_src: str, out_path: str, fps: float = 30.0):
tmp_video = str(frames_dir.parent / "tmp_video.mp4")
run_cmd([
"ffmpeg", "-y", "-framerate", str(fps),
"-i", str(frames_dir / "%06d.png"),
"-c:v", "libx264", "-preset", "veryfast", "-pix_fmt", "yuv420p",
"-crf", "18", tmp_video
])
p = subprocess.run(
["ffprobe", "-v", "error", "-select_streams", "a", "-show_entries",
"stream=codec_type", "-of", "default=noprint_wrappers=1", audio_src],
stdout=subprocess.PIPE, stderr=subprocess.PIPE
)
if p.stdout.decode().strip():
run_cmd([
"ffmpeg", "-y", "-i", tmp_video, "-i", audio_src,
"-c:v", "copy", "-c:a", "aac",
"-map", "0:v:0", "-map", "1:a:0", out_path
])
os.remove(tmp_video)
else:
shutil.move(tmp_video, out_path)
# Simple upscaling function using torch interpolation as fallback
def simple_upscale(img: np.ndarray, scale: int) -> np.ndarray:
"""Simple bicubic upscaling using OpenCV"""
h, w = img.shape[:2]
return cv2.resize(img, (w * scale, h * scale), interpolation=cv2.INTER_CUBIC)
@spaces.GPU(duration=120)
def enhance_with_realesrgan(frames_dir: str, scale: int = 4) -> int:
"""
Enhance frames using Real-ESRGAN via Spandrel.
Separated function with GPU decorator for cleaner ZeroGPU handling.
"""
from spandrel import ImageModelDescriptor, ModelLoader
frames_path = Path(frames_dir)
frame_files = sorted(frames_path.glob("*.png"))
total = len(frame_files)
if total == 0:
return 0
# Download and load model
if scale == 2:
model_path = hf_hub_download(repo_id="ai-forever/Real-ESRGAN", filename="RealESRGAN_x2.pth")
else:
model_path = hf_hub_download(repo_id="ai-forever/Real-ESRGAN", filename="RealESRGAN_x4.pth")
model = ModelLoader().load_from_file(model_path)
assert isinstance(model, ImageModelDescriptor)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = model.to(device).eval()
print(f"Model loaded on {device}, processing {total} frames...")
for idx, frame_path in enumerate(frame_files):
# Read image
img = cv2.imread(str(frame_path))
if img is None:
continue
img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# Convert to tensor
tensor = torch.from_numpy(img_rgb).permute(2, 0, 1).float().div(255.0)
tensor = tensor.unsqueeze(0).to(device)
# Process
with torch.no_grad():
output = model(tensor)
# Convert back
output = output.squeeze(0).cpu().clamp(0, 1).mul(255).byte()
output = output.permute(1, 2, 0).numpy()
output_bgr = cv2.cvtColor(output, cv2.COLOR_RGB2BGR)
# Save
cv2.imwrite(str(frame_path), output_bgr)
if (idx + 1) % 5 == 0:
print(f"Processed {idx + 1}/{total}")
return total
def process_video(video_file, scale: int = 4) -> Tuple[str, str]:
"""Main video processing - handles file I/O outside GPU function"""
if video_file is None:
return "⚠️ Please upload a video file.", None
ts = int(time.time() * 1000)
base_dir = TEMP_DIR / f"job_{ts}"
base_dir.mkdir(parents=True, exist_ok=True)
in_path = base_dir / "input_video"
try:
shutil.copy(video_file, in_path)
except Exception as e:
return f"Error: {e}", None
try:
duration, w, h, fps = probe_video(str(in_path))
except Exception as e:
shutil.rmtree(base_dir, ignore_errors=True)
return f"Error probing video: {e}", None
if duration <= 0:
shutil.rmtree(base_dir, ignore_errors=True)
return "Could not determine video duration.", None
# Limit for ZeroGPU - process max ~30 seconds of video
max_frames = int(fps * 30) # ~30 seconds worth
print(f"Video: {w}x{h}, {duration:.1f}s, {fps:.1f}fps")
frames_dir = base_dir / "frames"
try:
extract_frames(str(in_path), frames_dir)
except Exception as e:
shutil.rmtree(base_dir, ignore_errors=True)
return f"Failed extracting frames: {e}", None
frame_files = sorted(frames_dir.glob("*.png"))
num_frames = len(frame_files)
# Limit frames if too many
if num_frames > max_frames:
print(f"Limiting from {num_frames} to {max_frames} frames")
for f in frame_files[max_frames:]:
f.unlink()
num_frames = max_frames
print(f"Processing {num_frames} frames...")
try:
enhanced = enhance_with_realesrgan(str(frames_dir), scale)
print(f"Enhanced {enhanced} frames")
except Exception as e:
print(f"Enhancement failed: {e}")
# Fallback to simple upscaling
print("Using fallback bicubic upscaling...")
try:
for fp in sorted(frames_dir.glob("*.png")):
img = cv2.imread(str(fp))
if img is not None:
upscaled = simple_upscale(img, scale)
cv2.imwrite(str(fp), upscaled)
except Exception as e2:
shutil.rmtree(base_dir, ignore_errors=True)
return f"Enhancement failed: {e}", None
out_video = base_dir / "enhanced_output.mp4"
try:
reassemble_video(frames_dir, str(in_path), str(out_video), fps)
except Exception as e:
shutil.rmtree(base_dir, ignore_errors=True)
return f"Failed reassembling: {e}", None
shutil.rmtree(frames_dir, ignore_errors=True)
try:
_, out_w, out_h, _ = probe_video(str(out_video))
return f"βœ… Done! {w}x{h} β†’ {out_w}x{out_h}", str(out_video)
except:
return "βœ… Done!", str(out_video)
# Gradio UI
with gr.Blocks(title="AI Video Enhancer", theme=gr.themes.Soft()) as demo:
gr.Markdown("# 🎬 AI Video Enhancer")
gr.Markdown("Upscale videos using Real-ESRGAN AI enhancement.")
# LOGIN BUTTON - This allows ZeroGPU to recognize your Pro account
gr.LoginButton()
with gr.Row():
with gr.Column(scale=2):
video_in = gr.File(label="Upload video", file_types=[".mp4", ".avi", ".mov", ".mkv", ".webm"])
scale_choice = gr.Radio(choices=[2, 4], value=4, label="Upscale Factor")
btn = gr.Button("πŸš€ Enhance", variant="primary")
status = gr.Textbox(label="Status", interactive=False)
with gr.Column(scale=1):
out_video = gr.Video(label="Result")
gr.Markdown("**Note:** Limited to ~30 seconds for ZeroGPU. Longer videos will be truncated.")
btn.click(fn=process_video, inputs=[video_in, scale_choice], outputs=[status, out_video])
if __name__ == "__main__":
demo.launch()