mirror of
https://github.com/lllyasviel/stable-diffusion-webui-forge.git
synced 2026-07-21 21:01:24 +08:00
upscaler
This commit is contained in:
parent
3d8212a718
commit
9d70067221
@ -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:
|
||||
|
||||
@ -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',
|
||||
),
|
||||
]
|
||||
@ -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
|
||||
|
||||
@ -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',
|
||||
)
|
||||
@ -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):
|
||||
|
||||
@ -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,
|
||||
),
|
||||
]
|
||||
@ -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"), {
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
@ -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)")
|
||||
|
||||
Loading…
Reference in New Issue
Block a user