flowframes / rife_core.py
Nekochu's picture
RIFE 4.9 frame interpolation on ZeroGPU
da72558
Raw History Blame Contribute Delete
50 kB
"""Streaming RIFE frame interpolation.
Two frames are resident at a time, never the clip. ffmpeg decodes raw RGB into a
pipe, each consecutive pair is interpolated, and the results go straight out to a
second ffmpeg encoding a segment. Memory is therefore O(1) in clip length, so a
5 minute 4K source costs time and nothing else.
Work is always cut into segments. On ZeroGPU that keeps each GPU call inside its
duration budget; on CPU the same seams act as checkpoints, so a cancel keeps the
minutes already spent. Segment length is derived from a measured per-frame rate
rather than a guess.
"""
from __future__ import annotations
import ctypes
import json
import os
import pathlib
import queue
import shutil
import statistics
import subprocess
import threading
import tempfile
import time
from dataclasses import dataclass
from typing import Callable
import numpy as np
HERE = pathlib.Path(__file__).parent
MODEL = HERE / "models" / "rife49.onnx"
# fp16 IS NOT USED, ON EITHER PATH, AND THE SELECTION CODE IS GONE ON PURPOSE.
# Converting to fp16 and running it on the CUDA EP was predicted at ~1.8x faster
# and came out 3.6x SLOWER, A/B'd inside a single GPU allocation on the same
# container: fp16 0.1096 s/frame vs fp32 0.0306 s/frame. keep_io_types inserts
# Cast nodes throughout, and GridSample/Resize appear to have no fp16 CUDA
# kernel, so the graph thrashes between precisions. The cosine gate PASSED at
# 0.99994549, so quality checks alone would have shipped the regression.
# On CPU it is a trap for a different reason: those cores have no half ALU, so
# every op promotes back to fp32 and you pay the conversion on top.
# This used to be a path guarded by a filename that deliberately did not exist,
# which is a disable that a later tidy-up "fixes" by accident. Re-adding it means
# re-running that A/B first.
ALIGN = 32 # RIFE downsamples 5 times; both dims must be a multiple of 32
FFMPEG = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE = shutil.which("ffprobe") or "ffprobe"
# Set False by container_cpus() when every cgroup read fails, meaning the count
# is the HOST's and only a guess at this container's allowance.
CPUS_TRUSTED = True
def cgroup_report() -> str:
"""What the cgroup actually says, for the startup line.
`container_cpus` returning 12 is ambiguous on its own: it means either the
container really may use 12, or every cgroup read failed and it fell back to
the host's count clamped by THREAD_CAP. Those are opposite situations and
only one of them is a bug, so print the raw values.
"""
bits = []
for path in ("/sys/fs/cgroup/cpu.max",
"/sys/fs/cgroup/cpu/cpu.cfs_quota_us",
"/sys/fs/cgroup/cpu/cpu.cfs_period_us",
os.environ.get("ZEROGPU_PROC_SELF_CGROUP_PATH", "")):
if not path:
continue
try:
bits.append(f"{path}={pathlib.Path(path).read_text().strip()[:40]}")
except Exception as e: # noqa: BLE001
bits.append(f"{path}=<{type(e).__name__}>")
bits.append(f"os.cpu_count()={os.cpu_count()}")
return " ".join(bits)
def container_cpus(default: int = 2) -> int:
"""CPUs this container may actually use. os.cpu_count() reports the host, so
on a Spaces box it over-reports badly and ORT would spawn threads that only
fight each other for the two vCPUs the container is allowed."""
for quota_p, period_p in (("/sys/fs/cgroup/cpu.max", None),
("/sys/fs/cgroup/cpu/cpu.cfs_quota_us",
"/sys/fs/cgroup/cpu/cpu.cfs_period_us")):
try:
raw = pathlib.Path(quota_p).read_text().split()
quota = raw[0]
period = raw[1] if period_p is None else pathlib.Path(period_p).read_text().strip()
if quota != "max" and int(quota) > 0:
return max(1, int(quota) // int(period))
except Exception:
continue
# Reached only when no cgroup file answered. That is the case this function
# exists to catch, so say so rather than silently returning the host count:
# ZeroGPU sets ZEROGPU_PROC_SELF_CGROUP_PATH, which hints the cgroup is not
# where this looks, and an over-reported count means ORT spawns threads that
# only fight each other.
global CPUS_TRUSTED
CPUS_TRUSTED = False
n = max(1, os.cpu_count() or default)
print(f"[rife] cgroup cpu quota unreadable, falling back to os.cpu_count()"
f"={n}; if this container is smaller than that, threads are "
f"oversubscribed", flush=True)
return n
# Measured sweep at 720p: this graph scales to about 12 threads and then falls
# apart, 32 threads being 4.3x slower than 12. A Space gets 2 vCPUs so the cgroup
# read decides it there; the cap only stops a fat host oversubscribing.
# Two different jobs, previously conflated into one cap of 12.
#
# When the cgroup answers, it states this container's real allowance and should
# be believed: ZeroGPU reports 16, and measured in one boot at 720p, 16 threads
# run 0.637 s/frame against 12 threads' 0.790 - 19% faster. The old cap of 12
# came from a sweep on a dev box that compared 12 with 32 and never tried 16.
#
# When no cgroup answers we are reading the HOST's core count, which on a fat
# machine over-reports badly, and ORT would spawn threads that only fight each
# other. That is the case the 12 was really protecting against, and it still
# does, because the same sweep showed 32 threads running 4.3x slower than 12.
_raw_cpus = container_cpus()
CPUS = _raw_cpus if CPUS_TRUSTED else min(_raw_cpus, 12)
# intra_op_num_threads governs ORT's pool only. OpenMP and the BLAS backends size
# themselves from the host core count, which on a Spaces box is dozens of cores
# the container will never get. Must be set before onnxruntime is imported.
for _v in ("OMP_NUM_THREADS", "MKL_NUM_THREADS", "OPENBLAS_NUM_THREADS"):
os.environ.setdefault(_v, str(CPUS))
@dataclass
class Clip:
width: int
height: int
fps: float
frames: int
duration: float
def probe(path: str) -> Clip:
# avg_frame_rate must be REQUESTED. Reading st.get("avg_frame_rate") without
# naming it here returned None on every call, so the whole VFR fix silently
# fell through to r_frame_rate and did nothing. Nothing raised.
#
# Bounded, too: this runs on every job before the first yield, with the lock
# held and no Stop button in existence yet, so a malformed container that
# sends ffprobe into a long analyzeduration scan would otherwise block the
# whole Space with nothing on screen.
try:
out = subprocess.run(
[FFPROBE, "-v", "error", "-select_streams", "v:0", "-show_entries",
"stream=width,height,r_frame_rate,avg_frame_rate,nb_frames"
":format=duration",
"-of", "json", path],
capture_output=True, text=True, check=True, timeout=30).stdout
except subprocess.TimeoutExpired:
# Not "Ffprobe ...": the tool's name is lowercase, so capitalising it to
# open the sentence looks like a typo. Reword instead.
raise RuntimeError("Could not read this file within 30s (ffprobe timed "
"out). It may be malformed; try re-muxing it, for "
"example ffmpeg -i in -c copy out.mp4") from None
except subprocess.CalledProcessError as e:
# An upload ffprobe rejects outright used to escape as CalledProcessError,
# whose str() is the full argv, so the user was shown an absolute path to
# ffprobe.EXE on someone else's machine and no hint that their FILE is the
# problem. Keep ffprobe's own last line, drop the command.
why = " ".join((e.stderr or "").split())[-160:]
raise RuntimeError("ffmpeg cannot open this file as a video"
+ (f" ({why})" if why else "")
+ ". Upload a normal video container such as mp4, "
"mkv or webm.") from None
d = json.loads(out)
if not d.get("streams"):
raise RuntimeError("This file has no video stream that ffmpeg will read "
"(is it audio only?)")
st = d["streams"][0]
# avg_frame_rate is the real cadence; r_frame_rate is the smallest tick that
# times every frame, which on a VFR source is the MAXIMUM rate. Segment start
# times are seconds computed from this while frame counts come from
# -frames:v, so a wrong fps slides the seek a frame earlier every ~1000
# frames and each segment then re-renders content the previous one covered.
num, _, den = (st.get("avg_frame_rate") or "0/0").partition("/")
fps = (float(num) / float(den)) if float(den or 0) and float(num) else 0.0
if fps <= 0:
num, _, den = st["r_frame_rate"].partition("/")
den_f = float(den) if den else 1.0
fps = float(num) / den_f if den_f else 0.0
if fps <= 0:
# Every consumer of fps divides by it, so failing here names the file
# instead of surfacing an opaque ZeroDivisionError inside the planner.
raise RuntimeError("Could not determine a frame rate for this file "
"(ffprobe reports 0/0). Re-mux it, for example "
"ffmpeg -i in -c copy out.mp4")
dur = float(d.get("format", {}).get("duration") or 0.0)
n = int(st.get("nb_frames") or 0) or int(round(fps * dur))
if n == 1:
raise RuntimeError("This file has a single frame, so there is no gap to "
"interpolate. Use a clip with at least two frames.")
if n < 2:
# nb_frames and format.duration are both absent on raw h264 and some
# mkv/webm. n then fell to 0, the planner made one segment of one gap,
# and a two frame file came back reported as a success. Counting is slow,
# so it is the fallback and not the default.
try:
cnt = subprocess.run(
[FFPROBE, "-v", "error", "-select_streams", "v:0", "-count_frames",
"-show_entries", "stream=nb_read_frames", "-of", "json", path],
capture_output=True, text=True, timeout=120).stdout
except subprocess.TimeoutExpired:
cnt = ""
try:
n = int(json.loads(cnt)["streams"][0]["nb_read_frames"])
except Exception: # noqa: BLE001
n = 0
if n < 2:
raise RuntimeError(
f"Could not determine a frame count for this file "
f"(container reports {n} frames). Re-mux it, for example "
f"ffmpeg -i in -c copy out.mp4")
if dur <= 0 and fps > 0:
dur = n / fps
if dur > 0 and n > 0 and abs(n / dur - fps) / fps > 0.02:
print(f"[rife] {path}: declared {fps:.3f}fps but {n} frames over "
f"{dur:.2f}s averages {n / dur:.3f}fps. Variable rate source: the "
f"output is stamped on a uniform grid, so it is retimed.",
flush=True)
return Clip(int(st["width"]), int(st["height"]), fps, n, dur)
# ---------------------------------------------------------------- backends
class Backend:
"""Common shape: infer(img0, img1, t) on NCHW float32 in 0..1."""
name = "none"
def infer(self, a: np.ndarray, b: np.ndarray, t: float) -> np.ndarray: # pragma: no cover
raise NotImplementedError
def preload_cuda_libs() -> int:
"""dlopen the CUDA runtime libraries so ONNX Runtime can find them.
ORT reports "Require cuDNN 9.* and CUDA 12.*" whenever any of them is missing
from the loader path, and then silently runs on CPU. The libraries do exist in
the image (torch ships them under site-packages/nvidia) but nothing puts that
directory on the path, and LD_LIBRARY_PATH cannot help because glibc reads it
once at process start. dlopening them RTLD_GLOBAL into this process does.
cuDNN 9 is split across several sub-libraries and libcudnn.so.9 depends on
them, so ordering matters: dependencies first, aggregate last.
"""
import ctypes, glob, sysconfig
site = sysconfig.get_paths().get("purelib", "")
root = os.path.join(site, "nvidia")
found = sorted(glob.glob(os.path.join(root, "*", "lib", "*.so*")))
print(f"[rife] nvidia .so files present: {len(found)}", flush=True)
if found:
names = sorted({os.path.basename(f).split(".so")[0] for f in found})
print(f"[rife] libs: {', '.join(names)}", flush=True)
# Dependencies before the aggregates that link them.
order = ["cudart", "cublasLt", "cublas", "cufft", "curand", "cusparse",
"cusolver", "nvrtc", "nvJitLink",
"cudnn_graph", "cudnn_engines_precompiled", "cudnn_engines_runtime_compiled",
"cudnn_heuristic", "cudnn_ops", "cudnn_adv", "cudnn_cnn", "cudnn"]
loaded, failed = [], []
for stem in order:
for so in sorted(glob.glob(os.path.join(root, "*", "lib", f"lib{stem}.so*"))):
try:
ctypes.CDLL(so, mode=ctypes.RTLD_GLOBAL)
loaded.append(os.path.basename(so))
except OSError as e:
failed.append(f"{os.path.basename(so)}: {str(e)[:60]}")
print(f"[rife] preloaded {len(loaded)} cuda libs, {len(failed)} failed", flush=True)
for f in failed[:5]:
print(f"[rife] FAILED {f}", flush=True)
return len(loaded)
def report_ffmpeg_hw() -> None:
"""What this image's ffmpeg can offload, printed once at startup.
The pip nvidia listing above cannot answer this: NVENC and NVDEC come from
the driver (libnvidia-encode, libnvcuvid) plus ffmpeg's own build flags, not
from site-packages. Bounded and swallowed, because a diagnostic must never be
able to stop the Space from starting.
"""
try:
enc = subprocess.run([FFMPEG, "-hide_banner", "-encoders"],
capture_output=True, text=True, timeout=20).stdout
dec = subprocess.run([FFMPEG, "-hide_banner", "-decoders"],
capture_output=True, text=True, timeout=20).stdout
acc = subprocess.run([FFMPEG, "-hide_banner", "-hwaccels"],
capture_output=True, text=True, timeout=20).stdout
except Exception as e: # noqa: BLE001
print(f"[rife] ffmpeg hw probe failed ({type(e).__name__})", flush=True)
return
nvenc = sorted({ln.split()[1] for ln in enc.splitlines()
if "nvenc" in ln and len(ln.split()) > 1})
nvdec = sorted({ln.split()[1] for ln in dec.splitlines()
if ("cuvid" in ln or "_qsv" in ln) and len(ln.split()) > 1})
hw = [w for w in acc.split() if w not in ("Hardware", "acceleration",
"methods:")]
print(f"[rife] ffmpeg hwaccels: {hw or 'none'}", flush=True)
print(f"[rife] ffmpeg nvenc encoders: {nvenc or 'none'}", flush=True)
print(f"[rife] ffmpeg hw decoders: {nvdec or 'none'}", flush=True)
def have_lib(soname: str) -> bool:
"""Can the dynamic loader find this shared object right now?
ctypes.CDLL does the real search - loader cache, LD_LIBRARY_PATH as it was at
process start, and anything already dlopened RTLD_GLOBAL - which is exactly
what ORT will do a moment later.
"""
try:
ctypes.CDLL(soname)
return True
except OSError:
return False
class OnnxBackend(Backend):
"""One model, best available execution provider.
Providers are tried one at a time. ORT's own provider list is all or nothing:
if the first entry fails to load it falls back to CPU alone and discards the
rest, so asking for TensorRT and CUDA together means a missing TensorRT takes
CUDA down with it. Whatever is finally selected is reported rather than
assumed, because a Space that quietly runs on CPU looks exactly like one that
does not.
fp32 on both paths; see the fp16 note at the top of this file.
"""
def __init__(self, threads: int = CPUS, prefer_gpu: bool = False,
trt_cache: str | None = None):
import onnxruntime as ort
o = ort.SessionOptions()
o.intra_op_num_threads = threads
o.inter_op_num_threads = 1
o.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
# The other three session knobs were measured at 736x1280, best of 12,
# and every alternative to the defaults lost: ORT_PARALLEL +61%,
# enable_mem_pattern=False +59%, enable_cpu_mem_arena=False +78%.
# ORT_PARALLEL losing badly is the useful one - it says this graph is a
# deep sequential chain with nothing to overlap, so an inter-op pool only
# buys synchronisation. Same story as the spin-wait measurement below.
# Spin-wait stays ON. Disabling it looks right on a 2 vCPU box (one core
# burned between ops) and measured the opposite: 1.171s vs 2.797s at two
# threads, 58% worse. RIFE is many small ops and the context-switch
# latency between them costs far more than the spin.
chain: list = []
if prefer_gpu:
preload_cuda_libs()
# No importlib.reload here. onnxruntime is a C extension and the CPU
# session built at import already holds bindings into it; reloading
# leaves that session pointing at a stale module. ORT resolves its
# provider shared objects lazily when a session is created, so the
# RTLD_GLOBAL preload above is already visible to it.
avail = set(ort.get_available_providers())
# Registered is not the same as usable. ORT ships the TensorRT EP
# inside the wheel, so it is always "available", but it dlopens
# libnvinfer.so.10 - TensorRT's own runtime, which this image does
# not have. Attempting it anyway cost a failed session construction
# on every fork, and ZeroGPU forks per GPU call. Look for the library
# first: still free if the image ever ships it, free if it does not.
if "TensorrtExecutionProvider" in avail and not have_lib("libnvinfer.so.10"):
print("[rife] TensorRT EP registered but libnvinfer.so.10 is "
"absent, skipping it", flush=True)
avail.discard("TensorrtExecutionProvider")
if "TensorrtExecutionProvider" in avail:
trt = {"trt_fp16_enable": True, "trt_engine_cache_enable": True}
if trt_cache:
os.makedirs(trt_cache, exist_ok=True)
trt["trt_engine_cache_path"] = trt_cache
chain.append([("TensorrtExecutionProvider", trt),
("CUDAExecutionProvider", {})])
if "CUDAExecutionProvider" in avail:
chain.append([("CUDAExecutionProvider", {})])
print(f"[rife] ORT sees providers: {sorted(avail)}", flush=True)
chain.append([("CPUExecutionProvider", {})])
model = str(MODEL)
self.s = None
for attempt in chain:
want = attempt[0][0]
try:
use = model if want != "CPUExecutionProvider" else str(MODEL)
sess = ort.InferenceSession(use, o, providers=attempt)
except Exception as e: # noqa: BLE001
print(f"[rife] {want} raised: {str(e)[:110]}", flush=True)
continue
# ORT does not raise when a provider fails to initialise. It prints
# "EP Error", quietly falls back to CPU inside the constructor and
# hands back a perfectly valid session. Constructing without error is
# therefore no evidence the provider was used: ask what it picked.
got = sess.get_providers()[0]
if got != want and want != "CPUExecutionProvider":
print(f"[rife] asked {want}, got {got}, trying next", flush=True)
del sess
continue
self.s = sess
break
if self.s is None:
raise RuntimeError("No usable execution provider")
self.provider = self.s.get_providers()[0]
self.name = {
"TensorrtExecutionProvider": "GPU TensorRT",
"CUDAExecutionProvider": "GPU CUDA",
"CPUExecutionProvider": f"CPU ONNX fp32 x{threads}",
}.get(self.provider, self.provider)
self.on_gpu = self.provider != "CPUExecutionProvider"
# Names are converter artifacts. Interrogate, never hardcode.
self.inames = [i.name for i in self.s.get_inputs()]
def infer(self, a, b, t):
return self.s.run(None, {self.inames[0]: a, self.inames[1]: b,
self.inames[2]: np.array([t], dtype=np.float32)})[0]
class Buffers:
"""Preallocated padded frame buffers for the streaming loop.
The naive path costs five full-size passes per frame: frombuffer, transpose,
astype, divide, ascontiguousarray, then np.pad again because 480, 720 and
1080 are none of them multiples of 32. At 1080p that is well over 100 MB of
memory traffic per decoded frame on a box with two vCPUs, and it is pure
serial overhead that the model rate measurement never sees.
Here the padded buffers are allocated once per segment and reused. Decoded
bytes are scaled straight into the interior view in a single pass, and only
the edge strips are re-replicated rather than the whole frame.
"""
def __init__(self, h: int, w: int):
self.h, self.w = h, w
self.ph = (h + ALIGN - 1) // ALIGN * ALIGN
self.pw = (w + ALIGN - 1) // ALIGN * ALIGN
# uint8 NHWC: exactly what the decoder emits and what the graph now
# takes, so a frame is copied once and never converted on this CPU.
self.a = np.zeros((1, self.ph, self.pw, 3), dtype=np.uint8)
self.b = np.zeros((1, self.ph, self.pw, 3), dtype=np.uint8)
self._u8 = np.empty((h, w, 3), dtype=np.uint8)
def load(self, dst: np.ndarray, raw: bytes) -> np.ndarray:
"""Copy decoded RGB straight into dst's interior, then fix the edges.
No scaling and no transpose any more: the graph takes uint8 NHWC and does
both on the device. This was 0.42 s of a 3.44 s segment at 852x480.
The decoder stays on packed rgb24 deliberately. Asking it for gbrp would
remove this transpose, but `yuv420p -> gbrp` and `yuv420p -> rgb24` are
different chroma upsampling paths in swscale and disagree by up to 154
on a colour-bar source, mean 1.87, with no dither or rounding flag
closing the gap. That is a change to what the user gets, for 0.74 s a
segment. The encoder half IS free, but it cannot be taken alone: the
source-frame pass-through writes decoder bytes straight to the encoder,
so the two pipes must agree on a format. See emit() for the full result.
"""
src = np.frombuffer(raw, dtype=np.uint8).reshape(self.h, self.w, 3)
dst[0, :self.h, :self.w, :] = src # one contiguous copy
if self.ph > self.h: # replicate bottom edge
dst[0, self.h:, :self.w, :] = dst[0, self.h - 1:self.h, :self.w, :]
if self.pw > self.w: # replicate right edge
dst[0, :, self.w:, :] = dst[0, :, self.w - 1:self.w, :]
return dst
def emit(self, x: np.ndarray) -> bytes:
"""Pack the model's uint8 NCHW output into the rgb24 the encoder wants.
The scale, round and clip used to happen here: three float passes over
2.76 M elements per frame, on two vCPUs, while the GPU finished the whole
frame in 21 ms. Measured on the Space that was 1.13-1.20 s of the 2.3 s
of non-model time in a 179-gap segment, the largest term after inference.
Mul/Add/Clip/Cast now live in the graph and run on the device, verified
bit identical against this function's old output on real model results.
What is left is the NCHW to HWC pack the encoder's rgb24 requires.
Planar (gbrp) on both pipes would remove this and the matching transpose
in load(), worth ~1.5 s of the 1.94 s of non-model time in a segment. It
was measured and declined: the ENCODER side is byte identical with gbrp
(max diff 0 feeding the same pixels both ways), but the DECODER is not -
`yuv420p -> gbrp` and `yuv420p -> rgb24` take different chroma upsampling
paths in swscale and disagree by up to 154 on a colour-bar source, mean
1.87, with no dither or rounding flag closing it. And the two pipes
cannot disagree: the source-frame pass-through writes decoder bytes
straight to the encoder, so a packed decoder with a planar encoder feeds
packed pixels into a planar pipe. Both formats or neither.
"""
# The model emits uint8 now. A float32 array here would truncate to
# near-black without raising, which is the kind of failure that reaches
# a user as a corrupt video and no error at all.
assert x.dtype == np.uint8, f"emit expects uint8 from the model, got {x.dtype}"
self._u8[...] = x[0, :, :self.h, :self.w].transpose(1, 2, 0)
return self._u8.tobytes()
def measure_rate(be: Backend, h: int, w: int, samples: int = 3) -> float:
"""Seconds per interpolated frame at this resolution, measured on this box.
Everything downstream is estimated from this rather than from a table, so a
slower vCPU or a faster GPU changes the plan instead of surprising the user.
"""
ph, pw = (h + ALIGN - 1) // ALIGN * ALIGN, (w + ALIGN - 1) // ALIGN * ALIGN
# NOT np.zeros. A freshly calloc'd array can sit on the shared zero page, and
# reusing one across samples reads it straight out of cache: measured at
# 640x360, reused zeros time at 0.1415 s against 0.2130 s for real decoded
# frames, 33% low. The single-sample version escaped that only by allocating
# fresh arrays each call. Sampling over one buffer would have walked into it
# and made every estimate optimistic. Random data times like real data
# (0.2130 vs 0.2130) and is reproducible.
rng = np.random.default_rng(0)
a = rng.integers(0, 256, (1, ph, pw, 3), dtype=np.uint8)
b = rng.integers(0, 256, (1, ph, pw, 3), dtype=np.uint8)
be.infer(a, b, 0.5) # warm
# Median of a few, not one. plan_gaps sizes every segment from this number,
# so its variance is a correctness property: a single sample at 640x360 was
# observed from 0.157 to 0.232 s across runs, a 1.5x spread, and a plan built
# on the low end makes segments 1.5x larger than intended - which on ZeroGPU
# is how a 90 s budget overruns a 120 s cap. At 1280x720 the spread is ~6%,
# so the error is worst where frames are small, but this is the foundation of
# every estimate the user is shown.
ts = []
for _ in range(max(1, samples)):
t0 = time.perf_counter()
be.infer(a, b, 0.5)
ts.append(time.perf_counter() - t0)
return statistics.median(ts)
def profile_ops(h: int, w: int, samples: int = 4) -> str:
"""One-line breakdown of where an inference goes, from ORT's own profiler.
Builds a SECOND session with profiling on rather than turning it on for the
live one: the profiler adds per-op bookkeeping, and measuring a session while
slowing it down is how you get numbers that describe the measurement.
"""
import glob
import onnxruntime as ort
try:
o = ort.SessionOptions()
o.enable_profiling = True
o.profile_file_prefix = os.path.join(tempfile.gettempdir(), "ffprof")
o.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
sess = ort.InferenceSession(str(MODEL), o, providers=[("CUDAExecutionProvider", {})])
if sess.get_providers()[0] != "CUDAExecutionProvider":
return "profile skipped: session did not land on CUDA"
rng = np.random.default_rng(0)
a = rng.integers(0, 256, (1, h, w, 3), dtype=np.uint8)
b = rng.integers(0, 256, (1, h, w, 3), dtype=np.uint8)
t = np.array([0.5], dtype=np.float32)
for _ in range(samples):
# The only hardcoded input names in the file; everything else binds
# off session.get_inputs(). They move with the graph.
sess.run(None, {"img0_u8": a, "img1_u8": b, "timestep": t})
path = sess.end_profiling()
del sess
with open(path, encoding="utf-8") as f:
ev = json.load(f)
tot, by, on_cpu, cpu_by = 0, {}, 0, {}
for e in ev:
if e.get("cat") != "Node" or "dur" not in e:
continue
args = e.get("args", {}) or {}
kind = args.get("op_name", "?")
by[kind] = by.get(kind, 0) + e["dur"]
# Which EP actually ran it. 38% memcpy has two possible causes and
# they need opposite fixes: inputs arriving on the host (fix by
# shrinking what crosses, or by IO binding), or an op with no CUDA
# kernel splitting the graph and forcing a round trip mid-inference
# (fix by replacing that op). Bucketing by provider tells them apart,
# and it is also how a uint8 input graph would be caught failing -
# ORT can place the leading Cast on the CPU EP, in which case the
# host casts to float32 and ships the same bytes as before, plus
# extra nodes.
if "CPU" in args.get("provider", ""):
on_cpu += e["dur"]
cpu_by[kind] = cpu_by.get(kind, 0) + e["dur"]
tot += e["dur"]
for f in glob.glob(o.profile_file_prefix + "*"):
try:
os.remove(f)
except OSError:
pass
if not tot:
return "profile produced no node events"
top = sorted(by.items(), key=lambda kv: -kv[1])[:6]
copies = sum(v for k, v in by.items() if "Memcpy" in k)
return ("ops " + ", ".join(f"{k} {v * 100 // tot}%" for k, v in top)
+ f" | memcpy {copies * 100 // tot}%"
+ f" | on CPU EP {on_cpu * 100 // tot}%"
+ (" [" + ", ".join(
f"{k} {v * 100 // tot}%" for k, v in
sorted(cpu_by.items(), key=lambda kv: -kv[1])[:5]) + "]"
if cpu_by else "")
+ f" | {tot / 1000 / samples:.1f} ms/run")
except Exception as e: # noqa: BLE001
return f"profile failed ({type(e).__name__}: {str(e)[:80]})"
def nvenc_works() -> bool:
"""Can this process actually ENCODE with NVENC right now?
`ffmpeg -encoders` listing h264_nvenc proves only that ffmpeg was built with
it. The runtime needs libnvidia-encode and a real device, and on ZeroGPU the
parent process has neither - a GPU exists only inside a @spaces.GPU fork. So
this has to run in the fork, and it has to encode something rather than ask.
"""
try:
r = subprocess.run(
[FFMPEG, "-v", "error", "-f", "lavfi", "-i", "nullsrc=s=256x256:d=0.1",
"-c:v", "h264_nvenc", "-f", "null", "-"],
capture_output=True, timeout=30)
return r.returncode == 0
except Exception: # noqa: BLE001
return False
# NVENC only pays above roughly 2 MP. Measured against x264 veryfast pinned to
# 2 threads, which is what the Space actually has, at MATCHED SSIM rather than
# matched knob number:
# 720p 0.9 MP 0.86x speed 1.02x size <- slower, so do not
# 1080p 2.1 MP 1.14x speed 0.92x size
# 1440p 3.7 MP 1.27x speed 0.85x size
# 2160p 8.3 MP 1.51x speed 0.82x size
# The threshold sits between the first two rows.
NVENC_MIN_PIXELS = 1_400_000
# cq is NOT crf. Passing crf straight through produced files 2.5-2.8x larger at
# the same number. Matching SSIM instead put the offset at +2: nvenc cq22 scores
# at or above x264 crf20 everywhere measured. Picking the offset by FILE SIZE
# would have said +6, and cq26 is a real quality drop (SSIM 0.941 against 0.967)
# that looks like a win on the size column alone.
NVENC_CQ_OFFSET = 2
def encoder_args(crf: int, hw: bool) -> list[str]:
"""x264, or NVENC when it is actually the faster of the two.
`-b:v 0` is required or vbr silently targets a bitrate and ignores -cq.
"""
if not hw:
return ["-c:v", "libx264", "-preset", "veryfast", "-crf", str(crf)]
cq = max(0, min(51, crf + NVENC_CQ_OFFSET))
return ["-c:v", "h264_nvenc", "-preset", "p4", "-tune", "hq",
"-rc", "vbr", "-cq", str(cq), "-b:v", "0"]
def interpolate_segment(src: str, dst: str, start: float, gaps: int, factor: int,
be: Backend, clip: Clip, crf: int = 18,
cancel: Callable[[], bool] | None = None,
on_frame: Callable[[int], None] | None = None,
write_tail: bool = True,
out_w: int = 0, out_h: int = 0,
hw_encode: bool = False) -> int:
"""Interpolate one time slice straight from decoder pipe to encoder pipe.
Segments overlap by exactly one source frame. A segment owns the GAPS that
start inside it, so it emits `prev` plus the generated frames for each gap and
leaves the trailing frame to whoever owns the next gap. Without that, every
seam would both duplicate a source frame and lose the frame belonging in the
gap across it; only the final segment writes its tail.
Two decoded frames are resident at a time, plus at most four more in the
encoder queue and the writer - about 100 MB at 4K rather than 50. Still O(1)
in clip length at any resolution, which is the property that matters; only
the constant grew when the writer thread arrived.
"""
w = out_w or clip.width
h = out_h or clip.height
fsize = w * h * 3
need = gaps + 1 # frames required to close them
# Scaling here rather than in a pre-pass: a whole-clip transcode before the
# first segment costs an extra encode AND decode over the entire source, is
# serial time before any progress shows, and adds a generation-loss hop.
vf = ([] if (w, h) == (clip.width, clip.height)
else ["-vf", f"scale={w}:{h}:flags=bicubic"])
# stderr goes to temp FILES, not pipes. A pipe nobody drains fills at ~64 KB,
# at which point ffmpeg blocks writing stderr, stops producing on stdout, and
# the unbounded dec.stdout.read() below waits forever. "-v error" hides
# warnings but not per-frame decode errors, and a slightly corrupt upload is
# routine input here, so that pipe filling is a question of when.
dec_err = tempfile.TemporaryFile()
enc_err = tempfile.TemporaryFile()
dec = subprocess.Popen(
[FFMPEG, "-v", "error", "-ss", f"{start:.6f}", "-i", src,
"-frames:v", str(need), *vf, "-f", "rawvideo", "-pix_fmt", "rgb24", "-"],
stdout=subprocess.PIPE, stderr=dec_err, bufsize=fsize * 2)
# -preset veryfast is MEASURED, not picked. Stage profiling puts the encoder
# at 1-2% of a CPU job and ~11% of a GPU one (ffmpeg does not get faster when
# the model does), so it is only worth anything on GPU, and only then if the
# trade is good. It is not. Benchmarked on real frames at fixed crf 20:
# 720p ultrafast -6.5% time +195% size superfast +20% time +100% size
# 2160p ultrafast -13% time +209% size superfast -14% time +102% size
# Three times the download to shave 13% off a stage that is a ninth of the
# job, about 1.5% overall. The slower presets cost 26-89% more time and did
# not come back smaller either. veryfast is the right point on this curve.
enc = subprocess.Popen(
[FFMPEG, "-v", "error", "-y", "-f", "rawvideo", "-pix_fmt", "rgb24",
"-s", f"{w}x{h}", "-r", f"{clip.fps * factor:.6f}", "-i", "-",
*encoder_args(crf, hw_encode and w * h >= NVENC_MIN_PIXELS),
"-pix_fmt", "yuv420p", dst],
stdin=subprocess.PIPE, stderr=enc_err)
written = 0
starved = False
# Stage accumulators. perf_counter calls are ~100 ns, so at a few per frame
# this is far below the millisecond effects being measured.
t_read = t_load = t_inf = t_emit = t_put = 0.0
bufs = Buffers(h, w)
prev_p, cur_p = bufs.a, bufs.b # ping-pong, no per-frame allocation
# The encoder gets its own thread so x264 overlaps inference instead of
# stalling it. A 4K frame is ~25 MB against a 64 KB default pipe buffer, so
# without this every write waits for ffmpeg to drain the last one.
wq: queue.Queue = queue.Queue(maxsize=3)
werr: list[BaseException] = []
def _writer():
while True:
item = wq.get()
if item is None:
return
try:
enc.stdin.write(item)
except BaseException as e: # noqa: BLE001
werr.append(e)
return # stop; put() notices below
wt = threading.Thread(target=_writer, daemon=True)
wt.start()
def put(b: bytes) -> None:
"""Hand a frame to the writer, refusing to block forever if it died."""
while True:
try:
wq.put(b, timeout=1.0)
return
except queue.Full:
if not wt.is_alive() or werr:
raise RuntimeError(
"the encoder stopped accepting frames"
+ (f": {werr[0]}" if werr else ""))
try:
raw_prev = dec.stdout.read(fsize)
if len(raw_prev) < fsize:
# Returning 0 from inside the try skipped the encoder check below and
# let concat splice straight over the gap: a jump cut whose only
# trace was one line in a 14-line ring buffer. A segment that decodes
# nothing is a bug in the bounds or a truncated source, so say so.
starved = True
else:
bufs.load(prev_p, raw_prev)
while not starved:
if cancel is not None and cancel():
break
_t = time.perf_counter()
raw_cur = dec.stdout.read(fsize)
t_read += time.perf_counter() - _t
if len(raw_cur) < fsize:
break
_t = time.perf_counter()
bufs.load(cur_p, raw_cur)
t_load += time.perf_counter() - _t
# A SOURCE frame goes out as the bytes that came in. It used to be
# scaled to float32 by load(), then scaled back by emit() - multiply,
# add, clip, transpose, copy - to reproduce what was already in hand.
# Verified byte-identical for all 256 levels: u/255 in float32, then
# clip(v*255 + 0.5), returns u exactly, so this changes no output.
# At factor 2 it is HALF of all emit work, at factor 4 a quarter.
_t = time.perf_counter()
put(raw_prev)
t_put += time.perf_counter() - _t
written += 1
for i in range(1, factor):
_t = time.perf_counter()
out = be.infer(prev_p, cur_p, i / factor)
t_inf += time.perf_counter() - _t
_t = time.perf_counter()
blob = bufs.emit(out)
t_emit += time.perf_counter() - _t
_t = time.perf_counter()
put(blob)
t_put += time.perf_counter() - _t
written += 1
if on_frame is not None:
on_frame(written)
prev_p, cur_p = cur_p, prev_p # swap, never copy
raw_prev = raw_cur
if write_tail and not starved:
put(raw_prev); written += 1
finally:
# Drain before closing stdin, or the frames still queued are discarded
# and the segment silently comes up short. The sentinel is put without a
# timeout guard on purpose: if the writer is gone the queue has room.
try:
wq.put(None, timeout=5.0)
except Exception as e: # noqa: BLE001
# Not swallowed. A dropped sentinel means the writer never learns to
# stop, so frames may still be queued and unwritten - which would
# otherwise return a short segment as a complete one.
werr.append(RuntimeError(f"could not signal the writer to finish: {e}"))
wt.join(timeout=120.0)
if wt.is_alive():
# join() returns identically whether the thread finished or the
# timeout expired, so this has to be asked explicitly.
werr.append(RuntimeError("the encoder writer did not finish in 120s"))
if written > 1:
tot = t_read + t_load + t_inf + t_emit + t_put
print(f"[rife] stages {w}x{h} {written}f: infer {t_inf:.2f}s "
f"read {t_read:.2f}s load {t_load:.2f}s emit {t_emit:.2f}s "
f"put {t_put:.2f}s | non-model {(tot - t_inf) * 1000 / written:.1f} "
f"ms/frame", flush=True)
for closer in (lambda: enc.stdin.close(), lambda: dec.stdout.close()):
try:
closer()
except Exception:
pass
rc = enc.wait() # safe now: stderr is a file, so it cannot block
try:
dec.wait(timeout=10)
except Exception:
dec.kill()
try:
# Order matters, all three ways.
# Starvation first: it is decided by the very FIRST read, before any
# cancel check, so it is a property of the bounds or the source and never
# of the cancellation. Swallowing it under a cancel would lose the one
# diagnostic built for bad bounds exactly when a user interrupts a job
# that was already going wrong. It also makes the encoder exit non-zero
# for want of input, so checking rc ahead of it reported "encoder failed"
# for what is really a decode fault.
if werr:
# Raised HERE, not in the finally. Raising from inside the finally
# jumped over enc.stdin.close(), dec.stdout.close(), enc.wait() and
# dec.wait(), so a failed encoder left the DECODER alive, blocked
# writing into a pipe with no reader, holding fds on the source and
# on its stderr temp file - one orphan pair per failed segment, and
# the shrink-and-retry path can try the same range four times.
#
# Ordered after starvation because a starved segment never calls
# put(), so werr cannot be set there, and before the cancel check
# because a broken write means the file is incomplete whatever the
# user pressed.
raise RuntimeError(f"encoder write failed: {werr[0]}")
if starved:
raise RuntimeError(
f"Decoder produced no frames at {start:.3f}s of a "
f"{clip.duration:.3f}s source. The file is probably truncated "
f"or corrupt: its header claims more than it contains. "
f"ffmpeg said: {tail(dec_err)}")
# Then the clean stop. A cancel landing between the loop's own check and
# the first frame leaves the encoder with nothing, and rc below would be
# reporting that cancellation rather than a fault.
if written == 0 and cancel is not None and cancel():
return 0
# A dead encoder used to read as success: written was returned
# regardless, so concat later spliced a truncated part.
if rc != 0:
raise RuntimeError(f"Encoder failed (exit {rc}). "
f"ffmpeg said: {tail(enc_err)}")
finally:
for f in (dec_err, enc_err):
try:
f.close()
except Exception: # noqa: BLE001
pass
return written
def tail(f, n: int = 400) -> str:
"""The last COMPLETE stderr line or two, for an error message.
Slicing the last n bytes lands mid-word: a truncated h264 file produced
"decoder produced no frames at 10.783s of a 20.000s source ze (2176 > 104)."
where "ze (2176 > 104)" is the tail of "Invalid NAL unit size (2176 > 104)".
ffmpeg also repeats the same complaint per frame, so the last couple of
distinct lines carry everything the earlier thousands do.
"""
try:
f.seek(0)
raw = f.read()[-4096:].decode("utf-8", "replace")
except Exception: # noqa: BLE001
return ""
lines, seen = [], set()
for ln in reversed([l.strip() for l in raw.splitlines() if l.strip()]):
if ln in seen:
continue
seen.add(ln)
lines.append(ln)
if len(lines) == 2 or sum(map(len, lines)) > n:
break
return " | ".join(reversed(lines))[:n]
# NOT corrected by a per-gap constant, and the measurements are why.
#
# plan_gaps charges `rate * (factor - 1)` per gap. A real segment also does a
# decode read, two buffer passes and `factor` x264 frames, so it under-charges.
# Measured, as measured/assumed on real segments:
#
# CPU 1280x720 f2 inference 0.9103 s/frame 1.028x -> 0.0258 s/gap missing
# CPU 1280x720 f4 inference 0.9103 s/frame 1.008x -> 0.0217 s/gap missing
# GPU 852x480 f2 inference 0.0579 s/frame 1.456x -> 0.0264 s/gap missing
# CPU 426x240 f2 inference 0.1023 s/frame 1.042x -> 0.0043 s/gap missing
#
# A single additive 0.026 s/gap fits the first three to within 0.4%. It is 6x too
# large for the fourth, where it would over-charge by 21% - and on ZeroGPU an
# over-charge is not free either, since 21% more segments is 21% more GPU
# allocations against a daily run budget. Four points that are neither constant
# nor pixel-linear (480p and 720p agree despite 2.25x the pixels) are not a
# model, so no constant is applied.
#
# What IS established: the under-charge is largest on a fast device, exactly
# where the 120 s ZeroGPU cap lives. That is handled by sizing the budget for it
# in app.py rather than by inventing a curve through four points. Every segment
# now logs its own ratio, so the data to build a real model accumulates from
# production instead of from one afternoon's harness.
def plan_gaps(clip: Clip, factor: int, rate_s: float, budget_s: float) -> int:
"""How many gaps one segment may own so it fits the compute budget.
Counted in gaps, not seconds. Second-based boundaries round against the frame
rate and the error accumulates across segments, which showed up as a handful
of extra frames on a seven segment run.
"""
per_gap = rate_s * max(1, factor - 1)
if per_gap <= 0:
return max(1, clip.frames - 1)
return max(1, int(budget_s / per_gap))
def segment_bounds(clip: Clip, gaps_per_seg: int,
start_frame: int = 0) -> list[tuple[float, int, bool]]:
"""(start_seconds, gaps_owned, is_last) for each segment, frame exact.
start_frame re-cuts only the tail of a clip, which is what a mid-job switch
from GPU to CPU needs: the segments already rendered keep their bounds and
the rest are re-sized for the slower device.
"""
total = max(1, clip.frames - 1)
out: list[tuple[float, int, bool]] = []
f = start_frame
while f < total:
n = min(gaps_per_seg, total - f)
# Half a frame early, deliberately. -ss lands on the first frame whose
# pts >= the target, and the target is a float division, so a target that
# rounds a hair above frame f's true pts silently starts the segment at
# f+1. Aiming between f-1 and f cannot reach f+1 and cannot fall short of
# f, which removes the rounding question entirely.
out.append((max(0.0, (f - 0.5) / clip.fps), n, f + n >= total))
f += n
return out
def concat(parts: list[str], dst: str) -> None:
lst = tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False)
for p in parts:
lst.write(f"file '{pathlib.Path(p).as_posix()}'\n")
lst.close()
try:
r = subprocess.run([FFMPEG, "-v", "error", "-y", "-f", "concat", "-safe",
"0", "-i", lst.name, "-c", "copy", dst],
capture_output=True)
if r.returncode != 0:
# check=True raises CalledProcessError, whose str() is the whole
# argv - an absolute path to ffmpeg and no reason. Same shape probe()
# was fixed to stop producing. Carry ffmpeg's own words instead.
why = " ".join((r.stderr or b"").decode("utf-8", "replace").split())
raise RuntimeError(
"could not join the finished segments"
+ (f": {why[-200:]}" if why else " and ffmpeg said nothing"))
finally:
# In a finally: a failing concat used to leak its list file, and concat
# fails on exactly the runs where something else already went wrong.
try:
os.unlink(lst.name)
except OSError:
pass
def mux_audio(video: str, original: str, dst: str) -> bool:
"""Carry the original audio across. Interpolation changes frame rate, not
duration, so the track still lines up."""
r = subprocess.run(
[FFMPEG, "-v", "error", "-y", "-i", video, "-i", original,
"-c", "copy", "-map", "0:v:0", "-map", "1:a:0?", "-shortest", dst],
capture_output=True)
return r.returncode == 0 and os.path.exists(dst) and os.path.getsize(dst) > 0