diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..bee8a64 --- /dev/null +++ b/.gitignore @@ -0,0 +1 @@ +__pycache__ diff --git a/__init__.py b/__init__.py index e69de29..7bd32dc 100644 --- a/__init__.py +++ b/__init__.py @@ -0,0 +1,66 @@ +import sys +import requests +from tqdm import tqdm +from .config import * + +# import pydevd_pycharm +# pydevd_pycharm.settrace('49.7.62.197', port=10090, stdoutToServer=True, stderrToServer=True) + +sys.path.append(utils_path) + +from .node import * + +def urldownload_progressbar(url, file_path): + response = requests.get(url, stream=True) + total_size = int(response.headers.get('content-length', 0)) + progress_bar = tqdm(total=total_size, unit='B', unit_scale=True) + with open(file_path, 'wb') as f: + for chunk in response.iter_content(1024): + if chunk: + f.write(chunk) + progress_bar.update(len(chunk)) + + progress_bar.close() + +print("Start Setting weights") +for url, filename in zip(urls, filenames): + if os.path.exists(filename): + continue + print(f"Start Downloading: {url}") + os.makedirs(os.path.dirname(filename), exist_ok=True) + urldownload_progressbar(url, filename) + +NODE_CLASS_MAPPINGS = { + "RetainFace": RetainFace, + "FaceFusion": FaceFusion, + "RatioMerge2Image": RatioMerge2Image, + "MaskMerge2Image": MaskMerge2Image, + "ReplaceBoxImg": ReplaceBoxImg, + "ExpandMaskBox": ExpandMaskFaceWidth, + "BoxCropImage": BoxCropImage, + "ColorTransfer": ColorTransfer, + "FaceSkin": FaceSkin, + "MaskDilateErode": MaskDilateErode, + "SkinRetouching": SkinRetouching, + "PortraitEnhancement": PortraitEnhancement, + "ResizeImage": ResizeImage, + "GetImageInfo": GetImageInfo, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "RetainFace": "RetainFace", + "FaceFusion": "FaceFusion", + "RatioMerge2Image": "RatioMerge2Image", + "MaskMerge2Image": "MaskMerge2Image", + "ReplaceBoxImg": "ReplaceBoxImg", + "ExpandMaskBox": "ExpandMaskBox", + "BoxCropImage": "BoxCropImage", + "ColorTransfer": "ColorTransfer", + "FaceSkin": "FaceSkin", + "MaskDilateErode": "MaskDilateErode", + "SkinRetouching": "SkinRetouching", + "PortraitEnhancement": "PortraitEnhancement", + "ResizeImage": "ResizeImage", + "GetImageInfo": "GetImageInfo", +} + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] diff --git a/config.py b/config.py new file mode 100644 index 0000000..c082052 --- /dev/null +++ b/config.py @@ -0,0 +1,35 @@ +import os, glob +from folder_paths import folder_names_and_paths + +root_path = os.path.dirname(__file__) +utils_path = os.path.join(os.path.dirname(__file__), "utils") +models_path = os.path.join(os.path.dirname(__file__), "models") +# save_dirs +urls = [ + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/control_v11p_sd15_openpose.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/control_v11p_sd15_canny.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/control_v11f1e_sd15_tile.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/control_sd15_random_color.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/FilmVelvia3.safetensors", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/body_pose_model.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/facenet.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/hand_pose_model.pth", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/vae-ft-mse-840000-ema-pruned.ckpt", + "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/face_skin.pth", +] +filenames = [ + os.path.join(folder_names_and_paths['controlnet'][0][0], "control_v11p_sd15_openpose.pth"), + os.path.join(folder_names_and_paths['controlnet'][0][0], "control_v11p_sd15_canny.pth"), + os.path.join(folder_names_and_paths['controlnet'][0][0], "control_v11f1e_sd15_tile.pth"), + os.path.join(folder_names_and_paths['controlnet'][0][0], "control_sd15_random_color.pth"), + os.path.join(folder_names_and_paths['loras'][0][0], "FilmVelvia3.safetensors"), + os.path.join(folder_names_and_paths['controlnet'][0][0], "body_pose_model.pth"), + os.path.join(folder_names_and_paths['controlnet'][0][0], "facenet.pth"), + os.path.join(folder_names_and_paths['controlnet'][0][0], "hand_pose_model.pth"), + os.path.join(folder_names_and_paths['vae'][0][0], "VAE/vae-ft-mse-840000-ema-pruned.ckpt"), + os.path.join(models_path, "face_skin.pth"), +] +# prompts +validation_prompt = "easyphoto_face, easyphoto, 1person" +DEFAULT_POSITIVE = '(cloth:1.7), (best quality), (realistic, photo-realistic:1.2), detailed skin, beautiful, cool, finely detail, light smile, extremely detailed CG unity 8k wallpaper, huge filesize, best quality, realistic, photo-realistic, ultra high res, raw photo, put on makeup' +DEFAULT_NEGATIVE = '(bags under the eyes:1.5), (Bags under eyes:1.5), (glasses:1.5), (naked:2.0), nude, (nsfw:2.0), breasts, penis, cum, (worst quality:2), (low quality:2), (normal quality:2), over red lips, hair, teeth, lowres, watermark, badhand, (normal quality:2), lowres, bad anatomy, bad hands, normal quality, mural,' diff --git a/face_process_utils.py b/face_process_utils.py new file mode 100644 index 0000000..dd9dd35 --- /dev/null +++ b/face_process_utils.py @@ -0,0 +1,536 @@ +import cv2 +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +import torchvision.transforms as transforms +from PIL import Image +from skimage import transform +from protrait.img_utils import np_to_mask + +import pydevd_pycharm +pydevd_pycharm.settrace('49.7.62.197', port=10090, stdoutToServer=True, stderrToServer=True) + +def safe_get_box_mask_keypoints(image, retinaface_result, crop_ratio, face_seg, mask_type): + ''' + Inputs: + image Input image + retinaface_result The detection results of retinaface + crop_ratio The proportion of facial clipping and expansion + face_seg Facial segmentation model + mask_type The type of facial segmentation methods, one is loop and the other is skin, and the result of facial segmentation is the skin or frame of the face + + Outputs: + retinaface_box After box amplification, the box relative to the original image + retinaface_keypoints Points relative to the original image + retinaface_mask_pil Segmentation Results + ''' + h, w, c = np.shape(image) + if len(retinaface_result['boxes']) != 0: + retinaface_boxs = [] + retinaface_keypoints = [] + retinaface_mask_pils = [] + retinaface_masks = [] + for index in range(len(retinaface_result['boxes'])): + # 获得retinaface的box并做扩充 + retinaface_box = np.array(retinaface_result['boxes'][index]) + face_width = retinaface_box[2] - retinaface_box[0] + face_height = retinaface_box[3] - retinaface_box[1] + retinaface_box[0] = np.clip(np.array(retinaface_box[0], np.int32) - face_width * (crop_ratio - 1) / 2, 0, w - 1) + retinaface_box[1] = np.clip(np.array(retinaface_box[1], np.int32) - face_height * (crop_ratio - 1) / 2, 0, h - 1) + retinaface_box[2] = np.clip(np.array(retinaface_box[2], np.int32) + face_width * (crop_ratio - 1) / 2, 0, w - 1) + retinaface_box[3] = np.clip(np.array(retinaface_box[3], np.int32) + face_height * (crop_ratio - 1) / 2, 0, h - 1) + retinaface_box = np.array(retinaface_box, np.int32) + retinaface_boxs.append(retinaface_box) + + # 检测关键点 + retinaface_keypoint = np.reshape(retinaface_result['keypoints'][index], [5, 2]) + retinaface_keypoint = np.array(retinaface_keypoint, np.float32) + retinaface_keypoints.append(retinaface_keypoint) + + # mask部分 + retinaface_crop = image.crop(np.int32(retinaface_box)) + retinaface_mask = np.zeros_like(np.array(image, np.uint8)) + if mask_type == "skin": + retinaface_sub_mask = face_seg(retinaface_crop) + retinaface_mask[retinaface_box[1]:retinaface_box[3], retinaface_box[0]:retinaface_box[2]] = np.expand_dims(retinaface_sub_mask, -1) + else: + retinaface_mask[retinaface_box[1]:retinaface_box[3], retinaface_box[0]:retinaface_box[2]] = 255 + retinaface_masks.append(retinaface_mask) + retinaface_mask_pil = Image.fromarray(np.uint8(retinaface_mask)) + retinaface_mask_pils.append(retinaface_mask_pil) + + retinaface_boxs = np.array(retinaface_boxs) + argindex = np.argsort(retinaface_boxs[:, 0]) + retinaface_boxs = [retinaface_boxs[index] for index in argindex] + retinaface_keypoints = [retinaface_keypoints[index] for index in argindex] + retinaface_mask_pils = [retinaface_mask_pils[index] for index in argindex] + retinaface_mask_np = [retinaface_masks[index] for index in argindex] + mask_tensor = np_to_mask(retinaface_mask_np[0]) + return retinaface_boxs, retinaface_keypoints, retinaface_mask_pils, mask_tensor + + else: + retinaface_box = np.array([]) + retinaface_keypoints = np.array([]) + retinaface_mask = np.zeros_like(np.array(image, np.uint8)) + retinaface_mask_pil = Image.fromarray(np.uint8(retinaface_mask)) + + return retinaface_box, retinaface_keypoints, retinaface_mask_pil, retinaface_mask + +def crop_and_paste(source_image, source_image_mask, target_image, source_five_point, target_five_point, source_box): + """ + Applies a face replacement by cropping and pasting one face onto another image. + + Args: + source_image (PIL.Image): The source image containing the face to be pasted. + source_image_mask (PIL.Image): The mask representing the face in the source image. + target_image (PIL.Image): The target image where the face will be pasted. + source_five_point (numpy.ndarray): Five key points of the face in the source image. + target_five_point (numpy.ndarray): Five key points of the corresponding face in the target image. + source_box (list): Coordinates of the bounding box around the face in the source image. + + Returns: + PIL.Image: The resulting image with the pasted face. + + Notes: + The function takes a source image, its corresponding mask, a target image, key points, and the bounding box + around the face in the source image. It then aligns and pastes the face from the source image onto the + corresponding location in the target image, taking into account the key points and bounding box. + """ + source_five_point = np.reshape(source_five_point, [5, 2]) - np.array(source_box[:2]) + target_five_point = np.reshape(target_five_point, [5, 2]) + + crop_source_image = source_image.crop(np.int32(source_box)) + crop_source_image_mask = source_image_mask.crop(np.int32(source_box)) + source_five_point, target_five_point = np.array(source_five_point), np.array(target_five_point) + + tform = transform.SimilarityTransform() + # 程序直接估算出转换矩阵M + tform.estimate(source_five_point, target_five_point) + M = tform.params[0:2, :] + + warped = cv2.warpAffine(np.array(crop_source_image), M, np.shape(target_image)[:2][::-1], borderValue=0.0) + warped_mask = cv2.warpAffine(np.array(crop_source_image_mask), M, np.shape(target_image)[:2][::-1], borderValue=0.0) + + mask = np.float32(warped_mask == 0) + output = mask * np.float32(target_image) + (1 - mask) * np.float32(warped) + return output + +def call_face_crop(retinaface_detection, image, crop_ratio, prefix="tmp"): + # retinaface detect + retinaface_result = retinaface_detection(image) + # get mask and keypoints + retinaface_box, retinaface_keypoints, retinaface_mask_pil, retinaface_mask_tensor = safe_get_box_mask_keypoints(image, retinaface_result, crop_ratio, None, "crop") + + return retinaface_box, retinaface_keypoints, retinaface_mask_pil, retinaface_mask_tensor + +def color_transfer(sc, dc): + """ + Transfer color distribution from of sc, referred to dc. + + Args: + sc (numpy.ndarray): input image to be transfered. + dc (numpy.ndarray): reference image + + Returns: + numpy.ndarray: Transferred color distribution on the sc. + """ + + def get_mean_and_std(img): + x_mean, x_std = cv2.meanStdDev(img) + x_mean = np.hstack(np.around(x_mean, 2)) + x_std = np.hstack(np.around(x_std, 2)) + return x_mean, x_std + + sc = cv2.cvtColor(sc, cv2.COLOR_BGR2LAB) # 转换颜色空间为clelab + s_mean, s_std = get_mean_and_std(sc) + dc = cv2.cvtColor(dc, cv2.COLOR_BGR2LAB) # 转换颜色空间为clelab + t_mean, t_std = get_mean_and_std(dc) + img_n = ((sc - s_mean) * (t_std / s_std)) + t_mean + np.putmask(img_n, img_n > 255, 255) + np.putmask(img_n, img_n < 0, 0) + dst = cv2.cvtColor(cv2.convertScaleAbs(img_n), cv2.COLOR_LAB2BGR) + return dst + +def conv3x3(in_planes, out_planes, stride=1): + """3x3 convolution with padding""" + return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, + padding=1, bias=False) + +class BasicBlock(nn.Module): + def __init__(self, in_chan, out_chan, stride=1): + super(BasicBlock, self).__init__() + self.conv1 = conv3x3(in_chan, out_chan, stride) + self.bn1 = nn.BatchNorm2d(out_chan) + self.conv2 = conv3x3(out_chan, out_chan) + self.bn2 = nn.BatchNorm2d(out_chan) + self.relu = nn.ReLU(inplace=True) + self.downsample = None + if in_chan != out_chan or stride != 1: + self.downsample = nn.Sequential( + nn.Conv2d(in_chan, out_chan, + kernel_size=1, stride=stride, bias=False), + nn.BatchNorm2d(out_chan), + ) + + def forward(self, x): + residual = self.conv1(x) + residual = F.relu(self.bn1(residual)) + residual = self.conv2(residual) + residual = self.bn2(residual) + + shortcut = x + if self.downsample is not None: + shortcut = self.downsample(x) + + out = shortcut + residual + out = self.relu(out) + return out + +def create_layer_basic(in_chan, out_chan, bnum, stride=1): + layers = [BasicBlock(in_chan, out_chan, stride=stride)] + for i in range(bnum - 1): + layers.append(BasicBlock(out_chan, out_chan, stride=1)) + return nn.Sequential(*layers) + +class Resnet18(nn.Module): + def __init__(self): + super(Resnet18, self).__init__() + self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3, + bias=False) + self.bn1 = nn.BatchNorm2d(64) + self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1) + self.layer1 = create_layer_basic(64, 64, bnum=2, stride=1) + self.layer2 = create_layer_basic(64, 128, bnum=2, stride=2) + self.layer3 = create_layer_basic(128, 256, bnum=2, stride=2) + self.layer4 = create_layer_basic(256, 512, bnum=2, stride=2) + + def forward(self, x): + x = self.conv1(x) + x = F.relu(self.bn1(x)) + x = self.maxpool(x) + + x = self.layer1(x) + feat8 = self.layer2(x) # 1/8 + feat16 = self.layer3(feat8) # 1/16 + feat32 = self.layer4(feat16) # 1/32 + return feat8, feat16, feat32 + + def get_params(self): + wd_params, nowd_params = [], [] + for name, module in self.named_modules(): + if isinstance(module, (nn.Linear, nn.Conv2d)): + wd_params.append(module.weight) + if not module.bias is None: + nowd_params.append(module.bias) + elif isinstance(module, nn.BatchNorm2d): + nowd_params += list(module.parameters()) + return wd_params, nowd_params + +class ConvBNReLU(nn.Module): + def __init__(self, in_chan, out_chan, ks=3, stride=1, padding=1, *args, **kwargs): + super(ConvBNReLU, self).__init__() + self.conv = nn.Conv2d(in_chan, + out_chan, + kernel_size=ks, + stride=stride, + padding=padding, + bias=False) + self.bn = nn.BatchNorm2d(out_chan) + self.init_weight() + + def forward(self, x): + x = self.conv(x) + x = F.relu(self.bn(x)) + return x + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + +class BiSeNetOutput(nn.Module): + def __init__(self, in_chan, mid_chan, n_classes, *args, **kwargs): + super(BiSeNetOutput, self).__init__() + self.conv = ConvBNReLU(in_chan, mid_chan, ks=3, stride=1, padding=1) + self.conv_out = nn.Conv2d(mid_chan, n_classes, kernel_size=1, bias=False) + self.init_weight() + + def forward(self, x): + x = self.conv(x) + x = self.conv_out(x) + return x + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + + def get_params(self): + wd_params, nowd_params = [], [] + for name, module in self.named_modules(): + if isinstance(module, nn.Linear) or isinstance(module, nn.Conv2d): + wd_params.append(module.weight) + if not module.bias is None: + nowd_params.append(module.bias) + elif isinstance(module, nn.BatchNorm2d): + nowd_params += list(module.parameters()) + return wd_params, nowd_params + +class AttentionRefinementModule(nn.Module): + def __init__(self, in_chan, out_chan, *args, **kwargs): + super(AttentionRefinementModule, self).__init__() + self.conv = ConvBNReLU(in_chan, out_chan, ks=3, stride=1, padding=1) + self.conv_atten = nn.Conv2d(out_chan, out_chan, kernel_size=1, bias=False) + self.bn_atten = nn.BatchNorm2d(out_chan) + self.sigmoid_atten = nn.Sigmoid() + self.init_weight() + + def forward(self, x): + feat = self.conv(x) + atten = F.avg_pool2d(feat, feat.size()[2:]) + atten = self.conv_atten(atten) + atten = self.bn_atten(atten) + atten = self.sigmoid_atten(atten) + out = torch.mul(feat, atten) + return out + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + +class ContextPath(nn.Module): + def __init__(self, *args, **kwargs): + super(ContextPath, self).__init__() + self.resnet = Resnet18() + self.arm16 = AttentionRefinementModule(256, 128) + self.arm32 = AttentionRefinementModule(512, 128) + self.conv_head32 = ConvBNReLU(128, 128, ks=3, stride=1, padding=1) + self.conv_head16 = ConvBNReLU(128, 128, ks=3, stride=1, padding=1) + self.conv_avg = ConvBNReLU(512, 128, ks=1, stride=1, padding=0) + + self.init_weight() + + def forward(self, x): + H0, W0 = x.size()[2:] + feat8, feat16, feat32 = self.resnet(x) + H8, W8 = feat8.size()[2:] + H16, W16 = feat16.size()[2:] + H32, W32 = feat32.size()[2:] + + avg = F.avg_pool2d(feat32, feat32.size()[2:]) + avg = self.conv_avg(avg) + avg_up = F.interpolate(avg, (H32, W32), mode='nearest') + + feat32_arm = self.arm32(feat32) + feat32_sum = feat32_arm + avg_up + feat32_up = F.interpolate(feat32_sum, (H16, W16), mode='nearest') + feat32_up = self.conv_head32(feat32_up) + + feat16_arm = self.arm16(feat16) + feat16_sum = feat16_arm + feat32_up + feat16_up = F.interpolate(feat16_sum, (H8, W8), mode='nearest') + feat16_up = self.conv_head16(feat16_up) + + return feat8, feat16_up, feat32_up # x8, x8, x16 + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + + def get_params(self): + wd_params, nowd_params = [], [] + for name, module in self.named_modules(): + if isinstance(module, (nn.Linear, nn.Conv2d)): + wd_params.append(module.weight) + if not module.bias is None: + nowd_params.append(module.bias) + elif isinstance(module, nn.BatchNorm2d): + nowd_params += list(module.parameters()) + return wd_params, nowd_params + +### This is not used, since I replace this with the resnet feature with the same size +class SpatialPath(nn.Module): + def __init__(self, *args, **kwargs): + super(SpatialPath, self).__init__() + self.conv1 = ConvBNReLU(3, 64, ks=7, stride=2, padding=3) + self.conv2 = ConvBNReLU(64, 64, ks=3, stride=2, padding=1) + self.conv3 = ConvBNReLU(64, 64, ks=3, stride=2, padding=1) + self.conv_out = ConvBNReLU(64, 128, ks=1, stride=1, padding=0) + self.init_weight() + + def forward(self, x): + feat = self.conv1(x) + feat = self.conv2(feat) + feat = self.conv3(feat) + feat = self.conv_out(feat) + return feat + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + + def get_params(self): + wd_params, nowd_params = [], [] + for name, module in self.named_modules(): + if isinstance(module, nn.Linear) or isinstance(module, nn.Conv2d): + wd_params.append(module.weight) + if not module.bias is None: + nowd_params.append(module.bias) + elif isinstance(module, nn.BatchNorm2d): + nowd_params += list(module.parameters()) + return wd_params, nowd_params + +class FeatureFusionModule(nn.Module): + def __init__(self, in_chan, out_chan, *args, **kwargs): + super(FeatureFusionModule, self).__init__() + self.convblk = ConvBNReLU(in_chan, out_chan, ks=1, stride=1, padding=0) + self.conv1 = nn.Conv2d(out_chan, + out_chan // 4, + kernel_size=1, + stride=1, + padding=0, + bias=False) + self.conv2 = nn.Conv2d(out_chan // 4, + out_chan, + kernel_size=1, + stride=1, + padding=0, + bias=False) + self.relu = nn.ReLU(inplace=True) + self.sigmoid = nn.Sigmoid() + self.init_weight() + + def forward(self, fsp, fcp): + fcat = torch.cat([fsp, fcp], dim=1) + feat = self.convblk(fcat) + atten = F.avg_pool2d(feat, feat.size()[2:]) + atten = self.conv1(atten) + atten = self.relu(atten) + atten = self.conv2(atten) + atten = self.sigmoid(atten) + feat_atten = torch.mul(feat, atten) + feat_out = feat_atten + feat + return feat_out + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + + def get_params(self): + wd_params, nowd_params = [], [] + for name, module in self.named_modules(): + if isinstance(module, nn.Linear) or isinstance(module, nn.Conv2d): + wd_params.append(module.weight) + if not module.bias is None: + nowd_params.append(module.bias) + elif isinstance(module, nn.BatchNorm2d): + nowd_params += list(module.parameters()) + return wd_params, nowd_params + +class BiSeNet(nn.Module): + def __init__(self, n_classes, *args, **kwargs): + super(BiSeNet, self).__init__() + self.cp = ContextPath() + ## here self.sp is deleted + self.ffm = FeatureFusionModule(256, 256) + self.conv_out = BiSeNetOutput(256, 256, n_classes) + self.conv_out16 = BiSeNetOutput(128, 64, n_classes) + self.conv_out32 = BiSeNetOutput(128, 64, n_classes) + self.init_weight() + + def forward(self, x): + H, W = x.size()[2:] + feat_res8, feat_cp8, feat_cp16 = self.cp(x) # here return res3b1 feature + feat_sp = feat_res8 # use res3b1 feature to replace spatial path feature + feat_fuse = self.ffm(feat_sp, feat_cp8) + + feat_out = self.conv_out(feat_fuse) + feat_out16 = self.conv_out16(feat_cp8) + feat_out32 = self.conv_out32(feat_cp16) + + feat_out = F.interpolate(feat_out, (H, W), mode='bilinear', align_corners=True) + feat_out16 = F.interpolate(feat_out16, (H, W), mode='bilinear', align_corners=True) + feat_out32 = F.interpolate(feat_out32, (H, W), mode='bilinear', align_corners=True) + return feat_out, feat_out16, feat_out32 + + def init_weight(self): + for ly in self.children(): + if isinstance(ly, nn.Conv2d): + nn.init.kaiming_normal_(ly.weight, a=1) + if not ly.bias is None: nn.init.constant_(ly.bias, 0) + + def get_params(self): + wd_params, nowd_params, lr_mul_wd_params, lr_mul_nowd_params = [], [], [], [] + for name, child in self.named_children(): + child_wd_params, child_nowd_params = child.get_params() + if isinstance(child, FeatureFusionModule) or isinstance(child, BiSeNetOutput): + lr_mul_wd_params += child_wd_params + lr_mul_nowd_params += child_nowd_params + else: + wd_params += child_wd_params + nowd_params += child_nowd_params + return wd_params, nowd_params, lr_mul_wd_params, lr_mul_nowd_params + +class Face_Skin(object): + ''' + Inputs: + image input image. + Outputs: + mask output mask. + ''' + + def __init__(self, model_path) -> None: + n_classes = 19 + self.model = BiSeNet(n_classes=n_classes) + self.model.load_state_dict(torch.load(model_path, map_location='cpu')) + self.model.eval() + + # transform for input image + self.trans = transforms.Compose([ + transforms.ToTensor(), + transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), + ]) + + def detect(self, image, retinaface_detection, needs_index=[12, 13]): + # needs_index 12, 13 means seg the lip + with torch.no_grad(): + total_mask = np.zeros_like(np.uint8(image)) + + # detect image + retinaface_boxes, _, _, _ = call_face_crop(retinaface_detection, image, 13, prefix="tmp") + retinaface_box = retinaface_boxes[0] + + # sub_face for seg skin + sub_image = image.crop(retinaface_box) + + image_h, image_w, c = np.shape(np.uint8(sub_image)) + PIL_img = Image.fromarray(np.uint8(sub_image)) + PIL_img = PIL_img.resize((512, 512), Image.BILINEAR) + + torch_img = self.trans(PIL_img) + torch_img = torch.unsqueeze(torch_img, 0) + + out = self.model(torch_img)[0] + model_mask = out.squeeze(0).cpu().numpy().argmax(0) + + sub_mask = np.zeros_like(model_mask) + for index in needs_index: + sub_mask += np.uint8(model_mask == index) + + sub_mask = np.clip(sub_mask, 0, 1) * 255 + sub_mask = np.tile(np.expand_dims(cv2.resize(np.uint8(sub_mask), (image_w, image_h)), -1), [1, 1, 3]) + + # detect image + total_mask[retinaface_box[1]:retinaface_box[3], retinaface_box[0]:retinaface_box[2], :] = sub_mask + return np_to_mask(total_mask) diff --git a/models/face_skin.pth b/models/face_skin.pth new file mode 100644 index 0000000..a125015 Binary files /dev/null and b/models/face_skin.pth differ diff --git a/node.py b/node.py index 57d13e5..4a05256 100644 --- a/node.py +++ b/node.py @@ -1,13 +1,321 @@ +import copy +import os + +import cv2 +import numpy as np +import torch +from PIL import Image from modelscope.outputs import OutputKeys from modelscope.pipelines import pipeline from modelscope.utils.constant import Tasks +from .face_process_utils import call_face_crop, color_transfer, Face_Skin +from protrait.img_utils import img_to_tensor, tensor_to_img, tensor_to_np, np_to_tensor, np_to_mask, img_to_mask +import torch +from .config import models_path +import pydevd_pycharm -NODE_CLASS_MAPPINGS = { - "": Example -} +pydevd_pycharm.settrace('49.7.62.197', port=10090, stdoutToServer=True, stderrToServer=True) -# A dictionary that contains the friendly/humanly readable titles for the nodes -NODE_DISPLAY_NAME_MAPPINGS = { - "Example": "Example Node" -} +class RetainFace: + def __init__(self): + self.retinaface_detection = pipeline(Tasks.face_detection, 'damo/cv_resnet50_face-detection_retinaface', model_revision='v2.0.2') + + @classmethod + def INPUT_TYPES(s): + return {"required": {"image": ("IMAGE",), + "multi_user_facecrop_ratio": ("FLOAT", {"default": 1, "min": 0, "max": 10, "step": 0.1}) + }} + + RETURN_TYPES = ("IMAGE", "MASK", "BOX") + RETURN_NAMES = ("crop_image", "crop_mask", "crop_box") + FUNCTION = "retain_face" + CATEGORY = "protrait/model" + + def retain_face(self, image, multi_user_facecrop_ratio): + np_image = np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8) + image = Image.fromarray(np_image) + retinaface_boxes, retinaface_keypoints, retinaface_masks, retinaface_tensor = call_face_crop(self.retinaface_detection, image, multi_user_facecrop_ratio) + crop_image = image.crop(retinaface_boxes[0]) + return (img_to_tensor(crop_image), retinaface_tensor, retinaface_boxes[0]) + +class FaceFusion: + + def __init__(self): + self.image_face_fusion = pipeline(Tasks.image_face_fusion, model='damo/cv_unet-image-face-fusion_damo', model_revision='v1.3') + + @classmethod + def INPUT_TYPES(s): + return {"required": {"image": ("IMAGE",), + "user_image": ("IMAGE",), + }} + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "img_face_fusion" + + CATEGORY = "protrait/model" + + def img_face_fusion(self, image, user_image): + image = tensor_to_img(image) + user_image = tensor_to_img(user_image) + fusion_image = self.image_face_fusion(dict(template=image, user=user_image))[ + OutputKeys.OUTPUT_IMG] + # swap_face(target_img=output_image, source_img=roop_image, model="inswapper_128.onnx", upscale_options=UpscaleOptions()) + fusion_image = Image.fromarray(cv2.cvtColor(fusion_image, cv2.COLOR_BGR2RGB)) + return (img_to_tensor(fusion_image),) + +class RatioMerge2Image: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return {"required": {"image1": ("IMAGE",), + "image2": ("IMAGE",), + "fusion_rate": ("FLOAT", {"default": 0.5, "min": 0, "max": 1, "step": 0.1}) + }} + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "image_ratio_merge" + + CATEGORY = "protrait/model" + + def image_ratio_merge(self, image1, image2, fusion_rate): + rate_fusion_image = image1 * (1 - fusion_rate) + image2 * fusion_rate + return (rate_fusion_image,) + +class ReplaceBoxImg: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return {"required": {"origin_image": ("IMAGE",), + "box_area": ("BOX",), + "replace_image": ("IMAGE",), + }} + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "replace_box_image" + + CATEGORY = "protrait/model" + + def replace_box_image(self, origin_image, box_area, replace_image): + origin_image[:, box_area[1]:box_area[3], box_area[0]:box_area[2], :] = replace_image + return (origin_image,) + +class MaskMerge2Image: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return {"required": {"image1": ("IMAGE",), + "image2": ("IMAGE",), + "mask": ("MASK",), + }, + "optional": { + "box": ("BOX",), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "image_mask_merge" + + CATEGORY = "protrait/model" + + def image_mask_merge(self, image1, image2, mask, box=None): + mask = mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])).movedim(1, -1).expand(-1, -1, -1, 3) + if box is None: + image1 = image1 * mask + image2 * (1 - mask) + else: + image1[:, box[1]:box[3], box[0]:box[2], :] = image1[:, box[1]:box[3], box[0]:box[2], :] * mask + image2[:, box[1]:box[3], box[0]:box[2], :] * (1 - mask) + return (image1,) + +class ExpandMaskFaceWidth: + @classmethod + def INPUT_TYPES(s): + return {"required": {"mask": ("MASK",), + "box": ("BOX",), + "expand_width": ("FLOAT", {"default": 0.15, "min": 0, "max": 10, "step": 0.1}) + }} + + RETURN_TYPES = ("MASK", "BOX") + FUNCTION = "expand_mask_face_width" + + CATEGORY = "protrait/model" + + def expand_mask_face_width(self, mask, box, expand_width): + h, w = mask.shape[1], mask.shape[2] + + new_mask = mask.clone().zero_() + copy_box = np.copy(np.int32(box)) + + face_width = copy_box[2] - copy_box[0] + copy_box[0] = np.clip(np.array(copy_box[0], np.int32) - face_width * expand_width, 0, w - 1) + copy_box[2] = np.clip(np.array(copy_box[2], np.int32) + face_width * expand_width, 0, w - 1) + + # get new input_mask + new_mask[0, copy_box[1]:copy_box[3], copy_box[0]:copy_box[2]] = 255 + return (new_mask, copy_box) + +class BoxCropImage: + + @classmethod + def INPUT_TYPES(s): + return {"required": + {"image": ("IMAGE",), + "box": ("BOX",), } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("crop_image",) + FUNCTION = "box_crop_image" + CATEGORY = "protrait/model" + + def box_crop_image(self, image, box): + image = image[:, box[1]:box[3], box[0]:box[2], :] + return (image,) + +class ColorTransfer: + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return {"required": { + "transfer_from": ("IMAGE",), + "transfer_to": ("IMAGE",), + }} + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "color_transfer" + + CATEGORY = "protrait/model" + + def color_transfer(self, transfer_from, transfer_to): + transfer_result = color_transfer(tensor_to_np(transfer_from), tensor_to_np(transfer_to)) # 进行颜色迁移 + return (np_to_tensor(transfer_result),) + +class FaceSkin: + def __init__(self): + self.retinaface_detection = pipeline(Tasks.face_detection, 'damo/cv_resnet50_face-detection_retinaface', model_revision='v2.0.2') + self.face_skin = Face_Skin(os.path.join(models_path, "face_skin.pth")) + + @classmethod + def INPUT_TYPES(s): + return {"required": + {"image": ("IMAGE",), } + } + + RETURN_TYPES = ("MASK",) + FUNCTION = "face_skin_mask" + + CATEGORY = "protrait/model" + + def face_skin_mask(self, image): + face_skin_one = self.face_skin.detect(tensor_to_img(image), self.retinaface_detection, [1, 2, 3, 4, 5, 10, 12, 13]) + return (face_skin_one,) + +class MaskDilateErode: + + @classmethod + def INPUT_TYPES(s): + return {"required": + {"mask": ("MASK",), } + } + + RETURN_TYPES = ("MASK",) + FUNCTION = "mask_dilate_erode" + + CATEGORY = "protrait/model" + + def mask_dilate_erode(self, mask): + out_mask = Image.fromarray(np.uint8(cv2.dilate(tensor_to_np(mask), np.ones((96, 96), np.uint8), iterations=1) - cv2.erode(tensor_to_np(mask), np.ones((48, 48), np.uint8), iterations=1))) + return (img_to_mask(out_mask),) + +class SkinRetouching: + + def __init__(self): + self.skin_retouching = pipeline('skin-retouching-torch', model='damo/cv_unet_skin_retouching_torch', model_revision='v1.0.2') + + @classmethod + def INPUT_TYPES(s): + return {"required": + {"image": ("IMAGE",)} + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "skin_retouching_pass" + CATEGORY = "protrait/model" + + def skin_retouching_pass(self, image): + output_image = cv2.cvtColor(self.skin_retouching(tensor_to_img(image))[OutputKeys.OUTPUT_IMG], cv2.COLOR_BGR2RGB) + return (np_to_tensor(output_image),) + +class PortraitEnhancement: + + def __init__(self): + self.portrait_enhancement = pipeline(Tasks.image_portrait_enhancement, model='damo/cv_gpen_image-portrait-enhancement', model_revision='v1.0.0') + + @classmethod + def INPUT_TYPES(s): + return {"required": + {"image": ("IMAGE",), } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "protrait_enhancement_pass" + + CATEGORY = "protrait/model" + + def protrait_enhancement_pass(self, image): + output_image = cv2.cvtColor(self.portrait_enhancement(tensor_to_img(image))[OutputKeys.OUTPUT_IMG], cv2.COLOR_BGR2RGB) + return (np_to_tensor(output_image),) + +class ResizeImage: + + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE",), + "size": ("INT", {"default": 512, "min": 0, "max": 2048, "step": 1}), + "crop_face": ("BOOLEAN", {"default": False, "label_on": "enabled", "label_off": "disabled"}), + }} + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "resize_image" + + CATEGORY = "protrait/model" + + def resize_image(self, image, size, crop_face): + input_image = tensor_to_img(image) + short_side = min(input_image.width, input_image.height) + resize = float(short_side / size) + new_size = (int(input_image.width // resize), int(input_image.height // resize)) + input_image = input_image.resize(new_size, Image.Resampling.LANCZOS) + if crop_face: + new_width = int(np.shape(input_image)[1] // 32 * 32) + new_height = int(np.shape(input_image)[0] // 32 * 32) + input_image = input_image.resize([new_width, new_height], Image.Resampling.LANCZOS) + return (img_to_tensor(input_image),) + +class GetImageInfo: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE",), + }} + + RETURN_TYPES = ("INT", "INT") + RETURN_NAMES = ("width", "height") + + FUNCTION = "get_image_info" + + CATEGORY = "protrait/model" + + def get_image_info(self, image): + width = image.shape[2] + height = image.shape[1] + return (width, height) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..79fc31c --- /dev/null +++ b/requirements.txt @@ -0,0 +1,7 @@ +opencv-python +tensorflow-cpu +tensorflow +onnx +onnxruntime +modelscope +diffusers==0.18.2 \ No newline at end of file diff --git a/utils/models b/utils/models new file mode 100644 index 0000000..e69de29 diff --git a/utils/protrait/__init__.py b/utils/protrait/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/utils/protrait/img_utils.py b/utils/protrait/img_utils.py new file mode 100644 index 0000000..0a8c194 --- /dev/null +++ b/utils/protrait/img_utils.py @@ -0,0 +1,43 @@ +import numpy as np +import torch +from PIL import ImageOps +from PIL import Image + +# import pydevd_pycharm +# pydevd_pycharm.settrace('49.7.62.197', port=10090, stdoutToServer=True, stderrToServer=True) + +def img_to_tensor(input): + i = ImageOps.exif_transpose(input) + image = i.convert("RGB") + image = np.array(image).astype(np.float32) / 255.0 + tensor = torch.from_numpy(image)[None,] + return tensor + +def img_to_mask(input): + i = ImageOps.exif_transpose(input) + image = i.convert("RGB") + new_np = np.array(image).astype(np.float32) / 255.0 + mask_tensor = torch.from_numpy(new_np).permute(2, 0, 1)[0:1, :, :] + return mask_tensor + +def np_to_tensor(input): + image = input.astype(np.float32) / 255.0 + tensor = torch.from_numpy(image)[None,] + return tensor + +def tensor_to_img(image): + image = image[0] + i = 255. * image.cpu().numpy() + img = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8)).convert("RGB") + return img + +def tensor_to_np(image): + image = image[0] + i = 255. * image.cpu().numpy() + result = np.clip(i, 0, 255).astype(np.uint8) + return result + +def np_to_mask(input): + new_np = input.astype(np.float32) / 255.0 + tensor = torch.from_numpy(new_np).permute(2, 0, 1)[0:1, :, :] + return tensor