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