swift-vr / app.py
multimodalart's picture
multimodalart HF Staff
Update app.py
884f557 verified
Raw
History Blame Contribute Delete
7.47 kB
"""SwiftVR: Real-Time One-Step Generative Video Restoration.
Gradio Space demo that loads the SwiftVR model from H-oliday/SwiftVR on the
Hugging Face Hub and restores / upscales user-uploaded videos in real time.
"""
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces
import sys
import tempfile
import traceback
import math
# Make the bundled swiftvr package importable
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import torch
import torch.nn.functional as F
import gradio as gr
import numpy as np
import imageio
from PIL import Image
from huggingface_hub import snapshot_download
import decord
decord.bridge.set_bridge("torch")
from swiftvr import SwiftVRPipeline
from swiftvr.io import (
get_video_info, iter_video_clips_fixed_scheme,
preprocess_clip_uint8, crop_spatial_padding_ntchw,
ntchw_to_uint8_frames, open_stream_video_writer,
)
from swiftvr.streaming.chunk import ChunkType
CHECKPOINT_REPO = "H-oliday/SwiftVR"
CHECKPOINT_DIR = "/tmp/swiftvr_checkpoints"
# --- Load model at module scope (ZeroGPU rule) --- #
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
print("Downloading SwiftVR checkpoints...")
snapshot_download(
repo_id=CHECKPOINT_REPO,
local_dir=CHECKPOINT_DIR,
repo_type="model",
)
print("Checkpoints downloaded. Loading model...")
pipe = SwiftVRPipeline.from_pretrained(CHECKPOINT_DIR)
pipe.to("cuda", dtype="bfloat16", attention_backend="sdpa")
print("SwiftVR model loaded and moved to CUDA with SDPA attention.")
def _aligned_pad(size, multiple=32):
return (multiple - size % multiple) % multiple
@spaces.GPU(duration=90)
def restore_video(
input_video: str,
upscale: int = 2,
clip_len: int = 24,
quality: int = 60,
):
"""Restore a low-quality video using SwiftVR.
Args:
input_video: Path to the input low-quality video file.
upscale: Upscale factor (2, 4, or 8).
clip_len: Processing chunk size (must be multiple of 4).
quality: Output quality 0-100 (maps to x265 CRF).
Returns:
Path to the restored output video, or an error message.
"""
if input_video is None:
return None, "Please upload or select a video first."
tmp_dir = tempfile.mkdtemp()
output_path = os.path.join(tmp_dir, "restored.mp4")
try:
# Single-threaded pipeline to avoid ZeroGPU fork conflicts
device = pipe.device
dtype = pipe.dtype
# Get video info
raw_total, lq_h, lq_w, src_fps = get_video_info(input_video, fallback_fps=30)
total_frames = 4 * ((raw_total - 1) // 4) + 1
# Compute target size
out_h = lq_h * upscale
out_w = lq_w * upscale
pad_h = _aligned_pad(out_h)
pad_w = _aligned_pad(out_w)
# Reset streaming state
pipe.tae_stream.reset()
pipe.dit_stream.reset()
pipe.dit_stream.overlap = 0
n_lat = clip_len // 4
prev_dit_out_cpu = None
# Open video writer
writer = open_stream_video_writer(
output_path, fps=src_fps, video_format="",
preset="", quality=quality)
# Process chunks (single-threaded, no pipeline threads)
clips = iter_video_clips_fixed_scheme(
input_video, clip_len=clip_len,
total_frames=total_frames,
crop_h=lq_h, crop_w=lq_w)
total_written = 0
for spec, cpu_rgb in clips:
# Move to GPU
gpu_rgb = cpu_rgb.to(device=device)
# Preprocess
clip_rgb = preprocess_clip_uint8(
gpu_rgb, out_h=out_h, out_w=out_w,
mode=pipe.upscale_mode, pad_h=pad_h, pad_w=pad_w,
dtype=dtype)
# Encode
z = pipe.tae_stream.encode_chunk_fixed(clip_rgb, spec)
# Denoise
if spec.ctype == ChunkType.LAST:
z_ntchw = pipe.dit_stream.denoise_last_chunk(
z, spec, pipe.prompt_emb, prev_dit_out_cpu,
n_lat, device, dtype)
else:
z_bcfhw = z.permute(0, 2, 1, 3, 4).contiguous()
z_den = pipe.dit_stream.denoise(z_bcfhw, pipe.prompt_emb)
z_ntchw = z_den.permute(0, 2, 1, 3, 4).contiguous()
prev_dit_out_cpu = z_bcfhw[:, :, -n_lat:].detach().cpu().clone()
# Decode
rgb_out = pipe.tae_stream.decode_chunk_fixed(z_ntchw, spec)
if rgb_out is not None and rgb_out.shape[1] > 0:
rgb_out = crop_spatial_padding_ntchw(rgb_out, pad_h, pad_w).detach()
frames = ntchw_to_uint8_frames(rgb_out)
if frames is not None:
for frame in frames:
writer.append_data(frame)
total_written += frames.shape[0]
# Clean up GPU memory
del gpu_rgb, clip_rgb, z, z_ntchw, rgb_out
torch.cuda.empty_cache()
writer.close()
msg = f"Restored {total_written} frames from {raw_total} input frames at {src_fps:.1f} fps"
return output_path, msg
except Exception as e:
err_msg = f"Error: {e}\n{traceback.format_exc()}"
print(err_msg)
return None, err_msg
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
with gr.Blocks() as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""
# SwiftVR: Real-Time One-Step Generative Video Restoration
Upload a low-quality video to restore and upscale it in real time using the SwiftVR model.
[Paper](https://arxiv.org/abs/2606.09516) | [GitHub](https://github.com/H-oliday/SwiftVR) | [Model](https://huggingface.co/H-oliday/SwiftVR)
"""
)
with gr.Row():
input_video = gr.Video(label="Input Video (low quality)", sources=["upload"])
output_video = gr.Video(label="Restored Video", interactive=False)
status = gr.Textbox(label="Status", interactive=False)
with gr.Accordion("Advanced Settings", open=False):
upscale = gr.Slider(
minimum=1, maximum=4, step=1, value=2,
label="Upscale Factor",
info="How many times to upscale the input video resolution. Use 2 for faster processing."
)
clip_len = gr.Slider(
minimum=4, maximum=48, step=4, value=24,
label="Clip Length",
info="Processing chunk size (must be multiple of 4). Larger = more context, slower."
)
quality = gr.Slider(
minimum=0, maximum=100, step=5, value=60,
label="Output Quality",
info="0-100, maps to x265 CRF. Higher = better quality, larger file."
)
run_btn = gr.Button("Restore Video", variant="primary")
gr.Examples(
examples=[
["examples/cat.mp4", 2, 24, 60],
],
inputs=[input_video, upscale, clip_len, quality],
outputs=[output_video, status],
fn=restore_video,
cache_examples=True,
cache_mode="lazy",
)
run_btn.click(
fn=restore_video,
inputs=[input_video, upscale, clip_len, quality],
outputs=[output_video, status],
)
if __name__ == "__main__":
demo.launch(mcp_server=True, show_error=True)