someone-in-the-world's picture
Claude Opus 5.5
Drop torch.compile from the upscale model; run eager (#50) (#70)
d7bfd39 unverified
Raw History Blame Contribute Delete
12.5 kB
"""`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