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)