feat: 🚧 wrapper for GFPGAN bg upscaler
this hooks into comfy's core upscaler model loader. It seems to work, as in it doesn't fail but it's not producing the proper results.
This commit is contained in:
@@ -3,6 +3,13 @@ import re
|
||||
import os
|
||||
|
||||
base_log_level = logging.DEBUG if os.environ.get("MTB_DEBUG") else logging.INFO
|
||||
print(f"Log level: {base_log_level}")
|
||||
|
||||
|
||||
# Custom object that discards the output
|
||||
class NullWriter:
|
||||
def write(self, text):
|
||||
pass
|
||||
|
||||
|
||||
class Formatter(logging.Formatter):
|
||||
@@ -32,15 +39,16 @@ class Formatter(logging.Formatter):
|
||||
|
||||
def mklog(name, level=base_log_level):
|
||||
logger = logging.getLogger(name)
|
||||
# set this to the highest level of all handlers
|
||||
logger.setLevel(level)
|
||||
# create console handler with a higher log level
|
||||
|
||||
for handler in logger.handlers:
|
||||
logger.removeHandler(handler)
|
||||
|
||||
ch = logging.StreamHandler()
|
||||
ch.setLevel(level)
|
||||
|
||||
ch.setFormatter(Formatter())
|
||||
|
||||
logger.addHandler(ch)
|
||||
|
||||
return logger
|
||||
|
||||
|
||||
|
||||
+70
-11
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
from gfpgan import GFPGANer
|
||||
import cv2
|
||||
import numpy as np
|
||||
@@ -6,8 +7,12 @@ from pathlib import Path
|
||||
import folder_paths
|
||||
from basicsr.utils import imwrite
|
||||
from PIL import Image
|
||||
from ..utils import pil2tensor, tensor2pil
|
||||
from ..utils import pil2tensor, tensor2pil, np2tensor, tensor2np
|
||||
import torch
|
||||
from munch import Munch
|
||||
from ..log import NullWriter, log
|
||||
from comfy import model_management
|
||||
import comfy
|
||||
|
||||
|
||||
class LoadFaceEnhanceModel:
|
||||
@@ -39,28 +44,78 @@ class LoadFaceEnhanceModel:
|
||||
),
|
||||
"upscale": ("INT", {"default": 2}),
|
||||
},
|
||||
"optional": {"bg_model_path": ("UPSCALE_MODEL", {"default": "None"})},
|
||||
"optional": {"bg_upsampler": ("UPSCALE_MODEL", {"default": None})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FACEENHANCE_MODEL",)
|
||||
RETURN_NAMES = ("model",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "face"
|
||||
|
||||
def load_model(self, model_name, upscale=2, bg_upsampler="realesrgan"):
|
||||
def load_model(self, model_name, upscale=2, bg_upsampler=None):
|
||||
basic = "RestoreFormer" not in model_name
|
||||
|
||||
root = self.get_models_root()
|
||||
|
||||
return (
|
||||
GFPGANer(
|
||||
model_path=(root / model_name).as_posix(),
|
||||
upscale=upscale,
|
||||
arch="clean" if basic else "RestoreFormer", # or original for v1.0 only
|
||||
channel_multiplier=2, # 1 for v1.0 only
|
||||
bg_upsampler=None, # TODO:"realesrgan",
|
||||
),
|
||||
if bg_upsampler is not None:
|
||||
log.warning(
|
||||
f"Upscale value overridden to {bg_upsampler.scale} from bg_upsampler"
|
||||
)
|
||||
upscale = bg_upsampler.scale
|
||||
bg_upsampler = BGUpscaleWrapper(bg_upsampler)
|
||||
|
||||
sys.stdout = NullWriter()
|
||||
model = GFPGANer(
|
||||
model_path=(root / model_name).as_posix(),
|
||||
upscale=upscale,
|
||||
arch="clean" if basic else "RestoreFormer", # or original for v1.0 only
|
||||
channel_multiplier=2, # 1 for v1.0 only
|
||||
bg_upsampler=bg_upsampler,
|
||||
)
|
||||
|
||||
sys.stdout = sys.__stdout__
|
||||
return (model,)
|
||||
|
||||
|
||||
class BGUpscaleWrapper:
|
||||
def __init__(self, upscale_model) -> None:
|
||||
self.upscale_model = upscale_model
|
||||
|
||||
def enhance(self, img: Image, outscale=2):
|
||||
device = model_management.get_torch_device()
|
||||
self.upscale_model.to(device)
|
||||
|
||||
tile = 128 + 64
|
||||
overlap = 8
|
||||
|
||||
imgt = np2tensor(img)
|
||||
imgt = imgt.movedim(-1, -3).to(device)
|
||||
|
||||
steps = imgt.shape[0] * comfy.utils.get_tiled_scale_steps(
|
||||
imgt.shape[3], imgt.shape[2], tile_x=tile, tile_y=tile, overlap=overlap
|
||||
)
|
||||
|
||||
log.debug(f"Steps: {steps}")
|
||||
|
||||
pbar = comfy.utils.ProgressBar(steps)
|
||||
|
||||
s = comfy.utils.tiled_scale(
|
||||
imgt,
|
||||
lambda a: self.upscale_model(a),
|
||||
tile_x=tile,
|
||||
tile_y=tile,
|
||||
overlap=overlap,
|
||||
upscale_amount=self.upscale_model.scale,
|
||||
pbar=pbar,
|
||||
)
|
||||
|
||||
self.upscale_model.cpu()
|
||||
s = torch.clamp(s.movedim(-3, -1), min=0, max=1.0)
|
||||
return (tensor2np(s),)
|
||||
|
||||
|
||||
import sys
|
||||
|
||||
|
||||
class RestoreFace:
|
||||
def __init__(self) -> None:
|
||||
@@ -103,6 +158,8 @@ class RestoreFace:
|
||||
width, height = image.size
|
||||
|
||||
source_img = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
|
||||
|
||||
sys.stdout = NullWriter()
|
||||
cropped_faces, restored_faces, restored_img = model.enhance(
|
||||
source_img,
|
||||
has_aligned=aligned,
|
||||
@@ -111,6 +168,8 @@ class RestoreFace:
|
||||
# TODO: weight has no effect in 1.3 and 1.4 (only tested these for now...)
|
||||
weight=weight,
|
||||
)
|
||||
sys.stdout = sys.__stdout__
|
||||
log.warning(f"Weight value has no effect for now. (value: {weight})")
|
||||
|
||||
if save_tmp_steps:
|
||||
self.save_intermediate_images(cropped_faces, restored_faces, height, width)
|
||||
|
||||
+30
-8
@@ -1,5 +1,6 @@
|
||||
# region imports
|
||||
from ifnude import detect
|
||||
import onnxruntime
|
||||
from pathlib import Path
|
||||
from PIL import Image
|
||||
from typing import List, Set, Tuple
|
||||
@@ -11,9 +12,10 @@ import numpy as np
|
||||
import os
|
||||
import tempfile
|
||||
import torch
|
||||
|
||||
from insightface.model_zoo.inswapper import INSwapper
|
||||
from ..utils import pil2tensor, tensor2pil
|
||||
from ..log import mklog
|
||||
from ..log import mklog, NullWriter
|
||||
import sys
|
||||
|
||||
# endregion
|
||||
|
||||
@@ -50,7 +52,15 @@ class LoadFaceSwapModel:
|
||||
folder_paths.models_dir, "insightface", faceswap_model
|
||||
)
|
||||
log.info(f"Loading model {model_path}")
|
||||
return (insightface.model_zoo.get_model(model_path),)
|
||||
return (
|
||||
INSwapper(
|
||||
model_path,
|
||||
onnxruntime.InferenceSession(
|
||||
path_or_bytes=model_path,
|
||||
providers=onnxruntime.get_available_providers(),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# region roop node
|
||||
@@ -71,6 +81,7 @@ class FaceSwap:
|
||||
"reference": ("IMAGE",),
|
||||
"faces_index": ("STRING", {"default": "0"}),
|
||||
"faceswap_model": ("FACESWAP_MODEL", {"default": "None"}),
|
||||
"allow_nsfw": (["true", "false"], {"default": "false"}),
|
||||
},
|
||||
"optional": {"debug": (["true", "false"], {"default": "false"})},
|
||||
}
|
||||
@@ -85,7 +96,8 @@ class FaceSwap:
|
||||
reference: torch.Tensor,
|
||||
faces_index: str,
|
||||
faceswap_model,
|
||||
debug: str,
|
||||
allow_nsfw="fase",
|
||||
debug="false",
|
||||
):
|
||||
def do_swap(img):
|
||||
img = tensor2pil(img)
|
||||
@@ -93,8 +105,11 @@ class FaceSwap:
|
||||
face_ids = {
|
||||
int(x) for x in faces_index.strip(",").split(",") if x.isnumeric()
|
||||
}
|
||||
|
||||
swapped = swap_face(ref, img, faceswap_model, face_ids)
|
||||
sys.stdout = NullWriter()
|
||||
swapped = swap_face(
|
||||
ref, img, faceswap_model, face_ids, allow_nsfw == "true"
|
||||
)
|
||||
sys.stdout = sys.__stdout__
|
||||
return pil2tensor(swapped)
|
||||
|
||||
batch_count = image.size(0)
|
||||
@@ -146,14 +161,19 @@ def swap_face(
|
||||
target_img: Image.Image,
|
||||
face_swapper_model=None,
|
||||
faces_index: Set[int] = None,
|
||||
allow_nsfw=False,
|
||||
) -> Image.Image:
|
||||
if faces_index is None:
|
||||
faces_index = {0}
|
||||
log.debug(f"Swapping faces: {faces_index}")
|
||||
result_image = target_img
|
||||
converted = convert_to_sd(target_img)
|
||||
scale, fn = converted[0], converted[1]
|
||||
if face_swapper_model is not None and not scale:
|
||||
nsfw, fn = converted[0], converted[1]
|
||||
|
||||
if nsfw and allow_nsfw:
|
||||
nsfw = False
|
||||
|
||||
if face_swapper_model is not None and not nsfw:
|
||||
if isinstance(source_img, str): # source_img is a base64 string
|
||||
import base64, io
|
||||
|
||||
@@ -175,7 +195,9 @@ def swap_face(
|
||||
for face_num in faces_index:
|
||||
target_face = get_face_single(target_img, face_index=face_num)
|
||||
if target_face is not None:
|
||||
sys.stdout = NullWriter()
|
||||
result = face_swapper_model.get(result, target_face, source_face)
|
||||
sys.stdout = sys.__stdout__
|
||||
else:
|
||||
log.warning(f"No target face found for {face_num}")
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from skimage.color import rgb2hsv, hsv2rgb
|
||||
import numpy as np
|
||||
import torchvision.transforms.functional as F
|
||||
from PIL import Image, ImageChops
|
||||
from ..utils import tensor2pil, pil2tensor, img_np_to_tensor, img_tensor_to_np
|
||||
from ..utils import tensor2pil, pil2tensor, np2tensor, tensor2np
|
||||
import cv2
|
||||
import torch
|
||||
from ..log import log
|
||||
@@ -362,7 +362,7 @@ class DeglazeImage:
|
||||
FUNCTION = "deglaze_image"
|
||||
|
||||
def deglaze_image(self, image):
|
||||
return (img_np_to_tensor(deglaze_np_img(img_tensor_to_np(image))),)
|
||||
return (np2tensor(deglaze_np_img(tensor2np(image))),)
|
||||
|
||||
|
||||
class MaskToImage:
|
||||
@@ -388,7 +388,7 @@ class MaskToImage:
|
||||
FUNCTION = "render_mask"
|
||||
|
||||
def render_mask(self, mask, color, background):
|
||||
mask = img_tensor_to_np(mask)
|
||||
mask = tensor2np(mask)
|
||||
mask = Image.fromarray(mask).convert("L")
|
||||
|
||||
image = Image.new("RGBA", mask.size, color=color)
|
||||
|
||||
Reference in New Issue
Block a user