merge v1.1.0

This commit is contained in:
toto
2023-10-24 21:39:54 +08:00
9 changed files with 1107 additions and 69 deletions
+64 -36
View File
@@ -13,8 +13,13 @@ If you have any questions or suggestions, you can reach us through:
- QQ Group: 10419777
- WeChat Group: <img src="./images/wechat.jpg" width="200">
## V1.1.0 Update
1. faceskin adds blur option
2. Add PM_FaceShapMatch node. See node introduction for details.
3. Add PM_MakeUpTransfer node. See node introduction for details.
3. Add a super-resolution model to the PM_PortraitEnhancement node. This super-resolution model can not highlight faces.
## Recent Updates
## V1.0.0 Update
1. Added log for model downloads.
2. Renamed nodes to resolve conflicts with other plugins.
3. Added "roop" model to the Facefusion PM node.
@@ -59,42 +64,65 @@ Click "Load" in the right panel of ComfyUI and select the ./workflow/easyphoto_w
## Node Introduction
* RetainFace PM: Processes images using the pipeline `damo/cv_resnet50_face-detection_retinaface` from Model Scope
* image: Input image
* multi_user_facecrop_ratio: Multiple for extracting the face area
* FaceFusion PM: Fuses two face in the image using the pipeline `damo/cv_unet-image-face-fusion_damo` from Model Scope
* image: Input image
* user_image: Image to be fused
* model: use ali model or roop model for fusion
* RatioMerge2Image PM: Merges two images according to a ratio
* image1: Input image
* Image2: Input image
* fusion_rate: Fusion ratio, maximum is 1, larger values lean towards image1
* MaskMerge2Image PM: Merges images using a mask
* image1: Input image
* image2: Input image
* mask: Mask to be replaced
* ReplaceBoxImg PM: Replaces the image in a box area
* origin_image: Original image
* box_area: Area
* replace_image: Image to be replaced in the area (resolution of box_area and replace_image must match)
* ExpandMaskFaceWidth PM: Proportionally expands the width of the mask
* mask: Input mask
* box: Box corresponding to the mask
* expand_width: Width expansion ratio based on the width of the box
* BoxCropImage PM: Crops images using a box
* ColorTransfer PM: Color transfer for images
* FaceSkin PM: Extracts the mask of the facial part of the image
* MaskDilateErode PM: Dilates and erodes the mask
* SkinRetouching PM: Processes images using the pipeline `damo/cv_gpen_image-portrait-enhancement` from Model Scope
* PortraitEnhancement PM: Processes images using the pipeline `damo/cv_gpen_image-portrait-enhancement` from Model Scope
* ImageResizeTarget PM: Resizes images to a target width and height
* ImageScaleShort PM: Reduces the width and height of the image's shorter side
* image: Input image
* size: Length to be scaled (proportionally scaled based on the shorter side of width and height)
* crop_face: Width and height must be multiples of 32 after scaling
* GetImageInfo PM: Extracts the width and height of the image
* RetainFace PM: Perform matting using models from Model Scope. [Link](https://www.modelscope.cn/models/damo/cv_resnet50_face-detection_retinaface/summary)
* image: Input image
* multi_user_facecrop_ratio: Multiplicative factor for extracting the head region.
* FaceFusion PM: Merge faces from two images.
* image: Input image
* user_image: The image with the face to be merged.
* model: Choose between Ali's model or Roop's model for merging.
* ali: [Link](https://www.modelscope.cn/models/damo/cv_unet-image-face-fusion_damo/summary)
* roop: [Link](https://github.com/deepinsight/insightface)
* RatioMerge2Image PM: Merge two images according to a specified ratio.
* image1: First input image
* image2: Second input image
* fusion_rate: Fusion ratio, ranging from 0 to 1, where higher values favor image1.
* MaskMerge2Image PM: Merge images using a mask.
* image1: First input image
* image2: Second input image
* mask: The mask to be applied for replacement.
* ReplaceBoxImg PM: Replace the image inside a specified box area.
* origin_image: The original image
* box_area: The area to be replaced
* replace_image: The image to replace (ensure the resolution matches box_area)
* ExpandMaskFaceWidth PM: Proportionally expand the width of the mask.
* mask: Input mask
* box: Corresponding box of the mask
* expand_width: The width expansion ratio, based on the box's width.
* BoxCropImage PM: Crop an image using a box.
* ColorTransfer PM: Perform color transfer on images.
* FaceSkin PM: Extract the mask of the facial region from an image.
* MaskDilateErode PM: Dilate and erode masks.
* Skin Retouching PM: Apply skin retouching using the following model.
* [Link](https://www.modelscope.cn/models/damo/cv_unet_skin-retouching/summary)
* Portrait Enhancement PM: Process images using the following model.
* model
* gpen: [Link](https://www.modelscope.cn/models/damo/cv_gpen_image-portrait-enhancement/summary)
* real_gan: [Link](https://www.modelscope.cn/models/bubbliiiing/cv_rrdb_image-super-resolution_x2/summary)
* ImageResizeTarget PM: Resize images to a target width and height.
* ImageScaleShort PM: Reduce the smaller dimension of an image proportionally.
* image: Input image
* size: Desired length for resizing (maintains the aspect ratio)
* crop_face: Ensure the resulting width and height are multiples of 32.
* GetImageInfo PM: Extract the width and height of an image.
* Face Shape Match PM: Apply a certain level of fusion between the diffused image and the original image to reduce differences around the face.
* Makeup Transfer PM: Use a GAN network model to perform makeup transfer.
## Contribution
If you find any issues or have suggestions for improvement, feel free to contribute. Follow these steps:
+20 -9
View File
@@ -15,8 +15,13 @@ English | [简体中文](./README_zh-CN.md)
- QQ 群:10419777
- 微信群: <img src="./images/wechat.jpg" width="200">
## v1.1.0 更新
1. faceskin 增加模糊选项
2. 增加 PM_FaceShapMatch节点 详情查看节点介绍
3. 增加 PM_MakeUpTransfer节点 详情查看节点介绍
3. PM_PortraitEnhancement节点增加一种超分模型,此超分模型可以对人脸不做高光
## 近期更新v1.0.0
## v1.0.0 更新
1. 增加模型下载的log
2. 节点重命名解决与其他插件冲突问题
@@ -24,10 +29,8 @@ English | [简体中文](./README_zh-CN.md)
4. 更新workflow
5. 加速第二次的模型加载
## 正在开发v1.1.0
1. 人脸mask模糊
2. 人脸换妆容
3. 人脸脸型的迁移
## 正在解决
1. 联系 modelscope 解决windows依赖问题
## 安装
@@ -64,13 +67,15 @@ Easyphoto工作位置: [./workflow/easyphoto.json](./workflows/easyphoto.json )
## 节点介绍
* RetainFace PM:使用Model Scope中的pipleline `damo/cv_resnet50_face-detection_retinaface`处理图像
* RetainFace PM:使用Model Scope中的模型进行抠图 [链接](https://www.modelscope.cn/models/damo/cv_resnet50_face-detection_retinaface/summary)
* image:输入图像
* multi_user_facecrop_ratio:提取头像区域的倍数
* FaceFusion PM:使用Model Scope中的pipleline `damo/cv_unet-image-face-fusion_damo`将两张图像的人脸进行融合
* FaceFusion PM:将两张图像的人脸进行融合
* image:输入图像
* user_image:要融合的头像
* model: 使用ali的模型还是roop模型进行融合
* ali:[链接](https://www.modelscope.cn/models/damo/cv_unet-image-face-fusion_damo/summary)
* roop: [链接](https://github.com/deepinsight/insightface)
* RatioMerge2Image PM: 按照比例融合两张图片
* image1:输入的图像
* Image2:输入的图像
@@ -91,14 +96,20 @@ Easyphoto工作位置: [./workflow/easyphoto.json](./workflows/easyphoto.json )
* ColorTransfer PM:对图片进行颜色迁移
* FaceSkin PM:提取图片中人脸的部分的Mask
* MaskDilateErode PM: 对Mask进行膨胀与腐蚀
* SkinRetouching PM:使用Model Scope中的pipleline `damo/cv_gpen_image-portrait-enhancement`处理图像
* PortraitEnhancement PM:使用Model Scope中的pipleline `damo/cv_gpen_image-portrait-enhancement`处理图像
* SkinRetouching PM:使用以下模型进行皮肤美化
* [链接](https://www.modelscope.cn/models/damo/cv_unet_skin-retouching/summary)
* PortraitEnhancement PM:使用以下模型处理图像
* model
* gpen : [链接](https://www.modelscope.cn/models/damo/cv_gpen_image-portrait-enhancement/summary)
* real_gan:[链接](https://www.modelscope.cn/models/bubbliiiing/cv_rrdb_image-super-resolution_x2/summary)
* ImageResizeTarget PM:将图片缩放到目标宽高
* ImageScaleShort PM: 将图片的宽高中小的部分缩减到
* image:输入图像
* size:要缩放的长度(按照宽高中最短的一边进行比例缩放)
* crop_face:缩放后宽高要以32为倍数
* GetImageInfo PM: 提取图片的宽高
* FaceShapMatchPM: 扩散后的图片和原图片进行一定的融合,减少脸旁边的差异
* MakeUpTransferPM: 使用gan网络模型对妆容进行一定的迁移
## 贡献
+5
View File
@@ -1,5 +1,6 @@
import sys
import os
main_path = os.path.dirname(__file__)
sys.path.append(main_path)
@@ -74,6 +75,8 @@ NODE_CLASS_MAPPINGS = {
"PM_ImageScaleShort": ImageScaleShortPM,
"PM_ImageResizeTarget": ImageResizeTargetPM,
"PM_GetImageInfo": GetImageInfoPM,
"PM_MakeUpTransfer": MakeUpTransferPM,
"PM_FaceShapMatch": FaceShapMatchPM,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"PM_RetinaFace": "RetinaFace PM",
@@ -91,6 +94,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PM_ImageScaleShort": "ImageScaleShort PM",
"PM_ImageResizeTarget": "ImageResizeTarget PM",
"PM_GetImageInfo": "GetImageInfo PM",
"PM_MakeUpTransfer": "MakeUpTransfer PM",
"PM_FaceShapMatch":"FaceShapMatch PM"
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+7 -1
View File
@@ -17,7 +17,10 @@ urls = [
"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",
"https://huggingface.co/ezioruan/inswapper_128.onnx/resolve/main/inswapper_128.onnx"
"https://huggingface.co/ezioruan/inswapper_128.onnx/resolve/main/inswapper_128.onnx",
"https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/face_landmarks.pth",
"https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/webui/makeup_transfer.pth",
]
filenames = [
os.path.join(folder_names_and_paths['checkpoints'][0][0], "Chilloutmix-Ni-pruned-fp16-fix.safetensors"),
@@ -32,6 +35,9 @@ filenames = [
os.path.join(folder_names_and_paths['vae'][0][0], "vae-ft-mse-840000-ema-pruned.ckpt"),
os.path.join(models_path, "face_skin.pth"),
os.path.join(models_path, "inswapper_128.onnx"),
os.path.join(models_path, "face_landmarks.pth"),
os.path.join(models_path, "makeup_transfer.pth"),
]
# prompts
validation_prompt = "easyphoto_face, easyphoto, 1person"
+17
View File
@@ -3,6 +3,7 @@ from modelscope.utils.constant import Tasks
import insightface
from insightface.app import FaceAnalysis
from .utils.face_process_utils import Face_Skin
from .utils.psgan_utils import PSGAN_Inference
from .config import *
@@ -13,6 +14,8 @@ face_skin = None
roop = None
skin_retouching = None
portrait_enhancement = None
psgan_interface = None
real_gan_sr = None
def get_retinaface_detection():
global retinaface_detection
@@ -55,3 +58,17 @@ def get_portrait_enhancement():
if portrait_enhancement is None:
portrait_enhancement = pipeline(Tasks.image_portrait_enhancement, model='damo/cv_gpen_image-portrait-enhancement', model_revision='v1.0.0')
return portrait_enhancement
def get_real_gan_sr():
global real_gan_sr
if real_gan_sr is None:
real_gan_sr = pipeline('image-super-resolution-x2', model='bubbliiiing/cv_rrdb_image-super-resolution_x2', model_revision="v1.0.2")
return real_gan_sr
def get_pagan_interface():
global psgan_interface
if psgan_interface is None:
face_landmarks_model_path = os.path.join(models_path, "face_landmarks.pth")
makeup_transfer_model_path = os.path.join(models_path, "makeup_transfer.pth")
psgan_interface = PSGAN_Inference("cuda", makeup_transfer_model_path, get_retinaface_detection(), get_face_skin(), face_landmarks_model_path)
return psgan_interface
+81 -12
View File
@@ -3,7 +3,7 @@ import numpy as np
from PIL import Image
from modelscope.outputs import OutputKeys
from .utils.face_process_utils import call_face_crop, color_transfer, Face_Skin
from .utils.img_utils import img_to_tensor, tensor_to_img, tensor_to_np, np_to_tensor, np_to_mask, img_to_mask
from .utils.img_utils import img_to_tensor, tensor_to_img, tensor_to_np, np_to_tensor, np_to_mask, img_to_mask, img_to_np
from .model_holder import *
# import pydevd_pycharm
@@ -13,7 +13,7 @@ class RetinaFacePM:
@classmethod
def INPUT_TYPES(s):
return {"required": {"image": ("IMAGE",),
"multi_user_facecrop_ratio": ("FLOAT", {"default": 1, "min": 0, "max": 10, "step": 0.1})
"multi_user_facecrop_ratio": ("FLOAT", {"default": 1, "min": 0, "max": 10, "step": 0.01})
}}
RETURN_TYPES = ("IMAGE", "MASK", "BOX")
@@ -189,17 +189,24 @@ class FaceSkinPM:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"image": ("IMAGE",), }
}
{
"image": ("IMAGE",),
"blur_edge": ("BOOLEAN", {"default": False, "label_on": "enabled", "label_off": "disabled"}),
"blur_threshold": ("INT", {"default": 32, "min": 0, "max": 64, "step": 1}),
},
}
RETURN_TYPES = ("MASK",)
FUNCTION = "face_skin_mask"
CATEGORY = "protrait/model"
def face_skin_mask(self, image):
face_skin_one = get_face_skin().detect(tensor_to_img(image), get_retinaface_detection(), [1, 2, 3, 4, 5, 10, 12, 13])
return (face_skin_one,)
def face_skin_mask(self, image, blur_edge, blur_threshold):
face_skin_img = get_face_skin()(tensor_to_img(image), get_retinaface_detection(), [[1, 2, 3, 4, 5, 10, 12, 13]])[0]
face_skin_np = img_to_np(face_skin_img)
if blur_edge:
face_skin_np = cv2.blur(face_skin_np, (blur_threshold, blur_threshold))
return (np_to_mask(face_skin_np),)
class MaskDilateErodePM:
@@ -236,20 +243,25 @@ class SkinRetouchingPM:
class PortraitEnhancementPM:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"image": ("IMAGE",), }
}
{
"image": ("IMAGE",),
"model": (["pgen", "real_gan"],),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "protrait_enhancement_pass"
CATEGORY = "protrait/model"
def protrait_enhancement_pass(self, image):
output_image = cv2.cvtColor(get_portrait_enhancement()(tensor_to_img(image))[OutputKeys.OUTPUT_IMG], cv2.COLOR_BGR2RGB)
def protrait_enhancement_pass(self, image, model):
if model == "pgen":
output_image = cv2.cvtColor(get_portrait_enhancement()(tensor_to_img(image))[OutputKeys.OUTPUT_IMG], cv2.COLOR_BGR2RGB)
elif model == "real_gan":
output_image = cv2.cvtColor(get_real_gan_sr()(tensor_to_img(image))[OutputKeys.OUTPUT_IMG], cv2.COLOR_BGR2RGB)
return (np_to_tensor(output_image),)
class ImageScaleShortPM:
@@ -317,3 +329,60 @@ class GetImageInfoPM:
width = image.shape[2]
height = image.shape[1]
return (width, height)
class MakeUpTransferPM:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"source_image": ("IMAGE",),
"makeup_image": ("IMAGE",),
}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "makeup_transfer"
CATEGORY = "protrait/model"
def makeup_transfer(self, source_image, makeup_image):
source_image = tensor_to_img(source_image).resize([256, 256])
makeup_image = tensor_to_img(makeup_image).resize([256, 256])
result = get_pagan_interface().transfer(source_image, makeup_image)
return (img_to_tensor(result),)
class FaceShapMatchPM:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"source_image": ("IMAGE",),
"match_image": ("IMAGE",),
"face_box": ("BOX",),
}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "faceshap_match"
CATEGORY = "protrait/model"
def faceshap_match(self, source_image, match_image, face_box):
# detect face area
source_image = tensor_to_img(source_image)
match_image = tensor_to_img(match_image)
face_skin_mask = get_face_skin()(source_image, get_retinaface_detection(), needs_index=[[1, 2, 3, 4, 5, 7, 8, 10, 11, 12, 13]])[0]
face_width = face_box[2] - face_box[0]
kernel_size = np.ones((int(face_width // 10), int(face_width // 10)), np.uint8)
# Fill small holes with a close operation
face_skin_mask = Image.fromarray(np.uint8(cv2.morphologyEx(np.array(face_skin_mask), cv2.MORPH_CLOSE, kernel_size)))
# Use dilate to reconstruct the surrounding area of the face
face_skin_mask = Image.fromarray(np.uint8(cv2.dilate(np.array(face_skin_mask), kernel_size, iterations=1)))
face_skin_mask = cv2.blur(np.float32(face_skin_mask), (32, 32)) / 255
# paste back to photo, Using I2I generation controlled solely by OpenPose, even with a very small denoise amplitude,
# still carries the risk of introducing NSFW and global incoherence.!!! important!!!
input_image_uint8 = np.array(source_image) * face_skin_mask + np.array(match_image) * (1 - face_skin_mask)
return (np_to_tensor(input_image_uint8),)
+23 -11
View File
@@ -496,19 +496,27 @@ class Face_Skin(object):
self.model.load_state_dict(torch.load(model_path, map_location='cpu'))
self.model.eval()
self.cuda = torch.cuda.is_available()
if self.cuda:
self.model.cuda()
# 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]):
# index => label
# 1:'skin', 2:'left_brow', 3:'right_brow', 4:'left_eye', 5:'right_eye', 6:'eye_g', 7:'left_ear', 8:'right_ear',
# 9:'ear_r', 10:'nose', 11:'mouth', 12:'upper_lip', 13:'low_lip', 14:'neck', 15:'neck_l', 16:'cloth',
# 17:'hair', 18:'hat'
def __call__(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_boxes, _, _, _ = call_face_crop(retinaface_detection, image, 1.5, prefix="tmp")
retinaface_box = retinaface_boxes[0]
# sub_face for seg skin
@@ -520,17 +528,21 @@ class Face_Skin(object):
torch_img = self.trans(PIL_img)
torch_img = torch.unsqueeze(torch_img, 0)
if self.cuda:
torch_img = torch_img.cuda()
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)
masks = []
for _needs_index in needs_index:
total_mask = np.zeros_like(np.uint8(image))
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])
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])
total_mask[retinaface_box[1]:retinaface_box[3], retinaface_box[0]:retinaface_box[2], :] = sub_mask
masks.append(Image.fromarray(np.uint8(total_mask)))
# detect image
total_mask[retinaface_box[1]:retinaface_box[3], retinaface_box[0]:retinaface_box[2], :] = sub_mask
return np_to_mask(total_mask)
return masks
+6
View File
@@ -13,6 +13,12 @@ def img_to_tensor(input):
tensor = torch.from_numpy(image)[None,]
return tensor
def img_to_np(input):
i = ImageOps.exif_transpose(input)
image = i.convert("RGB")
image_np = np.array(image).astype(np.float32)
return image_np
def img_to_mask(input):
i = ImageOps.exif_transpose(input)
image = i.convert("RGB")
+884
View File
@@ -0,0 +1,884 @@
#!/usr/bin/python
# -*- encoding: utf-8 -*-
import math
import os.path as osp
from collections import OrderedDict
import cv2
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
from torch import nn
from torch.autograd import Variable
from torch.nn import Parameter, functional
from torchvision import transforms
from torchvision.transforms import ToPILImage
pwd = osp.split(osp.realpath(__file__))[0]
# Preprocess part
def to_var(x, requires_grad=True):
if requires_grad:
return Variable(x).float()
else:
return Variable(x, requires_grad=requires_grad).float()
def copy_area(tar, src, lms):
rect = [int(min(lms[:, 1])) - PreProcess.eye_margin,
int(min(lms[:, 0])) - PreProcess.eye_margin,
int(max(lms[:, 1])) + PreProcess.eye_margin + 1,
int(max(lms[:, 0])) + PreProcess.eye_margin + 1]
tar[:, :, rect[1]:rect[3], rect[0]:rect[2]] = \
src[:, :, rect[1]:rect[3], rect[0]:rect[2]]
src[:, :, rect[1]:rect[3], rect[0]:rect[2]] = 0
class rectangle():
def __init__(self, left, top, right, bottom):
self.left_num = left
self.top_num = top
self.right_num = right
self.bottom_num = bottom
def left(self):
return self.left_num
def top(self):
return self.top_num
def right(self):
return self.right_num
def bottom(self):
return self.bottom_num
def height(self):
return self.bottom_num - self.top_num
def width(self):
return self.right_num - self.left_num
def crop(image: Image, face, up_ratio, down_ratio, width_ratio) -> (Image, 'face'):
width, height = image.size
face_height = face.height()
face_width = face.width()
delta_up = up_ratio * face_height
delta_down = down_ratio * face_height
delta_width = width_ratio * width
img_left = int(max(0, face.left() - delta_width))
img_top = int(max(0, face.top() - delta_up))
img_right = int(min(width, face.right() + delta_width))
img_bottom = int(min(height, face.bottom() + delta_down))
image = image.crop((img_left, img_top, img_right, img_bottom))
face = rectangle(face.left() - img_left, face.top() - img_top,
face.right() - img_left, face.bottom() - img_top)
center = [(img_right - img_left) / 2, (img_bottom - img_top) / 2]
width, height = image.size
# import ipdb; ipdb.set_trace()
crop_left = img_left
crop_top = img_top
crop_right = img_right
crop_bottom = img_bottom
if width > height:
left = int(center[0] - height / 2)
right = int(center[0] + height / 2)
if left < 0:
left, right = 0, height
elif right > width:
left, right = width - height, width
image = image.crop((left, 0, right, height))
face = rectangle(face.left() - left, face.top(),
face.right() - left, face.bottom())
crop_left += left
crop_right = crop_left + height
elif width < height:
top = int(center[1] - width / 2)
bottom = int(center[1] + width / 2)
if top < 0:
top, bottom = 0, width
elif bottom > height:
top, bottom = height - width, height
image = image.crop((0, top, width, bottom))
face = rectangle(face.left(), face.top() - top,
face.right(), face.bottom() - top)
crop_top += top
crop_bottom = crop_top + width
crop_face = rectangle(crop_left, crop_top, crop_right, crop_bottom)
return image, face, crop_face
class FaceParser:
def __init__(self, device="cpu", face_skin=None):
mapper = [0, 1, 2, 3, 4, 5, 0, 11, 12, 0, 6, 8, 7, 9, 13, 0, 0, 10, 0]
self.device = device
self.dic = torch.tensor(mapper, device=device)
self.net = face_skin.model
self.to_tensor = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
])
def parse(self, image: Image):
assert image.shape[:2] == (512, 512)
with torch.no_grad():
image = self.to_tensor(image).to(self.device)
image = torch.unsqueeze(image, 0)
out = self.net(image)[0]
parsing = out.squeeze(0).argmax(0)
mask = torch.zeros_like(parsing)
for index, num in enumerate(self.dic):
mask[parsing == index] = num
return mask.float()
class WingLoss(nn.Module):
def __init__(self, wing_w=10.0, wing_epsilon=2.0):
super(WingLoss, self).__init__()
self.wing_w = wing_w
self.wing_epsilon = wing_epsilon
self.wing_c = self.wing_w * (1.0 - math.log(1.0 + self.wing_w / self.wing_epsilon))
def forward(self, targets, predictions, euler_angle_weights=None):
abs_error = torch.abs(targets - predictions)
loss = torch.where(torch.le(abs_error, self.wing_w),
self.wing_w * torch.log(1.0 + abs_error / self.wing_epsilon), abs_error - self.wing_c)
loss_sum = torch.sum(loss, 1)
if euler_angle_weights is not None:
loss_sum *= euler_angle_weights
return torch.mean(loss_sum)
class LinearBottleneck(nn.Module):
def __init__(self, input_channels, out_channels, expansion, stride=1, activation=nn.ReLU6):
super(LinearBottleneck, self).__init__()
self.expansion_channels = input_channels * expansion
self.conv1 = nn.Conv2d(input_channels, self.expansion_channels, stride=1, kernel_size=1)
self.bn1 = nn.BatchNorm2d(self.expansion_channels)
self.depth_conv2 = nn.Conv2d(self.expansion_channels, self.expansion_channels, stride=stride, kernel_size=3,
groups=self.expansion_channels, padding=1)
self.bn2 = nn.BatchNorm2d(self.expansion_channels)
self.conv3 = nn.Conv2d(self.expansion_channels, out_channels, stride=1, kernel_size=1)
self.bn3 = nn.BatchNorm2d(out_channels)
self.activation = activation(inplace=True) # inplace=True
self.stride = stride
self.input_channels = input_channels
self.out_channels = out_channels
def forward(self, input):
residual = input
out = self.conv1(input)
out = self.bn1(out)
# out = self.activation(out)
out = self.depth_conv2(out)
out = self.bn2(out)
out = self.activation(out)
out = self.conv3(out)
out = self.bn3(out)
if self.stride == 1 and self.input_channels == self.out_channels:
out += residual
return out
class AuxiliaryNet(nn.Module):
def __init__(self, input_channels, nums_class=3, activation=nn.ReLU, first_conv_stride=2):
super(AuxiliaryNet, self).__init__()
self.input_channels = input_channels
# self.num_channels = [128, 128, 32, 128, 32]
self.num_channels = [512, 512, 512, 512, 1024]
self.conv1 = nn.Conv2d(self.input_channels, self.num_channels[0], kernel_size=3, stride=first_conv_stride,
padding=1)
self.bn1 = nn.BatchNorm2d(self.num_channels[0])
self.conv2 = nn.Conv2d(self.num_channels[0], self.num_channels[1], kernel_size=3, stride=1, padding=1)
self.bn2 = nn.BatchNorm2d(self.num_channels[1])
self.conv3 = nn.Conv2d(self.num_channels[1], self.num_channels[2], kernel_size=3, stride=2, padding=1)
self.bn3 = nn.BatchNorm2d(self.num_channels[2])
self.conv4 = nn.Conv2d(self.num_channels[2], self.num_channels[3], kernel_size=7, stride=1, padding=3)
self.bn4 = nn.BatchNorm2d(self.num_channels[3])
self.fc1 = nn.Linear(in_features=self.num_channels[3], out_features=self.num_channels[4])
self.fc2 = nn.Linear(in_features=self.num_channels[4], out_features=nums_class)
self.activation = activation(inplace=True)
self.init_params()
def init_params(self):
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode='fan_out')
if m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.BatchNorm2d):
nn.init.constant_(m.weight, 1)
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.01)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def forward(self, input):
out = self.conv1(input)
out = self.bn1(out)
out = self.activation(out)
out = self.conv2(out)
out = self.bn2(out)
out = self.activation(out)
out = self.conv3(out)
out = self.bn3(out)
out = self.activation(out)
out = self.conv4(out)
out = self.bn4(out)
out = self.activation(out)
out = functional.adaptive_avg_pool2d(out, 1).squeeze(-1).squeeze(-1)
out = self.fc1(out)
euler_angles_pre = self.fc2(out)
return euler_angles_pre
class MobileNetV2(nn.Module):
def __init__(self, input_channels=3, num_of_channels=None, nums_class=136, activation=nn.ReLU6):
super(MobileNetV2, self).__init__()
assert num_of_channels is not None
self.num_of_channels = num_of_channels
self.conv1 = nn.Conv2d(input_channels, self.num_of_channels[0], kernel_size=3, stride=2, padding=1)
self.bn1 = nn.BatchNorm2d(self.num_of_channels[0])
self.depth_conv2 = nn.Conv2d(self.num_of_channels[0], self.num_of_channels[0], kernel_size=3, stride=1,
padding=1, groups=self.num_of_channels[0])
self.bn2 = nn.BatchNorm2d(self.num_of_channels[0])
self.stage0 = self.make_stage(self.num_of_channels[0], self.num_of_channels[0], stride=2, stage=0, times=5,
expansion=2, activation=activation)
self.stage1 = self.make_stage(self.num_of_channels[0], self.num_of_channels[1], stride=2, stage=1, times=7,
expansion=4, activation=activation)
self.linear_bottleneck_end = nn.Sequential(LinearBottleneck(self.num_of_channels[1], self.num_of_channels[2],
expansion=2, stride=1, activation=activation))
self.conv3 = nn.Conv2d(self.num_of_channels[2], self.num_of_channels[3], kernel_size=3, stride=2, padding=1)
self.bn3 = nn.BatchNorm2d(self.num_of_channels[3])
self.conv4 = nn.Conv2d(self.num_of_channels[3], self.num_of_channels[4], kernel_size=7, stride=1)
self.bn4 = nn.BatchNorm2d(self.num_of_channels[4])
self.activation = activation(inplace=True)
self.in_features = 14 * 14 * self.num_of_channels[2] + 7 * 7 * self.num_of_channels[3] + 1 * 1 * self.num_of_channels[4]
self.fc = nn.Linear(in_features=self.in_features, out_features=nums_class)
self.init_params()
def init_params(self):
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode='fan_out')
if m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.BatchNorm2d):
nn.init.constant_(m.weight, 1)
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.01)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def make_stage(self, input_channels, out_channels, stride, stage, times, expansion, activation=nn.ReLU6):
modules = OrderedDict()
stage_name = 'LinearBottleneck{}'.format(stage)
module = LinearBottleneck(input_channels, out_channels, expansion=2,
stride=stride, activation=activation)
modules[stage_name + '_0'] = module
for i in range(times - 1):
module = LinearBottleneck(out_channels, out_channels, expansion=expansion, stride=1,
activation=activation)
module_name = stage_name + '_{}'.format(i + 1)
modules[module_name] = module
return nn.Sequential(modules)
def forward(self, input):
with torch.no_grad():
out = self.conv1(input)
out = self.bn1(out)
out = self.activation(out)
out = self.depth_conv2(out)
out = self.bn2(out)
out = self.activation(out)
out = self.stage0(out)
out1 = self.stage1(out)
out1 = self.linear_bottleneck_end(out1)
out2 = self.conv3(out1)
out2 = self.bn3(out2)
out2 = self.activation(out2)
out3 = self.conv4(out2)
out3 = self.bn4(out3)
out3 = self.activation(out3)
out1 = out1.contiguous().view(out1.size(0), -1)
out2 = out2.contiguous().view(out2.size(0), -1)
out3 = out3.contiguous().view(out3.size(0), -1)
multi_scale = torch.cat([out1, out2, out3], 1)
assert multi_scale.size(1) == self.in_features
pre_landmarks = self.fc(multi_scale)
return pre_landmarks, out
class PreProcess:
eye_margin = 16
diff_size = (64, 64)
def __init__(self, device="cpu", need_parser=True, retinaface_detection=None, face_skin=None, landmark_path=None):
self.device = device
self.img_size = 256
xs, ys = np.meshgrid(
np.linspace(
0, self.img_size - 1,
self.img_size
),
np.linspace(
0, self.img_size - 1,
self.img_size
)
)
xs = xs[None].repeat(68, axis=0)
ys = ys[None].repeat(68, axis=0)
fix = np.concatenate([ys, xs], axis=0)
self.fix = torch.Tensor(fix).to(self.device)
self.retinaface_detection = retinaface_detection
if need_parser:
self.face_parse = FaceParser(device=device, face_skin=face_skin)
self.landmark = MobileNetV2(num_of_channels=[64, 128, 16, 32, 128], nums_class=136)
self.landmark.load_state_dict(torch.load(landmark_path))
self.landmark.eval().to(self.device)
self.up_ratio = 0.6 / 0.85
self.down_ratio = 0.2 / 0.85
self.width_ratio = 0.2 / 0.85
self.lip_class = [7, 9]
self.face_class = [1, 6]
self.transform = transforms.Compose(
[
transforms.ToTensor(),
transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5])
]
)
def relative2absolute(self, lms):
return lms * self.img_size
def process(self, mask, lms, device="cpu"):
diff = to_var(
(self.fix.double() - torch.tensor(lms.transpose((1, 0)
).reshape(-1, 1, 1)).to(self.device)
).unsqueeze(0), requires_grad=False).to(self.device)
lms_eye_left = lms[42:48]
lms_eye_right = lms[36:42]
lms = lms.transpose((1, 0)).reshape(-1, 1, 1) # transpose to (y-x)
# lms = np.tile(lms, (1, 256, 256)) # (136, h, w)
diff = to_var((self.fix.double() - torch.tensor(lms).to(self.device)).unsqueeze(0), requires_grad=False).to(self.device)
mask_lip = (mask == self.lip_class[0]).float() + (mask == self.lip_class[1]).float()
mask_face = (mask == self.face_class[0]).float() + (mask == self.face_class[1]).float()
mask_eyes = torch.zeros_like(mask, device=device)
copy_area(mask_eyes, mask_face, lms_eye_left)
copy_area(mask_eyes, mask_face, lms_eye_right)
mask_eyes = to_var(mask_eyes, requires_grad=False).to(device)
mask_list = [mask_lip, mask_face, mask_eyes]
mask_aug = torch.cat(mask_list, 0) # (3, 1, h, w)
mask_re = F.interpolate(mask_aug, size=self.diff_size).repeat(1, diff.shape[1], 1, 1) # (3, 136, 64, 64)
diff_re = F.interpolate(diff, size=self.diff_size).repeat(3, 1, 1, 1) # (3, 136, 64, 64)
diff_re = diff_re * mask_re # (3, 136, 32, 32)
norm = torch.norm(diff_re, dim=1, keepdim=True).repeat(1, diff_re.shape[1], 1, 1)
norm = torch.where(norm == 0, torch.tensor(1e10, device=device), norm)
diff_re /= norm
return mask_aug, diff_re
def __call__(self, image: Image):
retinaface_result = self.retinaface_detection(image)
face = []
for box in retinaface_result['boxes']:
face.append(rectangle(*np.int32(box)))
if len(face) == 0:
return None, None, None
face_on_image = face[0]
image, face, crop_face = crop(image, face_on_image, self.up_ratio, self.down_ratio, self.width_ratio)
np_image = np.array(image)
mask = self.face_parse.parse(cv2.resize(np_image, (512, 512)))
# obtain face parsing result
mask = F.interpolate(
mask.view(1, 1, 512, 512),
(self.img_size, self.img_size),
mode="nearest")
mask = mask.type(torch.uint8)
mask = to_var(mask, requires_grad=False).to(self.device)
input = image.crop([face.left(), face.top(), face.right(), face.bottom()])
input = input.resize([112, 112])
input = np.expand_dims(np.array(input, np.float32) / 255.0, 0)
input = torch.Tensor(input.transpose((0, 3, 1, 2))).to(self.device)
pre_landmarks, _ = self.landmark(input)
lms = pre_landmarks[0].cpu().detach().numpy()
lms = lms.reshape(-1, 2) * [face.width(), face.height()] + np.int32([face.left(), face.top()])
lms = lms / [np.shape(image)[0], np.shape(image)[1]] * self.img_size
lms = lms[:, ::-1]
mask, diff = self.process(mask, lms, device=self.device)
image = image.resize((self.img_size, self.img_size), Image.ANTIALIAS)
image = self.transform(image)
real = to_var(image.unsqueeze(0))
return [real, mask, diff], face_on_image, crop_face
# Solver part (GAN part)
def l2normalize(v, eps=1e-12):
return v / (v.norm() + eps)
class SpectralNorm(object):
def __init__(self):
self.name = "weight"
self.power_iterations = 1
def compute_weight(self, module):
u = getattr(module, self.name + "_u")
v = getattr(module, self.name + "_v")
w = getattr(module, self.name + "_bar")
height = w.data.shape[0]
for _ in range(self.power_iterations):
v.data = l2normalize(torch.mv(torch.t(w.view(height, -1).data), u.data))
u.data = l2normalize(torch.mv(w.view(height, -1).data, v.data))
# sigma = torch.dot(u.data, torch.mv(w.view(height,-1).data, v.data))
sigma = u.dot(w.view(height, -1).mv(v))
return w / sigma.expand_as(w)
@staticmethod
def apply(module):
name = "weight"
fn = SpectralNorm()
try:
u = getattr(module, name + "_u")
v = getattr(module, name + "_v")
w = getattr(module, name + "_bar")
except AttributeError:
w = getattr(module, name)
height = w.data.shape[0]
width = w.view(height, -1).data.shape[1]
u = Parameter(w.data.new(height).normal_(0, 1), requires_grad=False)
v = Parameter(w.data.new(width).normal_(0, 1), requires_grad=False)
w_bar = Parameter(w.data)
# del module._parameters[name]
module.register_parameter(name + "_u", u)
module.register_parameter(name + "_v", v)
module.register_parameter(name + "_bar", w_bar)
# remove w from parameter list
del module._parameters[name]
setattr(module, name, fn.compute_weight(module))
# recompute weight before every forward()
module.register_forward_pre_hook(fn)
return fn
def remove(self, module):
weight = self.compute_weight(module)
delattr(module, self.name)
del module._parameters[self.name + '_u']
del module._parameters[self.name + '_v']
del module._parameters[self.name + '_bar']
module.register_parameter(self.name, Parameter(weight.data))
def __call__(self, module, inputs):
setattr(module, self.name, self.compute_weight(module))
def spectral_norm(module):
SpectralNorm.apply(module)
return module
def remove_spectral_norm(module):
name = 'weight'
for k, hook in module._forward_pre_hooks.items():
if isinstance(hook, SpectralNorm) and hook.name == name:
hook.remove(module)
del module._forward_pre_hooks[k]
return module
raise ValueError("spectral_norm of '{}' not found in {}"
.format(name, module))
# Defines the GAN loss which uses either LSGAN or the regular GAN.
# When LSGAN is used, it is basically same as MSELoss,
# but it abstracts away the need to create the target label tensor
# that has the same size as the input
class ResidualBlock(nn.Module):
"""Residual Block."""
def __init__(self, dim_in, dim_out, net_mode=None):
if net_mode == 'p' or (net_mode is None):
use_affine = True
elif net_mode == 't':
use_affine = False
super(ResidualBlock, self).__init__()
self.main = nn.Sequential(
nn.Conv2d(dim_in, dim_out, kernel_size=3, stride=1, padding=1, bias=False),
nn.InstanceNorm2d(dim_out, affine=use_affine),
nn.ReLU(inplace=True),
nn.Conv2d(dim_out, dim_out, kernel_size=3, stride=1, padding=1, bias=False),
nn.InstanceNorm2d(dim_out, affine=use_affine)
)
def forward(self, x):
return x + self.main(x)
class GetMatrix(nn.Module):
def __init__(self, dim_in, dim_out):
super(GetMatrix, self).__init__()
self.get_gamma = nn.Conv2d(dim_in, dim_out, kernel_size=1, stride=1, padding=0, bias=False)
self.get_beta = nn.Conv2d(dim_in, dim_out, kernel_size=1, stride=1, padding=0, bias=False)
def forward(self, x):
gamma = self.get_gamma(x)
beta = self.get_beta(x)
return x, gamma, beta
class NONLocalBlock2D(nn.Module):
def __init__(self):
super(NONLocalBlock2D, self).__init__()
self.g = nn.Conv2d(in_channels=1, out_channels=1,
kernel_size=1, stride=1, padding=0)
def forward(self, source, weight):
"""(b, c, h, w)
src_diff: (3, 136, 32, 32)
"""
batch_size = source.size(0)
g_source = source.view(batch_size, 1, -1) # (N, C, H*W)
g_source = g_source.permute(0, 2, 1) # (N, H*W, C)
y = torch.bmm(weight.to_dense(), g_source)
y = y.permute(0, 2, 1).contiguous() # (N, C, H*W)
y = y.view(batch_size, 1, *source.size()[2:])
return y
class Generator(nn.Module):
"""Generator. Encoder-Decoder Architecture."""
def __init__(self):
super(Generator, self).__init__()
# -------------------------- PNet(MDNet) for obtaining makeup matrices --------------------------
layers = nn.Sequential(
nn.Conv2d(3, 64, kernel_size=7, stride=1, padding=3, bias=False),
nn.InstanceNorm2d(64, affine=True),
nn.ReLU(inplace=True)
)
self.pnet_in = layers
# Down-Sampling
curr_dim = 64
for i in range(2):
layers = nn.Sequential(
nn.Conv2d(curr_dim, curr_dim * 2, kernel_size=4, stride=2, padding=1, bias=False),
nn.InstanceNorm2d(curr_dim * 2, affine=True),
nn.ReLU(inplace=True),
)
setattr(self, f'pnet_down_{i + 1}', layers)
curr_dim = curr_dim * 2
# Bottleneck. All bottlenecks share the same attention module
self.atten_bottleneck_g = NONLocalBlock2D()
self.atten_bottleneck_b = NONLocalBlock2D()
self.simple_spade = GetMatrix(curr_dim, 1) # get the makeup matrix
for i in range(3):
setattr(self, f'pnet_bottleneck_{i + 1}', ResidualBlock(dim_in=curr_dim, dim_out=curr_dim, net_mode='p'))
# --------------------------- TNet(MANet) for applying makeup transfer ----------------------------
self.tnet_in_conv = nn.Conv2d(3, 64, kernel_size=7, stride=1, padding=3, bias=False)
self.tnet_in_spade = nn.InstanceNorm2d(64, affine=False)
self.tnet_in_relu = nn.ReLU(inplace=True)
# Down-Sampling
curr_dim = 64
for i in range(2):
setattr(self, f'tnet_down_conv_{i + 1}', nn.Conv2d(curr_dim, curr_dim * 2, kernel_size=4, stride=2, padding=1, bias=False))
setattr(self, f'tnet_down_spade_{i + 1}', nn.InstanceNorm2d(curr_dim * 2, affine=False))
setattr(self, f'tnet_down_relu_{i + 1}', nn.ReLU(inplace=True))
curr_dim = curr_dim * 2
# Bottleneck
for i in range(6):
setattr(self, f'tnet_bottleneck_{i + 1}', ResidualBlock(dim_in=curr_dim, dim_out=curr_dim, net_mode='t'))
# Up-Sampling
for i in range(2):
setattr(self, f'tnet_up_conv_{i + 1}', nn.ConvTranspose2d(curr_dim, curr_dim // 2, kernel_size=4, stride=2, padding=1, bias=False))
setattr(self, f'tnet_up_spade_{i + 1}', nn.InstanceNorm2d(curr_dim // 2, affine=False))
setattr(self, f'tnet_up_relu_{i + 1}', nn.ReLU(inplace=True))
curr_dim = curr_dim // 2
layers = nn.Sequential(
nn.Conv2d(curr_dim, 3, kernel_size=7, stride=1, padding=3, bias=False),
nn.Tanh()
)
self.tnet_out = layers
@staticmethod
def atten_feature(mask_s, weight, gamma_s, beta_s, atten_module_g, atten_module_b):
"""
feature size: (1, c, h, w)
mask_c(s): (3, 1, h, w)
diff_c: (1, 138, 256, 256)
return: (1, c, h, w)
"""
channel_num = gamma_s.shape[1]
mask_s_re = F.interpolate(mask_s, size=gamma_s.shape[2:]).repeat(1, channel_num, 1, 1)
gamma_s_re = gamma_s.repeat(3, 1, 1, 1)
gamma_s = gamma_s_re * mask_s_re # (3, c, h, w)
beta_s_re = beta_s.repeat(3, 1, 1, 1)
beta_s = beta_s_re * mask_s_re
gamma = atten_module_g(gamma_s, weight) # (3, c, h, w)
beta = atten_module_b(beta_s, weight)
gamma = (gamma[0] + gamma[1] + gamma[2]).unsqueeze(0) # (c, h, w) combine the three parts
beta = (beta[0] + beta[1] + beta[2]).unsqueeze(0)
return gamma, beta
def get_weight(self, mask_c, mask_s, fea_c, fea_s, diff_c, diff_s):
""" s --> source; c --> target
feature size: (1, 256, 64, 64)
diff: (3, 136, 32, 32)
"""
HW = 64 * 64
batch_size = 3
assert fea_s is not None # fea_s when i==3
# get 3 part fea using mask
channel_num = fea_s.shape[1]
mask_c_re = F.interpolate(mask_c, size=64).repeat(1, channel_num, 1, 1) # (3, c, h, w)
fea_c = fea_c.repeat(3, 1, 1, 1) # (3, c, h, w)
fea_c = fea_c * mask_c_re # (3, c, h, w) 3 stands for 3 parts
mask_s_re = F.interpolate(mask_s, size=64).repeat(1, channel_num, 1, 1)
fea_s = fea_s.repeat(3, 1, 1, 1)
fea_s = fea_s * mask_s_re
theta_input = torch.cat((fea_c * 0.01, diff_c), dim=1)
phi_input = torch.cat((fea_s * 0.01, diff_s), dim=1)
theta_target = theta_input.view(batch_size, -1, HW) # (N, C+136, H*W)
theta_target = theta_target.permute(0, 2, 1) # (N, H*W, C+136)
phi_source = phi_input.view(batch_size, -1, HW) # (N, C+136, H*W)
weight = torch.bmm(theta_target, phi_source) # (3, HW, HW)
with torch.no_grad():
v = weight.detach().nonzero().long().permute(1, 0)
# This clone is required to correctly release cuda memory.
weight_ind = v.clone()
del v
torch.cuda.empty_cache()
weight *= 200 # hyper parameters for visual feature
weight = F.softmax(weight, dim=-1)
weight = weight[weight_ind[0], weight_ind[1], weight_ind[2]]
ret = torch.sparse.FloatTensor(weight_ind, weight, torch.Size([3, HW, HW]))
return ret
def forward(self, c, s, mask_c, mask_s, diff_c, diff_s, gamma=None, beta=None, ret=False):
c, s, mask_c, mask_s, diff_c, diff_s = [x.squeeze(0) if x.ndim == 5 else x for x in [c, s, mask_c, mask_s, diff_c, diff_s]]
"""attention version
c: content, stands for source image. shape: (b, c, h, w)
s: style, stands for reference image. shape: (b, c, h, w)
mask_list_c: lip, skin, eye. (b, 1, h, w)
"""
# forward c in tnet(MANet)
c_tnet = self.tnet_in_conv(c)
s = self.pnet_in(s)
c_tnet = self.tnet_in_spade(c_tnet)
c_tnet = self.tnet_in_relu(c_tnet)
# down-sampling
for i in range(2):
if gamma is None:
cur_pnet_down = getattr(self, f'pnet_down_{i + 1}')
s = cur_pnet_down(s)
cur_tnet_down_conv = getattr(self, f'tnet_down_conv_{i + 1}')
cur_tnet_down_spade = getattr(self, f'tnet_down_spade_{i + 1}')
cur_tnet_down_relu = getattr(self, f'tnet_down_relu_{i + 1}')
c_tnet = cur_tnet_down_conv(c_tnet)
c_tnet = cur_tnet_down_spade(c_tnet)
c_tnet = cur_tnet_down_relu(c_tnet)
# bottleneck
for i in range(6):
if gamma is None and i <= 2:
cur_pnet_bottleneck = getattr(self, f'pnet_bottleneck_{i + 1}')
cur_tnet_bottleneck = getattr(self, f'tnet_bottleneck_{i + 1}')
# get s_pnet from p and transform
if i == 3:
if gamma is None: # not in test_mix
s, gamma, beta = self.simple_spade(s)
weight = self.get_weight(mask_c, mask_s, c_tnet, s, diff_c, diff_s)
gamma, beta = self.atten_feature(mask_s, weight, gamma, beta, self.atten_bottleneck_g, self.atten_bottleneck_b)
if ret:
return [gamma, beta]
# else: # in test mode
# gamma, beta = param_A[0]*w + param_B[0]*(1-w), param_A[1]*w + param_B[1]*(1-w)
c_tnet = c_tnet * (1 + gamma) + beta # apply makeup transfer using makeup matrices
if gamma is None and i <= 2:
s = cur_pnet_bottleneck(s)
c_tnet = cur_tnet_bottleneck(c_tnet)
# up-sampling
for i in range(2):
cur_tnet_up_conv = getattr(self, f'tnet_up_conv_{i + 1}')
cur_tnet_up_spade = getattr(self, f'tnet_up_spade_{i + 1}')
cur_tnet_up_relu = getattr(self, f'tnet_up_relu_{i + 1}')
c_tnet = cur_tnet_up_conv(c_tnet)
c_tnet = cur_tnet_up_spade(c_tnet)
c_tnet = cur_tnet_up_relu(c_tnet)
c_tnet = self.tnet_out(c_tnet)
return c_tnet
# Gan Solver
class Solver():
def __init__(self, device="cpu", inference=None):
self.G = Generator()
self.G.load_state_dict(torch.load(inference, map_location=torch.device(device)))
self.G = self.G.to(device).eval()
return
def generate(self, org_A, ref_B, lms_A=None, lms_B=None, mask_A=None, mask_B=None,
diff_A=None, diff_B=None, gamma=None, beta=None, ret=False):
"""org_A is content, ref_B is style"""
res = self.G(org_A, ref_B, mask_A, mask_B, diff_A, diff_B, gamma, beta, ret)
return res
def test(self, real_A, mask_A, diff_A, real_B, mask_B, diff_B):
cur_prama = None
with torch.no_grad():
cur_prama = self.generate(real_A, real_B, None, None, mask_A, mask_B,
diff_A, diff_B, ret=True)
fake_A = self.generate(real_A, real_B, None, None, mask_A, mask_B,
diff_A, diff_B, gamma=cur_prama[0], beta=cur_prama[1])
fake_A = fake_A.squeeze(0)
# normalize
min_, max_ = fake_A.min(), fake_A.max()
fake_A.add_(-min_).div_(max_ - min_ + 1e-5)
return ToPILImage()(fake_A.cpu())
# PostProcess part
class PostProcess:
def __init__(self):
self.denoise = False
self.img_size = 256
def __call__(self, source: Image, result: Image):
source = np.array(source)
result = np.array(result)
height, width = source.shape[:2]
small_source = cv2.resize(source, (self.img_size, self.img_size))
laplacian_diff = source.astype(np.float) - cv2.resize(small_source, (width, height)).astype(np.float)
result = (cv2.resize(result, (width, height)) + laplacian_diff).round().clip(0, 255).astype(np.uint8)
if self.denoise:
result = cv2.fastNlMeansDenoisingColored(result)
result = Image.fromarray(result).convert('RGB')
return result
class PSGAN_Inference:
"""
An inference wrapper for makeup transfer.
It takes two image `source` and `reference` in,
and transfers the makeup of reference to source.
"""
def __init__(self, device="cpu", model_path="assets/models/G.pth", retinaface_detection=None, face_skin=None, landmark_path=None):
"""
Args:
device (str): Device type and index, such as "cpu" or "cuda:2".
device_id (int): Specifying which device index
will be used for inference.
"""
self.device = device
self.solver = Solver(device, inference=model_path)
self.preprocess = PreProcess(device, retinaface_detection=retinaface_detection, face_skin=face_skin, landmark_path=landmark_path)
self.postprocess = PostProcess()
def transfer(self, source: Image, reference: Image):
"""
Args:
source (Image): The image where makeup will be transferred to.
reference (Image): Image containing targeted makeup.
Return:
Image: Transferred image.
"""
source_input, face, crop_face = self.preprocess(source)
reference_input, _, _ = self.preprocess(reference)
if not (source_input and reference_input):
return source
for i in range(len(source_input)):
source_input[i] = source_input[i].to(self.device)
for i in range(len(reference_input)):
reference_input[i] = reference_input[i].to(self.device)
# TODO: Abridge the parameter list.
result = self.solver.test(*source_input, *reference_input)
source_crop = source.crop((crop_face.left(), crop_face.top(), crop_face.right(), crop_face.bottom()))
result = self.postprocess(source_crop, result)
return result