This commit is contained in:
Haoming 2025-07-31 14:24:15 +08:00
parent 3d8212a718
commit 9d70067221
9 changed files with 331 additions and 494 deletions

View File

@ -25,7 +25,6 @@ import modules.textual_inversion.textual_inversion
from modules.shared import cmd_opts
from PIL import PngImagePlugin
from modules.realesrgan_model import get_realesrgan_models
from modules import devices
from typing import Any, Union, get_origin, get_args
import piexif
@ -226,7 +225,6 @@ class Api:
self.add_api_route("/sdapi/v1/sd-models", self.get_sd_models, methods=["GET"], response_model=list[models.SDModelItem])
self.add_api_route("/sdapi/v1/sd-modules", self.get_sd_vaes_and_text_encoders, methods=["GET"], response_model=list[models.SDModuleItem])
self.add_api_route("/sdapi/v1/face-restorers", self.get_face_restorers, methods=["GET"], response_model=list[models.FaceRestorerItem])
self.add_api_route("/sdapi/v1/realesrgan-models", self.get_realesrgan_models, methods=["GET"], response_model=list[models.RealesrganItem])
self.add_api_route("/sdapi/v1/prompt-styles", self.get_prompt_styles, methods=["GET"], response_model=list[models.PromptStyleItem])
self.add_api_route("/sdapi/v1/embeddings", self.get_embeddings, methods=["GET"], response_model=models.EmbeddingsResponse)
self.add_api_route("/sdapi/v1/refresh-embeddings", self.refresh_embeddings, methods=["POST"])
@ -713,9 +711,6 @@ class Api:
def get_face_restorers(self):
return [{"name":x.name(), "cmd_dir": getattr(x, "cmd_dir", None)} for x in shared.face_restorers]
def get_realesrgan_models(self):
return [{"name":x.name,"path":x.data_path, "scale":x.scale} for x in get_realesrgan_models(None)]
def get_prompt_styles(self):
styleList = []
for k in shared.prompt_styles.styles:

View File

@ -1,95 +0,0 @@
import os
from modules import modelloader, errors
from modules.shared import cmd_opts, opts
from modules.upscaler import Upscaler, UpscalerData
from modules.upscaler_utils import upscale_with_model
from modules_forge.utils import prepare_free_memory
class UpscalerDAT(Upscaler):
def __init__(self, user_path):
self.name = "DAT"
self.user_path = user_path
self.scalers = []
super().__init__()
for file in self.find_models(ext_filter=[".pt", ".pth", ".safetensors"]):
name = modelloader.friendly_name(file)
scaler_data = UpscalerData(name, file, upscaler=self, scale=None)
self.scalers.append(scaler_data)
for model in get_dat_models(self):
if model.name in opts.dat_enabled_models:
self.scalers.append(model)
def do_upscale(self, img, path):
prepare_free_memory()
try:
info = self.load_model(path)
except Exception:
errors.report(f"Unable to load DAT model {path}", exc_info=True)
return img
model_descriptor = modelloader.load_spandrel_model(
info.local_data_path,
device=self.device,
prefer_half=(not cmd_opts.no_half and not cmd_opts.upcast_sampling),
expected_architecture="DAT",
)
return upscale_with_model(
model_descriptor,
img,
tile_size=opts.DAT_tile,
tile_overlap=opts.DAT_tile_overlap,
)
def load_model(self, path):
for scaler in self.scalers:
if scaler.data_path == path:
if scaler.local_data_path.startswith("http"):
scaler.local_data_path = modelloader.load_file_from_url(
scaler.data_path,
model_dir=self.model_download_path,
hash_prefix=scaler.sha256,
)
if os.path.getsize(scaler.local_data_path) < 200:
# Re-download if the file is too small, probably an LFS pointer
scaler.local_data_path = modelloader.load_file_from_url(
scaler.data_path,
model_dir=self.model_download_path,
hash_prefix=scaler.sha256,
re_download=True,
)
if not os.path.exists(scaler.local_data_path):
raise FileNotFoundError(f"DAT data missing: {scaler.local_data_path}")
return scaler
raise ValueError(f"Unable to find model info: {path}")
def get_dat_models(scaler):
return [
UpscalerData(
name="DAT x2",
path="https://huggingface.co/w-e-w/DAT/resolve/main/experiments/pretrained_models/DAT/DAT_x2.pth",
scale=2,
upscaler=scaler,
sha256='7760aa96e4ee77e29d4f89c3a4486200042e019461fdb8aa286f49aa00b89b51',
),
UpscalerData(
name="DAT x3",
path="https://huggingface.co/w-e-w/DAT/resolve/main/experiments/pretrained_models/DAT/DAT_x3.pth",
scale=3,
upscaler=scaler,
sha256='581973e02c06f90d4eb90acf743ec9604f56f3c2c6f9e1e2c2b38ded1f80d197',
),
UpscalerData(
name="DAT x4",
path="https://huggingface.co/w-e-w/DAT/resolve/main/experiments/pretrained_models/DAT/DAT_x4.pth",
scale=4,
upscaler=scaler,
sha256='391a6ce69899dff5ea3214557e9d585608254579217169faf3d4c353caff049e',
),
]

