diff --git a/modules/face_restoration.py b/modules/face_restoration.py new file mode 100644 index 000000000..d17191fdf --- /dev/null +++ b/modules/face_restoration.py @@ -0,0 +1,17 @@ +from modules import shared + + +class FaceRestoration: + def name(self): + return "None" + + def restore(self, np_image): + return np_image + + +def restore_faces(np_image, p=None): + face_restorers = [x for x in shared.face_restorers if x.name() == shared.opts.face_restoration_model or shared.opts.face_restoration_model is None] + if len(face_restorers) == 0: + return np_image + face_restorer = face_restorers[0] + return face_restorer.restore(np_image, p) diff --git a/modules/postprocess/codeformer_model.py b/modules/postprocess/codeformer_model.py index 4493ab8a3..327332db1 100644 --- a/modules/postprocess/codeformer_model.py +++ b/modules/postprocess/codeformer_model.py @@ -4,6 +4,7 @@ import torch import modules.detailer from modules import shared, devices, modelloader, errors from modules.paths import models_path +from installer import install # codeformer people made a choice to include modified basicsr library to their project which makes # it utterly impossible to use it alongside with other libraries that also use basicsr, like GFPGAN. @@ -37,9 +38,13 @@ def setup_model(dirname): self.cmd_dir = dirname def create_models(self): - from modules.postprocess.codeformer_arch import CodeFormer - from facelib.utils.detailer_helper import FaceRestoreHelper - from facelib.detection.retinaface import retinaface + try: + from modules.postprocess.codeformer_arch import CodeFormer + from facelib.utils.detailer_helper import FaceRestoreHelper + from facelib.detection.retinaface import retinaface + except Exception as e: + shared.log.error(f"CodeFormer error: {e}") + return None, None if self.net is not None and self.face_helper is not None: self.net.to(devices.device) return self.net, self.face_helper @@ -108,7 +113,7 @@ def setup_model(dirname): have_codeformer = True global codeformer # pylint: disable=global-statement codeformer = FaceRestorerCodeFormer(dirname) - shared.detailers.append(codeformer) + shared.face_restorers.append(codeformer) except Exception as e: errors.display(e, 'codeformer') diff --git a/modules/postprocess/gfpgan_model.py b/modules/postprocess/gfpgan_model.py index 03274b7f4..ad0aa8221 100644 --- a/modules/postprocess/gfpgan_model.py +++ b/modules/postprocess/gfpgan_model.py @@ -108,6 +108,6 @@ def setup_model(dirname): def restore(self, np_image, p=None): # pylint: disable=unused-argument return gfpgan_fix_faces(np_image) - shared.detailers.append(FaceRestorerGFPGAN()) + shared.face_restorers.append(FaceRestorerGFPGAN()) except Exception as e: errors.log.error(f'GFPGan failed to initialize: {e}') diff --git a/modules/processing.py b/modules/processing.py index d582806cb..6078ad10c 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -4,7 +4,7 @@ import time from contextlib import nullcontext import numpy as np from PIL import Image, ImageOps -from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_hijack_freeu, sd_models, sd_vae, processing_helpers, timer +from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_hijack_freeu, sd_models, sd_vae, processing_helpers, timer, face_restoration from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, StableDiffusionProcessingControl # pylint: disable=unused-import from modules.processing_info import create_infotext @@ -354,6 +354,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed: else: sample = validate_sample(sample) image = Image.fromarray(sample) + sample = face_restoration.restore_faces(sample, p) if p.detailer: if not p.do_not_save_samples and shared.opts.save_images_before_detailer: info = create_infotext(p, p.prompts, p.seeds, p.subseeds, index=i) diff --git a/modules/shared.py b/modules/shared.py index 81befa466..e98dc8383 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -44,6 +44,7 @@ clip_model = None interrogator = modules.interrogate.InterrogateModels(os.path.join("models", "interrogate")) sd_upscalers = [] detailers = [] +face_restorers = [] yolo = None tab_names = [] extra_networks = [] @@ -799,6 +800,9 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), { "code_former_weight": OptionInfo(0.2, "CodeFormer weight parameter", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}), "detailer_unload": OptionInfo(False, "Move detailer model to CPU when complete"), + "postprocessing_sep_face_restore": OptionInfo("