atharvak30 commited on
Commit
d3d126c
Β·
verified Β·
1 Parent(s): 0a8f899

Upload 3 files

Browse files
Files changed (3) hide show
  1. README.md +17 -0
  2. app.py +255 -0
  3. requirements.txt +5 -0
README.md ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ title: AI Video Enhancer
3
+ emoji: 🎬
4
+ colorFrom: blue
5
+ colorTo: purple
6
+ sdk: gradio
7
+ sdk_version: "5.12.0"
8
+ app_file: app.py
9
+ pinned: false
10
+ license: mit
11
+ hf_oauth: true
12
+ hf_oauth_expiration_minutes: 480
13
+ ---
14
+
15
+ # 🎬 AI Video Enhancer
16
+
17
+ Upscale your videos using **Swin2SR** AI super-resolution.
app.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # app.py
2
+ # AI Video Enhancer with OAuth login for ZeroGPU quota
3
+
4
+ import os
5
+ import shutil
6
+ import subprocess
7
+ import tempfile
8
+ import time
9
+ from pathlib import Path
10
+
11
+ import gradio as gr
12
+ import spaces
13
+ import torch
14
+ import numpy as np
15
+ from PIL import Image
16
+
17
+ # ============== Configuration ==============
18
+ TEMP_DIR = Path(tempfile.gettempdir()) / "video_enhancer"
19
+ TEMP_DIR.mkdir(parents=True, exist_ok=True)
20
+
21
+ MODELS = {
22
+ 2: "caidas/swin2SR-classical-sr-x2-64",
23
+ 4: "caidas/swin2SR-classical-sr-x4-48",
24
+ }
25
+
26
+
27
+ # ============== Video Utilities ==============
28
+ def run_ffmpeg(cmd: list) -> str:
29
+ result = subprocess.run(cmd, capture_output=True, text=True)
30
+ if result.returncode != 0:
31
+ raise RuntimeError(f"FFmpeg error: {result.stderr}")
32
+ return result.stdout
33
+
34
+
35
+ def get_video_info(path: str) -> dict:
36
+ import json
37
+ cmd = [
38
+ "ffprobe", "-v", "quiet", "-print_format", "json",
39
+ "-show_streams", "-show_format", path
40
+ ]
41
+ result = subprocess.run(cmd, capture_output=True, text=True)
42
+ if result.returncode != 0:
43
+ raise RuntimeError("Could not read video info")
44
+
45
+ data = json.loads(result.stdout)
46
+ video_stream = next((s for s in data.get("streams", []) if s["codec_type"] == "video"), None)
47
+
48
+ if not video_stream:
49
+ raise RuntimeError("No video stream found")
50
+
51
+ fps_str = video_stream.get("r_frame_rate", "30/1")
52
+ if "/" in fps_str:
53
+ num, den = map(float, fps_str.split("/"))
54
+ fps = num / den if den != 0 else 30.0
55
+ else:
56
+ fps = float(fps_str)
57
+
58
+ return {
59
+ "width": int(video_stream.get("width", 0)),
60
+ "height": int(video_stream.get("height", 0)),
61
+ "fps": fps,
62
+ "duration": float(data.get("format", {}).get("duration", 0)),
63
+ "has_audio": any(s["codec_type"] == "audio" for s in data.get("streams", []))
64
+ }
65
+
66
+
67
+ def extract_frames(video_path: str, output_dir: Path, max_frames: int = None) -> int:
68
+ output_dir.mkdir(parents=True, exist_ok=True)
69
+ cmd = ["ffmpeg", "-y", "-i", video_path, "-vsync", "0"]
70
+ if max_frames:
71
+ cmd.extend(["-vframes", str(max_frames)])
72
+ cmd.append(str(output_dir / "frame_%06d.png"))
73
+ run_ffmpeg(cmd)
74
+ return len(list(output_dir.glob("*.png")))
75
+
76
+
77
+ def create_video(frames_dir: Path, output_path: str, fps: float, audio_source: str = None):
78
+ temp_video = str(Path(output_path).parent / "temp_no_audio.mp4")
79
+
80
+ run_ffmpeg([
81
+ "ffmpeg", "-y", "-framerate", str(fps),
82
+ "-i", str(frames_dir / "frame_%06d.png"),
83
+ "-c:v", "libx264", "-preset", "medium", "-crf", "18",
84
+ "-pix_fmt", "yuv420p", temp_video
85
+ ])
86
+
87
+ if audio_source:
88
+ try:
89
+ info = get_video_info(audio_source)
90
+ if info["has_audio"]:
91
+ run_ffmpeg([
92
+ "ffmpeg", "-y", "-i", temp_video, "-i", audio_source,
93
+ "-c:v", "copy", "-c:a", "aac", "-map", "0:v:0", "-map", "1:a:0",
94
+ output_path
95
+ ])
96
+ os.remove(temp_video)
97
+ return
98
+ except:
99
+ pass
100
+
101
+ shutil.move(temp_video, output_path)
102
+
103
+
104
+ # ============== Image Enhancement ==============
105
+ @spaces.GPU(duration=60)
106
+ def enhance_batch(images: list, scale: int = 2) -> list:
107
+ """Enhance a batch of images using Swin2SR."""
108
+ from transformers import Swin2SRForImageSuperResolution, Swin2SRImageProcessor
109
+
110
+ model_id = MODELS[scale]
111
+ processor = Swin2SRImageProcessor.from_pretrained(model_id)
112
+ model = Swin2SRForImageSuperResolution.from_pretrained(model_id)
113
+
114
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
115
+ model = model.to(device).eval()
116
+
117
+ enhanced = []
118
+
119
+ for img in images:
120
+ inputs = processor(img, return_tensors="pt").to(device)
121
+
122
+ with torch.no_grad():
123
+ outputs = model(**inputs)
124
+
125
+ output = outputs.reconstruction.squeeze().cpu().clamp(0, 1)
126
+ output = output.permute(1, 2, 0).numpy()
127
+ output = (output * 255).astype(np.uint8)
128
+
129
+ enhanced.append(Image.fromarray(output))
130
+
131
+ return enhanced
132
+
133
+
134
+ # ============== Main Processing ==============
135
+ def process_video(video_file, scale: int, profile: gr.OAuthProfile | None, progress=gr.Progress()):
136
+ """Main video processing function"""
137
+
138
+ # Show login status
139
+ user_msg = f"Processing as: {profile.name}" if profile else "Not logged in (limited quota)"
140
+ print(user_msg)
141
+
142
+ if video_file is None:
143
+ return "⚠️ Please upload a video file.", None
144
+
145
+ job_id = f"job_{int(time.time() * 1000)}"
146
+ work_dir = TEMP_DIR / job_id
147
+ work_dir.mkdir(parents=True, exist_ok=True)
148
+
149
+ input_path = work_dir / "input_video"
150
+ frames_dir = work_dir / "frames"
151
+ output_path = work_dir / "output.mp4"
152
+
153
+ try:
154
+ shutil.copy(video_file, input_path)
155
+
156
+ progress(0.05, "Analyzing video...")
157
+ info = get_video_info(str(input_path))
158
+
159
+ max_frames = min(int(info["fps"] * 10), 300)
160
+
161
+ progress(0.1, "Extracting frames...")
162
+ num_frames = extract_frames(str(input_path), frames_dir, max_frames)
163
+
164
+ if num_frames == 0:
165
+ return "❌ Could not extract frames from video.", None
166
+
167
+ progress(0.15, f"Processing {num_frames} frames...")
168
+
169
+ frame_files = sorted(frames_dir.glob("*.png"))
170
+ batch_size = 4
171
+
172
+ for batch_start in range(0, len(frame_files), batch_size):
173
+ batch_end = min(batch_start + batch_size, len(frame_files))
174
+ batch_files = frame_files[batch_start:batch_end]
175
+
176
+ batch_images = [Image.open(f).convert("RGB") for f in batch_files]
177
+
178
+ try:
179
+ enhanced_images = enhance_batch(batch_images, scale)
180
+
181
+ for img, path in zip(enhanced_images, batch_files):
182
+ img.save(path)
183
+
184
+ except Exception as e:
185
+ print(f"GPU enhancement failed: {e}, using fallback...")
186
+ for img, path in zip(batch_images, batch_files):
187
+ new_size = (img.width * scale, img.height * scale)
188
+ img.resize(new_size, Image.LANCZOS).save(path)
189
+
190
+ pct = 0.15 + 0.75 * (batch_end / len(frame_files))
191
+ progress(pct, f"Enhanced {batch_end}/{len(frame_files)} frames")
192
+
193
+ progress(0.92, "Creating video...")
194
+ create_video(frames_dir, str(output_path), info["fps"], str(input_path))
195
+
196
+ shutil.rmtree(frames_dir, ignore_errors=True)
197
+
198
+ out_info = get_video_info(str(output_path))
199
+
200
+ progress(1.0, "Done!")
201
+ return (
202
+ f"βœ… Enhanced: {info['width']}x{info['height']} β†’ "
203
+ f"{out_info['width']}x{out_info['height']} ({scale}x)",
204
+ str(output_path)
205
+ )
206
+
207
+ except Exception as e:
208
+ shutil.rmtree(work_dir, ignore_errors=True)
209
+ return f"❌ Error: {str(e)}", None
210
+
211
+
212
+ def greet(profile: gr.OAuthProfile | None) -> str:
213
+ if profile is None:
214
+ return "πŸ‘‹ **Not logged in** - Click 'Sign in with Hugging Face' above for full GPU quota!"
215
+ return f"πŸ‘‹ Welcome **{profile.name}**! You have Pro GPU quota."
216
+
217
+
218
+ # ============== Gradio Interface ==============
219
+ with gr.Blocks(title="AI Video Enhancer", theme=gr.themes.Soft()) as demo:
220
+ gr.Markdown("# 🎬 AI Video Enhancer")
221
+
222
+ # Login button and status
223
+ gr.LoginButton()
224
+ login_status = gr.Markdown()
225
+ demo.load(greet, inputs=None, outputs=login_status)
226
+
227
+ gr.Markdown("---")
228
+
229
+ with gr.Row():
230
+ with gr.Column(scale=2):
231
+ video_input = gr.File(
232
+ label="πŸ“ Upload Video",
233
+ file_types=[".mp4", ".avi", ".mov", ".mkv", ".webm"]
234
+ )
235
+
236
+ scale_selector = gr.Radio(
237
+ choices=[2, 4],
238
+ value=2,
239
+ label="πŸ” Upscale Factor"
240
+ )
241
+
242
+ enhance_btn = gr.Button("πŸš€ Enhance Video", variant="primary", size="lg")
243
+ status_text = gr.Textbox(label="Status", interactive=False)
244
+
245
+ with gr.Column(scale=1):
246
+ video_output = gr.Video(label="πŸŽ₯ Enhanced Video")
247
+
248
+ enhance_btn.click(
249
+ fn=process_video,
250
+ inputs=[video_input, scale_selector],
251
+ outputs=[status_text, video_output]
252
+ )
253
+
254
+ if __name__ == "__main__":
255
+ demo.launch()
requirements.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ spaces
2
+ torch
3
+ transformers
4
+ Pillow
5
+ numpy