mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
video upscaling using spandrel
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
+3
-3
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user