luca115's picture
Upload folder using huggingface_hub
e281960 verified
Raw
History Blame Contribute Delete
6.2 kB
import os
import subprocess
import tempfile
import cv2
import gradio as gr
import numpy as np
import spaces
import torch
import torch.nn.functional as F
from huggingface_hub import hf_hub_download
from imageio_ffmpeg import get_ffmpeg_exe
from train_log.IFNet_HDv3 import IFNet
from upsampler_theme import UPSAMPLER_THEME, UPSAMPLER_CSS, footer_html, header_html
# Official Practical-RIFE v4.25 weights (MIT), mirrored unchanged on our org
# so no third-party repo can drift underneath the Space.
WEIGHTS = hf_hub_download("Upsampler/rife-4-25", "flownet.pkl")
MAX_SECONDS = 6.5
MAX_SIDE = 1280
# Consecutive frames more different than this (mean abs diff on gray
# thumbnails) are treated as a hard cut: duplicate instead of interpolate,
# so cuts don't produce ghost blends.
SCENE_CUT_DIFF = 60.0
flownet = IFNet()
state = torch.load(WEIGHTS, map_location="cpu")
flownet.load_state_dict(
{k.replace("module.", ""): v for k, v in state.items() if "module." in k},
strict=False,
)
flownet = flownet.to("cuda").eval()
def probe(path):
cap = cv2.VideoCapture(path)
fps = cap.get(cv2.CAP_PROP_FPS) or 30
n = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
cap.release()
return fps, n, w, h
def get_duration(video, factor, progress=None):
if not video:
return 10
try:
fps, n, w, h = probe(video)
except Exception:
return 60
mult = int(factor[0]) if factor else 2
# ~35ms per generated frame at 720p on the ZeroGPU slice, plus decode,
# encode, and model-to-GPU overhead; never over-request.
est = 20 + int(n * (mult - 1) * 0.06 * max(1.0, (w * h) / (1280 * 720)))
return min(110, est)
def midpoint(img0, img1):
imgs = torch.cat((img0, img1), 1)
scale_list = [16, 8, 4, 2, 1]
_, _, merged = flownet(imgs, 0.5, scale_list)
return merged[-1]
@spaces.GPU(duration=get_duration)
def interpolate(video, factor, progress=gr.Progress()):
if not video:
raise gr.Error("Please upload a video first.")
fps, n, w, h = probe(video)
seconds = n / max(fps, 1)
if seconds > MAX_SECONDS:
raise gr.Error(
f"This clip is {seconds:.1f}s; the free tool accepts up to {MAX_SECONDS:.0f}s. "
"Trim it first and try again."
)
if max(w, h) > MAX_SIDE:
raise gr.Error(
f"This clip is {w}x{h}; the free tool accepts up to {MAX_SIDE}px on the longest side."
)
mult = int(factor[0])
cap = cv2.VideoCapture(video)
frames = []
ok, frame = cap.read()
while ok:
frames.append(frame)
ok, frame = cap.read()
cap.release()
if len(frames) < 2:
raise gr.Error("Could not read enough frames from this video.")
ph = ((h - 1) // 64 + 1) * 64
pw = ((w - 1) // 64 + 1) * 64
def to_tensor(frame):
t = torch.from_numpy(frame).to("cuda", non_blocking=True)
t = t.permute(2, 0, 1).float().unsqueeze(0) / 255.0
return F.pad(t, (0, pw - w, 0, ph - h))
def to_frame(t):
out = (t[0][:, :h, :w].permute(1, 2, 0) * 255.0).clamp(0, 255)
return out.byte().cpu().numpy()
thumbs = [
cv2.cvtColor(cv2.resize(f, (64, 36)), cv2.COLOR_BGR2GRAY).astype(np.float32)
for f in frames
]
out_path = os.path.join(tempfile.mkdtemp(), "interpolated.mp4")
ffmpeg = get_ffmpeg_exe()
enc = subprocess.Popen(
[
ffmpeg, "-y", "-f", "rawvideo", "-pix_fmt", "bgr24",
"-s", f"{w}x{h}", "-r", str(fps * mult), "-i", "-",
"-i", video, "-map", "0:v", "-map", "1:a?", "-c:a", "copy",
"-c:v", "libx264", "-preset", "fast", "-crf", "18",
"-pix_fmt", "yuv420p", "-shortest", out_path,
],
stdin=subprocess.PIPE,
stderr=subprocess.DEVNULL,
)
with torch.inference_mode():
for i in range(len(frames) - 1):
progress((i + 1) / len(frames), desc="Interpolating frames")
enc.stdin.write(frames[i].tobytes())
is_cut = float(np.abs(thumbs[i] - thumbs[i + 1]).mean()) > SCENE_CUT_DIFF
if is_cut:
for _ in range(mult - 1):
enc.stdin.write(frames[i].tobytes())
continue
t0, t1 = to_tensor(frames[i]), to_tensor(frames[i + 1])
if mult == 2:
enc.stdin.write(to_frame(midpoint(t0, t1)).tobytes())
else: # 4x: recursive midpoints -> quarter timesteps
mid = midpoint(t0, t1)
for t in (midpoint(t0, mid), mid, midpoint(mid, t1)):
enc.stdin.write(to_frame(t).tobytes())
enc.stdin.write(frames[-1].tobytes())
enc.stdin.close()
if enc.wait() != 0:
raise gr.Error("Video encoding failed. Please try a different clip.")
return out_path
with gr.Blocks(theme=UPSAMPLER_THEME, css=UPSAMPLER_CSS) as demo:
gr.HTML(
header_html(
"RIFE Video Frame Interpolation",
"Double or quadruple your video's frame rate for smooth slow motion",
)
)
with gr.Row():
with gr.Column(scale=1):
input_video = gr.Video(label="Video (up to 6s, 1280px)")
factor = gr.Radio(["2x frames", "4x frames"], value="2x frames", label="Interpolation")
run = gr.Button("Interpolate", variant="primary")
with gr.Column(scale=1):
output_video = gr.Video(label="Result", autoplay=True)
run.click(interpolate, inputs=[input_video, factor], outputs=output_video, api_name="interpolate")
gr.HTML(
footer_html(
"RIFE (Real-Time Intermediate Flow Estimation, Practical-RIFE v4.25) "
"generates new in-between frames for any video, turning choppy clips "
"into smooth slow motion and converting frame rates like 30fps to "
"60fps or 120fps, with scene-cut detection to avoid ghosting.",
"https://upsampler.com/free-video-frame-interpolation-no-signup",
"free video frame interpolation tool",
)
)
if __name__ == "__main__":
demo.launch(ssr_mode=False)