File size: 5,939 Bytes
f5e3ffa 7123409 f5e3ffa c8ec0f7 f5e3ffa 2682589 f5e3ffa 7123409 f5e3ffa 2682589 7123409 f5e3ffa 2682589 f5e3ffa fed6bc3 c8ec0f7 f5e3ffa fed6bc3 f5e3ffa fed6bc3 f5e3ffa 2682589 f5e3ffa c8ec0f7 f5e3ffa fed6bc3 c8ec0f7 fed6bc3 f5e3ffa c8ec0f7 f5e3ffa c8ec0f7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 | """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)
|