View File

@ -1,64 +1,75 @@
from modules import modelloader, devices, errors
import re
from functools import lru_cache
from PIL import Image
from modules import devices, errors, modelloader
from modules.shared import opts
from modules.upscaler import Upscaler, UpscalerData
from modules.upscaler_utils import upscale_with_model
from modules_forge.utils import prepare_free_memory
PREFER_HALF = opts.prefer_fp16_upscalers
if PREFER_HALF:
print("[Upscalers] Prefer Half-Precision:", PREFER_HALF)
class UpscalerESRGAN(Upscaler):
def __init__(self, dirname):
def __init__(self, dirname: str):
self.user_path = dirname
self.model_path = dirname
super().__init__(True)
self.name = "ESRGAN"
self.model_url = "https://github.com/cszn/KAIR/releases/download/v1.0/ESRGAN.pth"
self.model_name = "ESRGAN_4x"
self.model_name = "ESRGAN"
self.scalers = []
self.user_path = dirname
super().__init__()
model_paths = self.find_models(ext_filter=[".pt", ".pth", ".safetensors"])
scalers = []
if len(model_paths) == 0:
scaler_data = UpscalerData(self.model_name, self.model_url, self, 4)
scalers.append(scaler_data)
self.scalers.append(scaler_data)
for file in model_paths:
if file.startswith("http"):
name = self.model_name
else:
name = modelloader.friendly_name(file)
scaler_data = UpscalerData(name, file, self, 4)
if match := re.search(r"(\d)[xX]|[xX](\d)", name):
scale = int(match.group(1) or match.group(2))
else:
scale = 4
scaler_data = UpscalerData(name, file, self, scale)
self.scalers.append(scaler_data)
def do_upscale(self, img, selected_model):
def do_upscale(self, img: Image.Image, selected_model: str):
prepare_free_memory()
try:
model = self.load_model(selected_model)
except Exception:
errors.report(f"Unable to load ESRGAN model {selected_model}", exc_info=True)
errors.report(f"Unable to load {selected_model}", exc_info=True)
return img
model.to(devices.device_esrgan)
return esrgan_upscale(model, img)
def load_model(self, path: str):
if path.startswith("http"):
# TODO: this doesn't use `path` at all?
filename = modelloader.load_file_from_url(
url=self.model_url,
model_dir=self.model_download_path,
file_name=f"{self.model_name}.pth",
)
else:
filename = path
return modelloader.load_spandrel_model(
filename,
device=('cpu' if devices.device_esrgan.type == 'mps' else None),
expected_architecture='ESRGAN',
return upscale_with_model(
model=model,
img=img,
tile_size=opts.ESRGAN_tile,
tile_overlap=opts.ESRGAN_tile_overlap,
)
@lru_cache(maxsize=4, typed=False)
def load_model(self, path: str):
if not path.startswith("http"):
filename = path
else:
filename = modelloader.load_file_from_url(
url=path,
model_dir=self.model_download_path,
file_name=path.rsplit("/", 1)[-1],
)
def esrgan_upscale(model, img):
return upscale_with_model(
model,
img,
tile_size=opts.ESRGAN_tile,
tile_overlap=opts.ESRGAN_tile_overlap,
)
model = modelloader.load_spandrel_model(filename, device="cpu", prefer_half=PREFER_HALF)
model.to(devices.device_esrgan)
return model

View File

@ -1,45 +0,0 @@
import os
import sys
from modules import modelloader, devices
from modules.shared import opts
from modules.upscaler import Upscaler, UpscalerData
from modules.upscaler_utils import upscale_with_model
from modules_forge.utils import prepare_free_memory
class UpscalerHAT(Upscaler):
def __init__(self, dirname):
self.name = "HAT"
self.scalers = []
self.user_path = dirname
super().__init__()
for file in self.find_models(ext_filter=[".pt", ".pth"]):
name = modelloader.friendly_name(file)
scale = 4 # TODO: scale might not be 4, but we can't know without loading the model
scaler_data = UpscalerData(name, file, upscaler=self, scale=scale)
self.scalers.append(scaler_data)
def do_upscale(self, img, selected_model):
prepare_free_memory()
try:
model = self.load_model(selected_model)
except Exception as e:
print(f"Unable to load HAT model {selected_model}: {e}", file=sys.stderr)
return img
model.to(devices.device_esrgan) # TODO: should probably be device_hat
return upscale_with_model(
model,
img,
tile_size=opts.ESRGAN_tile, # TODO: should probably be HAT_tile
tile_overlap=opts.ESRGAN_tile_overlap, # TODO: should probably be HAT_tile_overlap
)
def load_model(self, path: str):
if not os.path.isfile(path):
raise FileNotFoundError(f"Model file {path} not found")
return modelloader.load_spandrel_model(
path,
device=devices.device_esrgan, # TODO: should probably be device_hat
expected_architecture='HAT',
)

View File

@ -22,7 +22,8 @@ from modules import sd_samplers, shared, script_callbacks, errors, stealth_infot
from modules.paths_internal import roboto_ttf_file
from modules.shared import opts
LANCZOS = (Image.Resampling.LANCZOS if hasattr(Image, 'Resampling') else Image.LANCZOS)
LANCZOS = getattr(Image, "Resampling", Image).LANCZOS
NEAREST = getattr(Image, "Resampling", Image).NEAREST
def get_font(fontsize: int):

View File

@ -1,108 +0,0 @@
import os
from modules import modelloader, errors
from modules.shared import cmd_opts, opts
from modules.upscaler import Upscaler, UpscalerData
from modules.upscaler_utils import upscale_with_model
from modules_forge.utils import prepare_free_memory
class UpscalerRealESRGAN(Upscaler):
def __init__(self, path):
self.name = "RealESRGAN"
self.user_path = path
super().__init__()
self.enable = True
self.scalers = []
scalers = get_realesrgan_models(self)
local_model_paths = self.find_models(ext_filter=[".pth"])
for scaler in scalers:
if scaler.local_data_path.startswith("http"):
filename = modelloader.friendly_name(scaler.local_data_path)
local_model_candidates = [local_model for local_model in local_model_paths if local_model.endswith(f"{filename}.pth")]
if local_model_candidates:
scaler.local_data_path = local_model_candidates[0]
if scaler.name in opts.realesrgan_enabled_models:
self.scalers.append(scaler)
def do_upscale(self, img, path):
prepare_free_memory()
if not self.enable:
return img
try:
info = self.load_model(path)
except Exception:
errors.report(f"Unable to load RealESRGAN model {path}", exc_info=True)
return img
model_descriptor = modelloader.load_spandrel_model(
info.local_data_path,
device=self.device,
prefer_half=(not cmd_opts.no_half and not cmd_opts.upcast_sampling),
expected_architecture="ESRGAN", # "RealESRGAN" isn't a specific thing for Spandrel
)
return upscale_with_model(
model_descriptor,
img,
tile_size=opts.ESRGAN_tile,
tile_overlap=opts.ESRGAN_tile_overlap,
# TODO: `outscale`?
)
def load_model(self, path):
for scaler in self.scalers:
if scaler.data_path == path:
if scaler.local_data_path.startswith("http"):
scaler.local_data_path = modelloader.load_file_from_url(
scaler.data_path,
model_dir=self.model_download_path,
)
if not os.path.exists(scaler.local_data_path):
raise FileNotFoundError(f"RealESRGAN data missing: {scaler.local_data_path}")
return scaler
raise ValueError(f"Unable to find model info: {path}")
def get_realesrgan_models(scaler: UpscalerRealESRGAN):
return [
UpscalerData(
name="R-ESRGAN General 4xV3",
path="https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.5.0/realesr-general-x4v3.pth",
scale=4,
upscaler=scaler,
),
UpscalerData(
name="R-ESRGAN General WDN 4xV3",
path="https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.5.0/realesr-general-wdn-x4v3.pth",
scale=4,
upscaler=scaler,
),
UpscalerData(
name="R-ESRGAN AnimeVideo",
path="https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.5.0/realesr-animevideov3.pth",
scale=4,
upscaler=scaler,
),
UpscalerData(
name="R-ESRGAN 4x+",
path="https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth",
scale=4,
upscaler=scaler,
),
UpscalerData(
name="R-ESRGAN 4x+ Anime6B",
path="https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.2.4/RealESRGAN_x4plus_anime_6B.pth",
scale=4,
upscaler=scaler,
),
UpscalerData(
name="R-ESRGAN 2x+",
path="https://github.com/xinntao/Real-ESRGAN/releases/download/v0.2.1/RealESRGAN_x2plus.pth",
scale=2,
upscaler=scaler,
),
]

View File

@ -102,12 +102,10 @@ options_templates.update(options_section(('saving-to-dirs', "Saving to a directo
options_templates.update(options_section(('upscaling', "Upscaling", "postprocessing"), {
"ESRGAN_tile": OptionInfo(192, "Tile size for ESRGAN upscalers.", gr.Slider, {"minimum": 0, "maximum": 512, "step": 16}).info("0 = no tiling"),
"ESRGAN_tile_overlap": OptionInfo(8, "Tile overlap for ESRGAN upscalers.", gr.Slider, {"minimum": 0, "maximum": 48, "step": 1}).info("Low values = visible seam"),
"realesrgan_enabled_models": OptionInfo(["R-ESRGAN 4x+", "R-ESRGAN 4x+ Anime6B"], "Select which Real-ESRGAN models to show in the web UI.", gr.CheckboxGroup, lambda: {"choices": shared_items.realesrgan_models_names()}),
"dat_enabled_models": OptionInfo(["DAT x2", "DAT x3", "DAT x4"], "Select which DAT models to show in the web UI.", gr.CheckboxGroup, lambda: {"choices": shared_items.dat_models_names()}),
"DAT_tile": OptionInfo(192, "Tile size for DAT upscalers.", gr.Slider, {"minimum": 0, "maximum": 512, "step": 16}).info("0 = no tiling"),
"DAT_tile_overlap": OptionInfo(8, "Tile overlap for DAT upscalers.", gr.Slider, {"minimum": 0, "maximum": 48, "step": 1}).info("Low values = visible seam"),
"composite_tiles_on_gpu": OptionInfo(False, "Composite the Tiles on GPU").info("improve performance and resource utilization"),
"upscaler_for_img2img": OptionInfo(None, "Upscaler for img2img", gr.Dropdown, lambda: {"choices": [x.name for x in shared.sd_upscalers]}),
"set_scale_by_when_changing_upscaler": OptionInfo(False, "Automatically set the Scale by factor based on the name of the selected Upscaler."),
"prefer_fp16_upscalers": OptionInfo(False, "Prefer to load Upscaler in half precision").info("increase speed; reduce quality; will try <b>fp16</b>, then <b>bf16</b>, then fall back to <b>fp32</b> if not supported").needs_restart(),
}))
options_templates.update(options_section(('face-restoration', "Face restoration", "postprocessing"), {

View File

@ -1,14 +1,14 @@
import os
from abc import abstractmethod
import PIL
from PIL import Image
import modules.shared
from modules import modelloader, shared
from modules import devices, modelloader, shared
from modules.images import LANCZOS, NEAREST
from modules.shared import cmd_opts, models_path, opts
LANCZOS = (Image.Resampling.LANCZOS if hasattr(Image, 'Resampling') else Image.LANCZOS)
NEAREST = (Image.Resampling.NEAREST if hasattr(Image, 'Resampling') else Image.NEAREST)
# hardcode
UPSCALE_ITERATIONS = 4
class Upscaler:
@ -20,132 +20,123 @@ class Upscaler:
filter = None
model = None
user_path = None
scalers: list
scalers: list["UpscalerData"] = []
tile = True
def __init__(self, create_dirs=False):
self.mod_pad_h = None
self.tile_size = modules.shared.opts.ESRGAN_tile
self.tile_pad = modules.shared.opts.ESRGAN_tile_overlap
self.device = modules.shared.device
self.scale: int = 1
self.tile_size: int = opts.ESRGAN_tile
self.tile_pad: int = opts.ESRGAN_tile_overlap
self.device = devices.device_esrgan
self.half: bool = not cmd_opts.no_half
self.model_download_path: str = None
self.img = None
self.output = None
self.scale = 1
self.half = not modules.shared.cmd_opts.no_half
self.pre_pad = 0
self.mod_scale = None
self.model_download_path = None
if self.model_path is None and self.name:
self.model_path = os.path.join(shared.models_path, self.name)
if self.model_path and create_dirs:
self.model_path = os.path.join(models_path, self.name)
if create_dirs and self.model_path:
os.makedirs(self.model_path, exist_ok=True)
try:
import cv2 # noqa: F401
self.can_tile = True
except Exception:
pass
@abstractmethod
def do_upscale(self, img: PIL.Image, selected_model: str):
return img
def upscale(self, img: PIL.Image, scale, selected_model: str = None):
self.scale = scale
dest_w = int((img.width * scale) // 8 * 8)
dest_h = int((img.height * scale) // 8 * 8)
for i in range(3):
if img.width >= dest_w and img.height >= dest_h and (i > 0 or scale != 1):
break
if shared.state.interrupted:
break
shape = (img.width, img.height)
img = self.do_upscale(img, selected_model)
if shape == (img.width, img.height):
break
if img.width != dest_w or img.height != dest_h:
img = img.resize((int(dest_w), int(dest_h)), resample=LANCZOS)
return img
def do_upscale(self, img: Image.Image, selected_model: str):
raise NotImplementedError
@abstractmethod
def load_model(self, path: str):
pass
raise NotImplementedError
def find_models(self, ext_filter=None) -> list:
return modelloader.load_models(model_path=self.model_path, model_url=self.model_url, command_path=self.user_path, ext_filter=ext_filter)
def upscale(self, img: Image.Image, scale: int, selected_model: str = None):
self.scale = scale
dest_w: int = (img.width * scale) // 8 * 8
dest_h: int = (img.height * scale) // 8 * 8
def update_status(self, prompt):
print(f"\nextras: {prompt}", file=shared.progress_print_out)
for _ in range(UPSCALE_ITERATIONS):
if shared.state.interrupted:
break
img = self.do_upscale(img, selected_model)
if ((img.width >= dest_w) and (img.height >= dest_h)) or (int(scale) == 1):
break
if (img.width != dest_w) or (img.height != dest_h):
img = img.resize((int(dest_w), int(dest_h)), LANCZOS)
return img
def find_models(self, ext_filter=None) -> list[str]:
return modelloader.load_models(
model_path=self.model_path,
model_url=self.model_url,
command_path=self.user_path,
ext_filter=ext_filter,
)
class UpscalerData:
name = None
data_path = None
scale: int = 4
scaler: Upscaler = None
name: str
data_path: str
scaler: Upscaler
scale: int
model: None
def __init__(self, name: str, path: str, upscaler: Upscaler = None, scale: int = 4, model=None, sha256: str = None):
def __init__(
self,
name: str,
path: str,
upscaler: Upscaler = None,
scale: int = 4,
model=None,
):
self.name = name
self.data_path = path
self.local_data_path = path
self.scaler = upscaler
self.scale = scale
self.model = model
self.sha256 = sha256
def __repr__(self):
return f"<UpscalerData name={self.name} path={self.data_path} scale={self.scale}>"
class UpscalerNone(Upscaler):
name = "None"
scalers = []
def load_model(self, path):
pass
def do_upscale(self, img, selected_model=None):
return img
def __init__(self, dirname=None):
super().__init__(False)
self.name = "None"
self.scalers = [UpscalerData("None", None, self)]
def load_model(self, _):
return
def do_upscale(self, img: Image.Image, *args, **kwargs):
return img
class UpscalerLanczos(Upscaler):
scalers = []
def do_upscale(self, img, selected_model=None):
return img.resize((int(img.width * self.scale), int(img.height * self.scale)), resample=LANCZOS)
def load_model(self, _):
pass
def __init__(self, dirname=None):
super().__init__(False)
self.name = "Lanczos"
self.scalers = [UpscalerData("Lanczos", None, self)]
def load_model(self, _):
return
def do_upscale(self, img: Image.Image, *args, **kwargs):
return img.resize(
size=(int(img.width * self.scale), int(img.height * self.scale)),
resample=LANCZOS,
)
class UpscalerNearest(Upscaler):
scalers = []
def do_upscale(self, img, selected_model=None):
return img.resize((int(img.width * self.scale), int(img.height * self.scale)), resample=NEAREST)
def load_model(self, _):
pass
def __init__(self, dirname=None):
super().__init__(False)
self.name = "Nearest"
self.scalers = [UpscalerData("Nearest", None, self)]
def load_model(self, _):
return
def do_upscale(self, img: Image.Image, *args, **kwargs):
return img.resize(
size=(int(img.width * self.scale), int(img.height * self.scale)),
resample=NEAREST,
)

View File

@ -1,4 +1,5 @@
import logging
from functools import wraps
from typing import Callable
import numpy as np
@ -11,44 +12,214 @@ from modules import devices, images, shared, torch_utils
logger = logging.getLogger(__name__)
def try_patch_spandrel():
try:
from spandrel.architectures.__arch_helpers.block import RRDB, ResidualDenseBlock_5C
_orig_init: Callable = ResidualDenseBlock_5C.__init__
_orig_5c_forward: Callable = ResidualDenseBlock_5C.forward
_orig_forward: Callable = RRDB.forward
@wraps(_orig_init)
def RDB5C_init(self, *args, **kwargs):
_orig_init(self, *args, **kwargs)
self.nf, self.gc = kwargs.get("nf", 64), kwargs.get("gc", 32)
@wraps(_orig_5c_forward)
def RDB5C_forward(self, x: torch.Tensor):
B, _, H, W = x.shape
nf, gc = self.nf, self.gc
buf = torch.empty((B, nf + 4 * gc, H, W), dtype=x.dtype, device=x.device)
buf[:, :nf].copy_(x)
x1 = self.conv1(x)
buf[:, nf : nf + gc].copy_(x1)
x2 = self.conv2(buf[:, : nf + gc])
if self.conv1x1:
x2.add_(self.conv1x1(x))
buf[:, nf + gc : nf + 2 * gc].copy_(x2)
x3 = self.conv3(buf[:, : nf + 2 * gc])
buf[:, nf + 2 * gc : nf + 3 * gc].copy_(x3)
x4 = self.conv4(buf[:, : nf + 3 * gc])
if self.conv1x1:
x4.add_(x2)
buf[:, nf + 3 * gc : nf + 4 * gc].copy_(x4)
x5 = self.conv5(buf)
return x5.mul_(0.2).add_(x)
@wraps(_orig_forward)
def RRDB_forward(self, x):
return self.RDB3(self.RDB2(self.RDB1(x))).mul_(0.2).add_(x)
ResidualDenseBlock_5C.__init__ = RDB5C_init
ResidualDenseBlock_5C.forward = RDB5C_forward
RRDB.forward = RRDB_forward
logger.info("Successfully patched Spandrel blocks")
except Exception as e:
logger.info(f"Failed to patch Spandrel blocks\n{type(e).__name__}: {e}")
try_patch_spandrel()
def _model(model: Callable, x: torch.Tensor) -> torch.Tensor:
if x.dtype == torch.float32 or model.architecture.name not in ("ATD", "DAT"):
return model(x)
# Spandrel does not correctly handle non-FP32 for ATD and DAT models
try:
# Force the upscaler to use the dtype it should for new tensors
torch.set_default_dtype(x.dtype)
# Using torch.device incurs a small amount of overhead, but makes sure we don't
# get errors when unsupported dtype tensors would be made on the CPU.
with torch.device(x.device):
return model(x)
finally:
torch.set_default_dtype(torch.float32)
def pil_rgb_to_tensor_bgr(img: Image.Image, param: torch.Tensor) -> torch.Tensor:
tensor = torch.from_numpy(np.asarray(img)).to(param.device)
tensor = tensor.to(param.dtype).mul_(1.0 / 255.0).permute(2, 0, 1)
return tensor[[2, 1, 0], ...].unsqueeze(0).contiguous()
def tensor_bgr_to_pil_rgb(tensor: torch.Tensor) -> Image.Image:
tensor = tensor[:, [2, 1, 0], ...]
tensor = tensor.squeeze(0).permute(1, 2, 0).mul_(255.0).round_().clamp_(0.0, 255.0)
return Image.fromarray(tensor.to(torch.uint8).cpu().numpy())
def pil_image_to_torch_bgr(img: Image.Image) -> torch.Tensor:
img = np.array(img.convert("RGB"))
img = img[:, :, ::-1] # flip RGB to BGR
img = np.transpose(img, (2, 0, 1)) # HWC to CHW
img = np.ascontiguousarray(img) / 255 # Rescale to [0, 1]
img = img[:, :, ::-1]
img = np.transpose(img, (2, 0, 1))
img = np.ascontiguousarray(img) / 255
return torch.from_numpy(img)
def torch_bgr_to_pil_image(tensor: torch.Tensor) -> Image.Image:
if tensor.ndim == 4:
# If we're given a tensor with a batch dimension, squeeze it out
# (but only if it's a batch of size 1).
if tensor.shape[0] != 1:
raise ValueError(f"{tensor.shape} does not describe a BCHW tensor")
tensor = tensor.squeeze(0)
assert tensor.ndim == 3, f"{tensor.shape} does not describe a CHW tensor"
# TODO: is `tensor.float().cpu()...numpy()` the most efficient idiom?
arr = tensor.float().cpu().clamp_(0, 1).numpy() # clamp
arr = 255.0 * np.moveaxis(arr, 0, 2) # CHW to HWC, rescale
arr = arr.round().astype(np.uint8)
arr = arr[:, :, ::-1] # flip BGR to RGB
arr = tensor.detach().float().cpu().numpy()
arr = 255.0 * np.moveaxis(arr, 0, 2)
arr = np.clip(arr, 0, 255).astype(np.uint8)
arr = arr[:, :, ::-1]
return Image.fromarray(arr, "RGB")
@torch.inference_mode()
def upscale_tensor_tiles(model: Callable, tensor: torch.Tensor, tile_size: int, overlap: int, desc: str) -> torch.Tensor:
_, _, H_in, W_in = tensor.shape
stride = tile_size - overlap
n_tiles_x, n_tiles_y = (W_in + stride - 1) // stride, (H_in + stride - 1) // stride
total_tiles = n_tiles_x * n_tiles_y
if tile_size <= 0 or total_tiles <= 4:
return _model(model, tensor)
device = tensor.device
dtype = tensor.dtype # Accumulate in native model dtype
accum = None
model_scale = None
H_out = W_out = None
last_mask = None
last_mask_key = None
def get_weight_mask(h, w, y, x):
"""Generate feathered mask for tile overlap"""
top, bottom, left, right = y > 0, y + h < H_out, x > 0, x + w < W_out
key = (h, w, top, bottom, left, right)
if key == last_mask_key:
return key, last_mask
elif overlap == 0:
mask = torch.ones((1, 1, h, w), device=device, dtype=dtype)
else:
ov_h, ov_w = min(overlap, h), min(overlap, w)
ramp_x, ramp_y = torch.ones(w, device=device, dtype=dtype), torch.ones(h, device=device, dtype=dtype)
fade_x, fade_y = torch.linspace(0, 1, ov_w, device=device, dtype=dtype), torch.linspace(0, 1, ov_h, device=device, dtype=dtype)
ramp_x[:ov_w].lerp_(fade_x, float(left))
ramp_x[-ov_w:].lerp_(fade_x.flip(0), float(right))
ramp_y[:ov_h].lerp_(fade_y, float(top))
ramp_y[-ov_h:].lerp_(fade_y.flip(0), float(bottom))
mask = (ramp_y[:, None] * ramp_x[None, :]).expand(1, 1, h, w)
return key, mask
with tqdm.tqdm(desc=desc, total=total_tiles) as pbar:
for tile_idx in range(total_tiles):
if shared.state.interrupted:
return None
# Loop in row-major or column-major, depending on aspect ratio to maximise hit-rate on cached mask
x_idx, y_idx = (tile_idx % n_tiles_x, tile_idx // n_tiles_x) if W_in >= H_in else (tile_idx // n_tiles_y, tile_idx % n_tiles_y)
x, y = x_idx * stride, y_idx * stride
tile = tensor[:, :, y : y + tile_size, x : x + tile_size]
out = _model(model, tile)
if model_scale is None:
model_scale = out.shape[-2] / tile.shape[-2]
H_out, W_out = int(H_in * model_scale), int(W_in * model_scale)
accum = torch.zeros((1, 4, H_out, W_out), dtype=dtype, device=device)
h_out, w_out = out.shape[-2:]
y_out, x_out = int(y * model_scale), int(x * model_scale)
ys, ye = y_out, y_out + h_out
xs, xe = x_out, x_out + w_out
last_mask_key, last_mask = get_weight_mask(h_out, w_out, y_out, x_out)
accum_slice = accum[:, :, ys:ye, xs:xe]
accum_slice[:, :3].addcmul_(out, last_mask)
accum_slice[:, 3:].add_(last_mask)
del tile, out
pbar.update(1)
del last_mask
return accum[:, :3].div_(accum[:, 3:].clamp_min_(1e-6))
def upscale_with_model_gpu(
model: Callable[[torch.Tensor], torch.Tensor],
img: Image.Image,
*,
tile_size: int,
tile_overlap: int = 0,
desc="tiled upscale",
) -> Image.Image:
tensor = pil_rgb_to_tensor_bgr(img, torch_utils.get_param(model))
out = upscale_tensor_tiles(model, tensor, tile_size, tile_overlap, desc)
return img if out is None else tensor_bgr_to_pil_rgb(out)
def upscale_pil_patch(model, img: Image.Image) -> Image.Image:
"""
Upscale a given PIL image using the given model.
"""
"""Upscale a given PIL image using the given model"""
param = torch_utils.get_param(model)
with torch.inference_mode():
tensor = pil_image_to_torch_bgr(img).unsqueeze(0) # add batch dimension
tensor = pil_image_to_torch_bgr(img).unsqueeze(0)
tensor = tensor.to(device=param.device, dtype=param.dtype)
with devices.without_autocast():
return torch_bgr_to_pil_image(model(tensor))
return torch_bgr_to_pil_image(_model(model, tensor))
def upscale_with_model(
def upscale_with_model_cpu(
model: Callable[[torch.Tensor], torch.Tensor],
img: Image.Image,
*,
@ -65,14 +236,20 @@ def upscale_with_model(
grid = images.split_grid(img, tile_size, tile_size, tile_overlap)
newtiles = []
with tqdm.tqdm(total=grid.tile_count, desc=desc, disable=not shared.opts.enable_upscale_progressbar) as p:
with tqdm.tqdm(
total=grid.tile_count,
desc=desc,
disable=not shared.opts.enable_upscale_progressbar,
) as p:
for y, h, row in grid.tiles:
newrow = []
for x, w, tile in row:
if shared.state.interrupted:
return img
break
logger.debug("Tile (%d, %d) %s...", x, y, tile)
output = upscale_pil_patch(model, tile)
scale_factor = output.width // tile.width
logger.debug("=> %s (scale factor %s)", output, scale_factor)
newrow.append([x * scale_factor, w * scale_factor, output])
p.update(1)
newtiles.append([y * scale_factor, h * scale_factor, newrow])
@ -88,103 +265,15 @@ def upscale_with_model(
return images.combine_grid(newgrid)
def tiled_upscale_2(
img: torch.Tensor,
model,
*,
tile_size: int,
tile_overlap: int,
scale: int,
device: torch.device,
desc="Tiled upscale",
):
# Alternative implementation of `upscale_with_model` originally used by
# SwinIR and ScuNET. It differs from `upscale_with_model` in that tiling and
# weighting is done in PyTorch space, as opposed to `images.Grid` doing it in
# Pillow space without weighting.
b, c, h, w = img.size()
tile_size = min(tile_size, h, w)
if tile_size <= 0:
logger.debug("Upscaling %s without tiling", img.shape)
return model(img)
stride = tile_size - tile_overlap
h_idx_list = list(range(0, h - tile_size, stride)) + [h - tile_size]
w_idx_list = list(range(0, w - tile_size, stride)) + [w - tile_size]
result = torch.zeros(
b,
c,
h * scale,
w * scale,
device=device,
dtype=img.dtype,
)
weights = torch.zeros_like(result)
logger.debug("Upscaling %s to %s with tiles", img.shape, result.shape)
with tqdm.tqdm(total=len(h_idx_list) * len(w_idx_list), desc=desc, disable=not shared.opts.enable_upscale_progressbar) as pbar:
for h_idx in h_idx_list:
if shared.state.interrupted or shared.state.skipped:
break
for w_idx in w_idx_list:
if shared.state.interrupted or shared.state.skipped:
break
# Only move this patch to the device if it's not already there.
in_patch = img[
...,
h_idx : h_idx + tile_size,
w_idx : w_idx + tile_size,
].to(device=device)
out_patch = model(in_patch)
result[
...,
h_idx * scale : (h_idx + tile_size) * scale,
w_idx * scale : (w_idx + tile_size) * scale,
].add_(out_patch)
out_patch_mask = torch.ones_like(out_patch)
weights[
...,
h_idx * scale : (h_idx + tile_size) * scale,
w_idx * scale : (w_idx + tile_size) * scale,
].add_(out_patch_mask)
pbar.update(1)
output = result.div_(weights)
return output
def upscale_2(
def upscale_with_model(
model: Callable[[torch.Tensor], torch.Tensor],
img: Image.Image,
model,
*,
tile_size: int,
tile_overlap: int,
scale: int,
desc: str,
):
"""
Convenience wrapper around `tiled_upscale_2` that handles PIL images.
"""
param = torch_utils.get_param(model)
tensor = pil_image_to_torch_bgr(img).to(dtype=param.dtype).unsqueeze(0) # add batch dimension
with torch.no_grad():
output = tiled_upscale_2(
tensor,
model,
tile_size=tile_size,
tile_overlap=tile_overlap,
scale=scale,
desc=desc,
device=param.device,
)
return torch_bgr_to_pil_image(output)
tile_overlap: int = 0,
desc="tiled upscale",
) -> Image.Image:
if shared.opts.composite_tiles_on_gpu:
return upscale_with_model_gpu(model, img, tile_size=tile_size, tile_overlap=tile_overlap, desc=f"{desc} (GPU Composite)")
else:
return upscale_with_model_cpu(model, img, tile_size=tile_size, tile_overlap=tile_overlap, desc=f"{desc} (CPU Composite)")