"""`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