File size: 12,469 Bytes
3546667 819311d 55d9178 819311d 3546667 7123409 819311d 678f3ac 819311d c8ec0f7 819311d 678f3ac 819311d c8ec0f7 819311d 2ba97e4 819311d 60ca9a8 2ba97e4 60ca9a8 3546667 819311d 678f3ac 819311d 7123409 819311d 7123409 819311d a5e5374 c8ec0f7 819311d 7123409 819311d 3546667 d7bfd39 819311d 678f3ac 819311d 60ca9a8 ea4ada4 60ca9a8 678f3ac 60ca9a8 819311d 678f3ac 60ca9a8 ea4ada4 819311d c8ec0f7 fed6bc3 c8ec0f7 fed6bc3 c8ec0f7 fed6bc3 c8ec0f7 60ca9a8 819311d c8ec0f7 819311d 3546667 819311d fed6bc3 55d9178 60ca9a8 819311d c8ec0f7 678f3ac 60ca9a8 fed6bc3 c8ec0f7 60ca9a8 fed6bc3 60ca9a8 ea4ada4 2ba97e4 ea4ada4 678f3ac 2ba97e4 b730447 ea4ada4 b730447 60ca9a8 ea4ada4 b730447 60ca9a8 678f3ac 819311d | 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 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 | """`4xLSDIRCompact` tiled super-resolution, GPU (fp16) with a CPU (fp32) fallback.
Inference structure (pre-pad, tile, merge, crop) is adapted from xinntao/Real-ESRGAN's `RealESRGANer`
(BSD-3-Clause License, see `LICENSES/REAL-ESRGAN-LICENSE`, `realesrgan/utils.py`), trimmed to the single code path
this Space needs: one fixed compact model, RGB frames in (numpy float32 `[0,1]`, `(H,W,C)` β the same
representation `_apply_interpolation` already produces, so no PIL round trip is needed at the boundary) and
`PIL.Image` out (no alpha channel, no 16-bit, no cv2/BGR round-trip, no `dni` model blending, no CLI
prefetch/IO-queue threads β none of those apply to frames coming straight out of the Wan 2.2 decode step).
The weights are `4xLSDIRCompact` (Philip Hofmann, CC BY 4.0, see `LICENSES/4xLSDIRCompact-LICENSE`) rather than
Real-ESRGAN's own `realesr-general-x4v3`: same `SRVGGNetCompact` architecture (just `num_conv=16` instead of 32,
so if anything it's cheaper per tile), but trained on clean LSDIR photos with no synthetic noise/blur/JPEG
augmentation. `realesr-general-x4v3` bakes in fairly aggressive denoising to cope with that augmentation, which
flattens real skin/fabric/hair texture into a smooth, CG-like look on frames that were never actually degraded
(e.g. straight out of a video decode step, as here) β `4xLSDIRCompact` doesn't have that bias.
Runs inside the caller's `@spaces.GPU` allocation (see app.py's `run_inference`), same as RIFE interpolation
(postprocess/interpolation.py) β moved off the Space's shared CPU per issue #11 to avoid CPU contention across
concurrent visitors. Falls back to CPU/fp32 automatically when no CUDA device is present (e.g. local runs).
The weight download 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 upscale request.
"""
from __future__ import annotations
import math
import os
from functools import lru_cache
from pathlib import Path
from typing import TYPE_CHECKING, Callable, cast
import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image
from safetensors.torch import load_file
from torch.hub import download_url_to_file, get_dir
from .profiling import UpscaleProfiler
from .srvgg_arch import SRVGGNetCompact
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
WEIGHTS_URL = "https://huggingface.co/Phips/4xLSDIRCompact/resolve/main/4xLSDIRCompact.safetensors"
SCALE = 4
NUM_CONV = 16
# Every distinct (tile-shape, batch-size) pair seen by the compiled model costs a first-call
# penalty β ~3.6s cold, ~0.5s even on a warm container β while repeat calls take ~15ms (issue
# #50). At 512, an 832x624 frame split into a 2x2 grid of 4 distinct edge-clipped tile shapes, so
# the tile grid itself was most of the upscale cost. 1024 covers the largest frame this Space
# generates (model/pipeline.py's MAX_DIM=832, +TILE_PAD) in a single tile, i.e. one shape per
# resolution. Memory headroom allows it: at TILE_SIZE=512 upscale added only ~1 GiB, with ~13 GiB
# still available after the diffusion pipeline on the 47 GiB dev-Space device (issue #51); an
# untiled batch of 4 roughly doubles the per-call patch area. Tiling still kicks in for anything
# larger.
TILE_SIZE = 1024
TILE_PAD = 10
# Frames are grouped into batches (of equal size) before tiling, and same-shaped tile patches
# β across both the frame batch and tile positions β are run through the model together instead
# of one at a time, since SRVGGNetCompact is small enough that a single frame leaves most of the
# GPU idle. Short batches (the tail of the video) are padded up to FRAME_BATCH_SIZE so they reuse
# the full batch's compiled shape instead of paying a first-call penalty for a new batch size.
FRAME_BATCH_SIZE = 4
MAX_TILE_BATCH = 16
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.float16 if device.type == "cuda" else torch.float32
# Set UPSCALE_PROFILE=1 on the dev Space to print per-model()-call timing and peak CUDA memory
# for each upscale_frames() invocation β see profiling.py. Off by default.
PROFILE = os.environ.get("UPSCALE_PROFILE", "") not in ("", "0")
ProgressCallback = Callable[[int, int], None]
def _weights_path() -> Path:
cache_dir = Path(get_dir()) / "checkpoints"
cache_dir.mkdir(parents=True, exist_ok=True)
return cache_dir / Path(WEIGHTS_URL).name
def ensure_weights_downloaded() -> None:
"""Download the 4xLSDIRCompact weights if not already cached. CPU/network only β no CUDA
needed β so this is safe (and intended) to call outside any @spaces.GPU allocation."""
weights_path = _weights_path()
if not weights_path.exists():
download_url_to_file(WEIGHTS_URL, str(weights_path))
@lru_cache(maxsize=1)
def _load_model() -> torch.nn.Module:
model: torch.nn.Module = SRVGGNetCompact( # type: ignore[no-untyped-call] # vendored, untyped (see srvgg_arch.py)
num_in_ch=3, num_out_ch=3, num_feat=64, num_conv=NUM_CONV, upscale=SCALE, act_type="prelu"
)
ensure_weights_downloaded()
state_dict = load_file(_weights_path())
model.load_state_dict(state_dict, strict=True)
model.to(device=device, dtype=dtype).eval()
# Deliberately not torch.compile'd. Measured on the dev Space (issue #50), 53 frames:
# compiled model time was ~11s cold / 2.7-5.4s warm, almost all of it a per-request
# first-call (compile/guard) penalty, vs. ~0.96s cold *and* warm for eager. Eager's steady
# compute is ~25% slower per pixel (~+0.2s/request at 832x624), far less than that penalty.
return model
def _pre_pad(tensor: torch.Tensor, pad: int) -> torch.Tensor:
return F.pad(tensor, (0, pad, 0, pad), mode="reflect")
@torch.no_grad()
def _tile_process(model: torch.nn.Module, img: torch.Tensor, profiler: UpscaleProfiler) -> torch.Tensor:
batch, channel, height, width = img.shape
output = img.new_zeros((batch, channel, height * SCALE, width * SCALE))
tiles_x = math.ceil(width / TILE_SIZE)
tiles_y = math.ceil(height / TILE_SIZE)
# Every tile job needs its own input patch (pad_*) and output placement (in_*/trim_*), but
# patches share the same H/W across the batch dim and across interior tile positions β only
# the last row/column of tiles is smaller, clipped against the image edge. Group jobs by
# patch shape so each model() call is a single batched forward pass over same-shaped patches
# (chunked to MAX_TILE_BATCH to bound memory) instead of one tile at a time.
jobs_by_shape: dict[tuple[int, int], list[tuple[int, int, int, int, int, int, int, int, int]]] = {}
with profiler.span("job_grouping"):
for b in range(batch):
for y in range(tiles_y):
for x in range(tiles_x):
in_x0, in_x1 = x * TILE_SIZE, min((x + 1) * TILE_SIZE, width)
in_y0, in_y1 = y * TILE_SIZE, min((y + 1) * TILE_SIZE, height)
pad_x0, pad_x1 = max(in_x0 - TILE_PAD, 0), min(in_x1 + TILE_PAD, width)
pad_y0, pad_y1 = max(in_y0 - TILE_PAD, 0), min(in_y1 + TILE_PAD, height)
shape = (pad_y1 - pad_y0, pad_x1 - pad_x0)
jobs_by_shape.setdefault(shape, []).append(
(b, in_x0, in_y0, in_x1, in_y1, pad_x0, pad_y0, pad_x1, pad_y1)
)
for shape, jobs in jobs_by_shape.items():
for chunk_start in range(0, len(jobs), MAX_TILE_BATCH):
chunk = jobs[chunk_start:chunk_start + MAX_TILE_BATCH]
patch = torch.cat(
[img[b:b + 1, :, pad_y0:pad_y1, pad_x0:pad_x1] for b, _, _, _, _, pad_x0, pad_y0, pad_x1, pad_y1 in chunk],
dim=0,
)
with profiler.timed(shape, patch.shape[0]):
tile_out = model(patch)
with profiler.span("tile_merge"):
for i, (b, in_x0, in_y0, in_x1, in_y1, pad_x0, pad_y0, pad_x1, pad_y1) in enumerate(chunk):
trim_x0, trim_y0 = (in_x0 - pad_x0) * SCALE, (in_y0 - pad_y0) * SCALE
trim_x1 = trim_x0 + (in_x1 - in_x0) * SCALE
trim_y1 = trim_y0 + (in_y1 - in_y0) * SCALE
output[b:b + 1, :, in_y0 * SCALE:in_y1 * SCALE, in_x0 * SCALE:in_x1 * SCALE] = (
tile_out[i:i + 1, :, trim_y0:trim_y1, trim_x0:trim_x1]
)
return output
def _frame_hw(frame: FrameArray | torch.Tensor) -> tuple[int, int]:
# numpy frames are (H, W, C); GPU tensor frames (from RIFE's as_tensor=True path) are (C, H, W).
h, w = frame.shape[-2:] if isinstance(frame, torch.Tensor) else frame.shape[:2]
return int(h), int(w)
def _frames_to_tensor(frames: list[FrameArray | torch.Tensor]) -> torch.Tensor:
if isinstance(frames[0], torch.Tensor):
# Already (C, H, W) fp16 GPU tensors (RIFE's as_tensor=True output) β just batch them,
# no host round trip.
return torch.stack(cast("list[torch.Tensor]", frames), dim=0).to(device=device, dtype=dtype)
array = np.stack(cast("list[FrameArray]", frames))
return torch.from_numpy(array).permute(0, 3, 1, 2).to(device=device, dtype=dtype)
def upscale_frames(
frames: list[FrameArray] | list[torch.Tensor], progress_callback: ProgressCallback | None = None
) -> list[Image.Image]:
"""Upscale every frame 4x with `4xLSDIRCompact`, tiled to bound peak memory.
frames are either numpy float32 arrays in `[0,1]`, shape `(H,W,C)`, or (C,H,W) fp16 GPU
tensors β the same representations `_apply_interpolation` can produce, so callers don't
need to round-trip through `PIL.Image`, or through the CPU at all, just to cross this
boundary.
Frames are processed in batches (of up to FRAME_BATCH_SIZE, split wherever frame size
changes) rather than one at a time, so the model sees a real batch dimension instead of
batch=1 on every call. progress_callback(done, total), if given, fires after each frame
finishes.
"""
if not frames:
return []
model = _load_model()
out_frames: list[Image.Image] = []
profiler = UpscaleProfiler(PROFILE, device)
i = 0
while i < len(frames):
h, w = _frame_hw(frames[i])
batch: list[FrameArray | torch.Tensor] = [frames[i]]
i += 1
while i < len(frames) and len(batch) < FRAME_BATCH_SIZE and _frame_hw(frames[i]) == (h, w):
batch.append(frames[i])
i += 1
with profiler.span("frames_to_tensor"):
tensor = _frames_to_tensor(batch)
with profiler.span("pre_pad"):
# Repeat the last frame to fill a short batch (see FRAME_BATCH_SIZE); the extra
# outputs are dropped below. Costs one ~15ms steady call vs. a first-call penalty.
if len(batch) < FRAME_BATCH_SIZE:
tensor = torch.cat([tensor, tensor[-1:].expand(FRAME_BATCH_SIZE - len(batch), -1, -1, -1)])
padded = _pre_pad(tensor, TILE_PAD)
upscaled = _tile_process(model, padded, profiler)
# drop batch padding, crop the pre-pad border (scaled) back off
upscaled = upscaled[: len(batch), :, : h * SCALE, : w * SCALE]
# Quantize to uint8 on the device before the host transfer: done per frame in numpy on the
# CPU, this was ~200ms/frame (~76% of non-model time, issue #61), and transferring float32
# moved 4x the bytes. Same arithmetic as before (fp32 *255, round-half-to-even like
# np.round), so output is unchanged.
with profiler.span("quantize"):
quantized = upscaled.clamp(0, 1).float().mul_(255.0).round_().to(torch.uint8)
with profiler.span("host_transfer"):
results = quantized.permute(0, 2, 3, 1).contiguous().cpu().numpy()
for result in results:
with profiler.span("pil_convert"):
out_frames.append(Image.fromarray(result))
if progress_callback is not None:
progress_callback(len(out_frames), len(frames))
profiler.report(len(frames), TILE_SIZE, FRAME_BATCH_SIZE, MAX_TILE_BATCH)
return out_frames
|