From 9dd1d73f6d6de7bbd0978e15476eddc8985abde6 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Fri, 21 Aug 2026 16:21:12 +0200 Subject: [PATCH] video upscaling using spandrel Signed-off-by: Vladimir Mandic --- modules/upscaler.py | 6 +-- modules/upscaler_spandrel.py | 37 ++++++++++++---- modules/video_models/video_save.py | 5 ++- modules/video_models/video_ui.py | 11 ++--- modules/video_models/video_upscale.py | 62 +++++++++++++++++++++++++-- 5 files changed, 100 insertions(+), 21 deletions(-) diff --git a/modules/upscaler.py b/modules/upscaler.py index 4ed470ac9..5d1ef81f5 100644 --- a/modules/upscaler.py +++ b/modules/upscaler.py @@ -98,7 +98,7 @@ class Upscaler: return scalers @abstractmethod - def do_upscale(self, img: Image.Image | Tensor, selected_model: str): + def do_upscale(self, img: Image.Image | Tensor, selected_model: str, output_type='pil'): return img def upscale(self, img: Image.Image | Tensor, scale, selected_model: str | None = None): @@ -158,11 +158,11 @@ class UpscalerData: custom: bool = False name = None data_path = None - scale: int = 4 + scale: int = 1 scaler: Upscaler | None = None model: None - def __init__(self, name: str, path: str | None = None, upscaler: Upscaler | None = None, scale: int = 0, model=None): + def __init__(self, name: str, path: str | None = None, upscaler: Upscaler | None = None, scale: int = 1, model=None): self.name = name self.data_path = path self.local_data_path = path diff --git a/modules/upscaler_spandrel.py b/modules/upscaler_spandrel.py index 56a67c574..67dd795c4 100644 --- a/modules/upscaler_spandrel.py +++ b/modules/upscaler_spandrel.py @@ -13,8 +13,12 @@ MODELS = { "Spandrel 2x RealPLKSR AnimeSharpV2": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/2x-AnimeSharpV2_RPLKSR_Sharp.pth", "Spandrel 2x RealESRGAN Compact": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/RealESRGAN-2x-Compact.pth", "Spandrel 2x RealESRGAN UltraCompact": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/RealESRGAN-2x-UltraCompact.pth", - "Spandrel 2x RealSAFMN++": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/Real-SAFMN-x2.pth", - "Spandrel 4x RealSAFMN++": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/Real-SAFMN-x4-v2.pth", + "Spandrel 2x RealSAFMN++": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/Real-SAFMN++.pth", + "Spandrel 2x RealSAFMN": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/Real-SAFMN-x2.pth", + "Spandrel 4x RealSAFMN": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/Real-SAFMN-x4-v2.pth", + "Spandrel 2x SAFMN PureScale": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/2x_SAFMN_PureScale.pth", + "Spandrel 2x SAFMN PureScale Sharper": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/2x_SAFMN_PureScale_sharper.pth", + "Spandrel 4x SAFMN PureScale": "https://huggingface.co/vladmandic/sdnext-upscalers/resolve/main/4x_SAFMN_PureScale.pth", } class UpscalerSpandrel(Upscaler): @@ -38,7 +42,7 @@ class UpscalerSpandrel(Upscaler): s.name = k break - def process(self, img: Image.Image) -> Image.Image: + def process(self, img: Image.Image, output_type='pil', quiet=False) -> Image.Image: if isinstance(img, Image.Image): from modules.image import convert img = img.convert('RGB') @@ -47,17 +51,34 @@ class UpscalerSpandrel(Upscaler): from modules.image import convert tensor = convert.to_tensor(img).unsqueeze(0).to(devices.device) elif isinstance(img, torch.Tensor): - tensor = img.to(devices.device) + tensor = img.to(devices.device).float() else: log.error(f'Spandrel: unsupported input type={type(img)}') return img t0 = time.time() with devices.inference_context(): + if tensor.max() > 1.0: # tensor is in [0,255] range, convert to [0,1] + tensor = tensor.div_(255.0) + if tensor.min() < 0: # tensor is in [-1,1] range, convert to [0,1] + tensor = (tensor + 1.0) / 2.0 tensor = self.model(tensor) tensor = tensor.clamp(0, 1).squeeze(0).cpu() t1 = time.time() - upscaled = convert.to_pil(tensor) - log.debug(f'Upscale: name="{self.selected}" input={img.size} output={upscaled.size} time={t1 - t0:.3f}') + if output_type == 'pil': + upscaled = convert.to_pil(tensor) + if not quiet: + log.debug(f'Upscale: name="{self.selected}" input={img.size} type={output_type} output={upscaled.size} time={t1 - t0:.3f}') + elif output_type == 'nd': + upscaled = (255.0 * tensor).float().numpy().astype(np.uint8) + if not quiet: + log.debug(f'Upscale: name="{self.selected}" input={img.shape} type={output_type} output={upscaled.shape} time={t1 - t0:.3f}') + elif output_type == 'tensor': + upscaled = tensor + if not quiet: + log.debug(f'Upscale: name="{self.selected}" input={img.shape} type={output_type} output={list(upscaled.shape)} time={t1 - t0:.3f}') + else: + upscaled = img + log.error(f'Upscale: type={output_type} unsupported') return upscaled def load_model(self, path: str): @@ -71,11 +92,11 @@ class UpscalerSpandrel(Upscaler): self.model = spandrel.ModelLoader().load_from_file(model.local_data_path) self.model.to(devices.device).eval() - def do_upscale(self, img: Image.Image | torch.Tensor | np.ndarray, selected_model=None): + def do_upscale(self, img: Image.Image | torch.Tensor | np.ndarray, selected_model=None, output_type='pil', quiet=False): try: if (self.model is None) or (self.selected != selected_model): self.load_model(selected_model) - return self.process(img) + return self.process(img, output_type=output_type, quiet=quiet) except Exception as e: log.error(f'Spandrel: {e}') errors.display(e, "Spandrel") diff --git a/modules/video_models/video_save.py b/modules/video_models/video_save.py index 916daf979..0ff991d73 100644 --- a/modules/video_models/video_save.py +++ b/modules/video_models/video_save.py @@ -326,7 +326,7 @@ def save_video( if upscale_upscaler is not None and len(upscale_upscaler) > 0: t_upscale = time.time() - pixels = upscale_video(pixels, scale=upscale_scale, upscaler=upscale_upscaler) + pixels = upscale_video(pixels, scale=upscale_scale, upscaler_name=upscale_upscaler) timer.process.add('upscale', time.time()-t_upscale) t_save = time.time() @@ -334,6 +334,7 @@ def save_video( pixels = pixels.unsqueeze(0) n, _c, t, h, w = pixels.shape size = pixels.element_size() * pixels.numel() + t_min, t_max = pixels.min().item(), pixels.max().item() log.debug(f'Video: video={mp4_video} export={mp4_frames} safetensors={mp4_sf} interpolate={mp4_interpolate}') if hasattr(audio, 'shape'): audio_txt = f'audio={audio.shape} aac={aac_sample_rate}' if audio is not None else 'no audio' @@ -341,7 +342,7 @@ def save_video( audio_txt = f'audio={audio.get("format", None)} packets={len(audio.get("frames", []))} ' else: audio_txt = None - log.debug(f'Video: encode={t} raw={size} latent={pixels.shape} {audio_txt} fps={mp4_fps} codec={mp4_codec} ext={mp4_ext} options="{mp4_opt}"') + log.debug(f'Video: encode={t} tensor={pixels.shape} min={t_min} max={t_max} bytes={size} {audio_txt} fps={mp4_fps} codec={mp4_codec} ext={mp4_ext} options="{mp4_opt}"') try: preparejob = shared.state.begin('Prepare video') if stream is not None: diff --git a/modules/video_models/video_ui.py b/modules/video_models/video_ui.py index 87c3d39e6..74af572aa 100644 --- a/modules/video_models/video_ui.py +++ b/modules/video_models/video_ui.py @@ -1,4 +1,5 @@ import os +import inspect import gradio as gr from modules import sd_models, ui_common, ui_sections, ui_symbols, call_queue from modules.logger import log @@ -56,11 +57,11 @@ def model_load(engine, model): def refresh_upscalers(): - exclude = ['latent', 'interpolation', 'vips'] - from modules import modelloader - upscalers = modelloader.load_upscalers() - upscalers = [x for x in upscalers if all(e not in x.lower() for e in exclude)] - return upscalers + from modules import shared, modelloader + modelloader.load_upscalers() # refresh + upscalers = [u for u in shared.sd_upscalers if 'output_type' in inspect.signature(u.scaler.do_upscale).parameters.keys()] + upscaler_names = ['None'] + [u.name for u in upscalers] + return upscaler_names def create_ui_outputs(): diff --git a/modules/video_models/video_upscale.py b/modules/video_models/video_upscale.py index 5417677fc..2be09c5c2 100644 --- a/modules/video_models/video_upscale.py +++ b/modules/video_models/video_upscale.py @@ -1,7 +1,63 @@ +import time +import inspect import torch from modules.logger import log +from modules import shared, upscaler -def upscale_video(pixels: torch.Tensor, scale: float = 1.0, upscaler: str = ""): - log.debug(f'Upscale video: scale={scale} upscaler="{upscaler}" shape={list(pixels.shape)} TODO') - return pixels +def load_upscaler(upscaler_name: str) -> upscaler.UpscalerData | None: + upscalers = [x for x in shared.sd_upscalers if x.name.lower().replace('-', ' ') == upscaler_name.lower().replace('-', ' ')] + # use inspect to check if upscaler.scaler method has output_type param, if not, then it is an old upscaler and we should not use it for video + upscalers = [u for u in upscalers if 'output_type' in inspect.signature(u.scaler.do_upscale).parameters.keys()] + if len(upscalers) == 0: # do force-refresh before failing + from modules.modelloader import load_upscalers + load_upscalers() + upscalers = [x for x in shared.sd_upscalers if x.name.lower().replace('-', ' ') == upscaler_name.lower().replace('-', ' ')] + upscalers = [u for u in upscalers if 'output_type' in inspect.signature(u.scaler.do_upscale).parameters.keys()] + if len(upscalers) > 0: + return upscalers[0] + else: + log.warning(f'Upscaler: invalid="{upscaler_name}"') + log.debug(f"Upscaler: available={[u.name for u in shared.sd_upscalers]}") + return None + + +def upscale_video(pixels: torch.Tensor, scale: float = 1.0, upscaler_name: str = ""): + if upscaler_name is None or upscaler_name == "" or upscaler_name.lower() == "none": + return pixels + model = load_upscaler(upscaler_name) + if model is None: + log.warning(f'Video upscale: upscaler="{upscaler_name}" not found') + return pixels + log.debug(f'Video upscale: scale={scale} upscaler="{upscaler_name}" cls={model.scaler.__class__.__name__} shape={list(pixels.shape)}') + # pixels: BCFHW [1, 3, 34, 480, 640] + if pixels.ndim == 5: + frames = pixels + elif pixels.ndim == 4: + frames = pixels.unsqueeze(0) + else: + log.warning(f'Video upscale: shape={list(pixels.shape)} unrecognized') + return pixels + if pixels.shape[1] != 3: + log.warning(f'Video upscale: shape={list(pixels.shape)} unrecognized') + return pixels + outputs = [] + t0 = time.time() + for idx in range(frames.shape[2]): + frame = frames[:, :, idx, :, :] # BCHW + w = int(frame.shape[-1] * scale) + h = int(frame.shape[-2] * scale) + # upscale + frame = model.scaler.do_upscale(frame, model.name, output_type='tensor', quiet=True) + frame = frame * 2.0 - 1.0 # upscaler returns 0:1, need -1:1 for video + if frame.ndim == 3: + frame = frame.unsqueeze(0) + # interpolate to exact size + if frame.shape[-1] != w or frame.shape[-2] != h: + frame = torch.nn.functional.interpolate(frame, size=(h, w), mode='lanczos', align_corners=False, antialias=True) + outputs.append(frame) + outputs = torch.stack(outputs, dim=2) + t1 = time.time() + frames = outputs.shape[2] + log.debug(f'Video upscale: frames={frames} width={outputs.shape[4]} height={outputs.shape[3]} fps={frames / (t1 - t0):.3f} time={t1 - t0:.3f}') + return outputs