someone-in-the-world's picture
Merge remote-tracking branch 'origin/develop' into add-mypy-strict-typecheck
c8ec0f7
Raw History Blame Contribute Delete
5.94 kB
"""RIFE v4.26 frame interpolation (FR-3): GPU, fp16, runs inside the same @spaces.GPU
allocation as generation.
Model code/weights come from `thornmaze/RIFE`, a packaged build of hzwer/Practical-RIFE
(MIT, see LICENSES/RIFE-LICENSE). The downloaded `train_log/RIFE_HDv3.py` module imports
`from model.warplayer import warp` (and, at module load time, from `model.loss` and
`model.pytorch_msssim`) expecting a sibling `model` package — see the vendored copies in
model/warplayer.py, model/loss.py, model/pytorch_msssim/.
The zip download/unzip itself needs no GPU, so `ensure_weights_downloaded()` is called from
app.py at process startup (like `load_pipeline()`), before any `@spaces.GPU` call — otherwise
a cold container pays for that fetch out of the metered GPU allocation on its first request.
"""
from __future__ import annotations
import subprocess
from functools import lru_cache
from pathlib import Path
from typing import TYPE_CHECKING, Any, Callable, cast
import numpy as np
import torch
import torch.nn.functional as F
if TYPE_CHECKING:
# numpy.typing needs numpy>=1.20; requirements.txt allows down to 1.16, so this must stay out
# of the runtime import path.
import numpy.typing as npt
FrameArray = npt.NDArray[np.float32]
else:
FrameArray = np.ndarray
RIFE_ZIP_URL = "https://huggingface.co/thornmaze/RIFE/resolve/main/RIFEv4.26_0921.zip"
RIFE_ZIP_PATH = Path("RIFEv4.26_0921.zip")
RIFE_DIR = Path("train_log")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
ProgressCallback = Callable[[int, int], None]
def ensure_weights_downloaded() -> None:
"""Download and unpack the RIFE weights if not already present. CPU/network only — no CUDA
needed — so this is safe (and intended) to call outside any @spaces.GPU allocation."""
if RIFE_DIR.exists():
return
subprocess.run(["wget", "-q", RIFE_ZIP_URL, "-O", str(RIFE_ZIP_PATH)], check=True)
subprocess.run(["unzip", "-o", str(RIFE_ZIP_PATH)], check=True)
@lru_cache(maxsize=1)
def _load_model() -> Any:
ensure_weights_downloaded()
from train_log.RIFE_HDv3 import Model # only importable after the zip is unpacked
model = Model()
model.load_model(str(RIFE_DIR), -1)
model.eval()
model.device()
model.flownet = model.flownet.half()
return model
@torch.no_grad()
def interpolate_frames(
frames_np: FrameArray | list[FrameArray],
multiplier: int,
scale: float = 1.0,
progress_callback: ProgressCallback | None = None,
as_tensor: bool = False,
) -> list[FrameArray] | list[torch.Tensor]:
"""Interpolate frames with RIFE to `multiplier`x the input frame rate.
Args:
frames_np: (T, H, W, C) array or list of (H, W, C) arrays, float32 in [0, 1].
multiplier: 2, 4, or 8. Values < 2 return the frames unchanged (as a list).
progress_callback: progress_callback(done, total) fired after each input-frame gap.
as_tensor: if True, skip the GPU->CPU conversion and return (C, H, W) fp16 GPU
tensors instead of (H, W, C) float32 numpy arrays — for callers (e.g. upscale)
that consume the result on GPU right after, so frames never round-trip to CPU
in between.
Returns:
List of (H, W, C) float32 numpy arrays in [0, 1], or (if as_tensor) list of
(C, H, W) fp16 GPU tensors in [0, 1].
"""
if multiplier < 2:
return list(frames_np) if isinstance(frames_np, np.ndarray) else frames_np
model = _load_model()
total_frames = len(frames_np)
height, width, _ = frames_np[0].shape
n_interp = multiplier - 1
tile = max(128, int(128 / scale))
padded_h = ((height - 1) // tile + 1) * tile
padded_w = ((width - 1) // tile + 1) * tile
padding = (0, padded_w - width, 0, padded_h - height)
def to_tensor(frame_np: FrameArray) -> torch.Tensor:
t = torch.from_numpy(frame_np).to(device)
t = t.permute(2, 0, 1).unsqueeze(0)
return F.pad(t, padding).half()
def from_tensor_np(tensor: torch.Tensor) -> FrameArray:
t = tensor[0, :, :height, :width]
t = t.permute(1, 2, 0)
return t.float().cpu().numpy()
def from_tensor_gpu(tensor: torch.Tensor) -> torch.Tensor:
# .contiguous() detaches from the padded working tensor's storage (a plain crop is
# just a view into it) so the padded buffer isn't kept alive for the batch's lifetime.
return tensor[0, :, :height, :width].contiguous()
from_tensor: Callable[[torch.Tensor], FrameArray | torch.Tensor]
from_tensor = from_tensor_gpu if as_tensor else from_tensor_np
def make_inference(i0: torch.Tensor, i1: torch.Tensor, n: int) -> list[torch.Tensor]:
if model.version >= 3.9:
return [model.inference(i0, i1, (i + 1) / (n + 1), scale) for i in range(n)]
middle = model.inference(i0, i1, scale)
if n == 1:
return [middle]
first_half = make_inference(i0, middle, n // 2)
second_half = make_inference(middle, i1, n // 2)
if n % 2:
return [*first_half, middle, *second_half]
return [*first_half, *second_half]
output_frames: list[FrameArray | torch.Tensor] = []
next_tensor = to_tensor(frames_np[0])
total_steps = total_frames - 1
for i in range(total_steps):
current_tensor = next_tensor
output_frames.append(from_tensor(current_tensor))
next_tensor = to_tensor(frames_np[i + 1])
for mid in make_inference(current_tensor, next_tensor, n_interp):
output_frames.append(from_tensor(mid))
if progress_callback is not None:
progress_callback(i + 1, total_steps)
output_frames.append(from_tensor(next_tensor))
del current_tensor, next_tensor
torch.cuda.empty_cache()
if as_tensor:
return cast("list[torch.Tensor]", output_frames)
return cast("list[FrameArray]", output_frames)