包含: - hair_service_sd: 主服务(换发型/换发色/生发,端口8801) - photo_service: LoRA调度+训练(端口32678) - hair_grow_service: 调试测试页(端口8888,含4个测试页) - 批量训练脚本(batch_train_hairstyles.py) - 发际线mask自动识别(hairline_mask.py,4种方案) - 手绘mask换发型(hair_swap_manual.py) - 文档:README.md + LARGE_FILES.md + docs/ 大文件(模型权重200G、训练数据123G)已排除,见 LARGE_FILES.md OSS/COS密钥已脱敏为环境变量,原文件备份在本地
297 lines
16 KiB
Python
297 lines
16 KiB
Python
#coding:utf-8
|
|
import os
|
|
|
|
|
|
|
|
import math
|
|
import shutil
|
|
import time
|
|
|
|
from step05_detect_fa_hairmatting_inplace import pkl_process
|
|
from PIL import Image, ImageFont, ImageDraw
|
|
import cv2
|
|
import numpy as np
|
|
import torch
|
|
import pickle
|
|
import json
|
|
import glob
|
|
import uuid
|
|
from gpt4v_caption import caption_image
|
|
from utils.call_hair_train import call_hair_train
|
|
import os
|
|
import os.path as osp
|
|
from common.logger import LogFactory
|
|
from core.process_modules import Process_Data, localtranslationwarpfastwithstrength,\
|
|
Generator_Hair, chinClass, Generator_Fusion_Res, Change_Hair_Color, GenderClassifyProcessor, BodySeg, \
|
|
localtranslationwarpfastwithstrength_v2, localtranslationwarpfastwithstrength_v2_soft, updateEndPosition
|
|
from process_modules import Get_Landmark as Get_Landmark_mtcnn
|
|
from core.face_enhance.face_enhancement import FaceEnhancement
|
|
from datetime import datetime
|
|
import random
|
|
from utils import landmark_processor
|
|
# from core.cos_module import COS_object as OSS_object
|
|
from core.oss_module import OSS_object
|
|
from core.faceseg.face_seg import FaceSeg
|
|
from common.logger import config
|
|
from process_modules import PersonProcessor_yolov5,KeypointsProcessor,Human_Keypoints,pt_conv_25_to_17
|
|
from gen_super_image import get_high_train_img
|
|
from core.MMCVFaceRecognitionServer import MomocvFaceRecognitionServer
|
|
import process_modules
|
|
|
|
from common.logger import config
|
|
|
|
version = config.get('default', 'version')
|
|
if version == "local":
|
|
train_gpu_nums = 2
|
|
else:
|
|
train_gpu_nums = 1
|
|
|
|
# webui_services = ["http://hairservice.tslead.net:32678/", "http://hairservice.tslead.net:32679/"]
|
|
|
|
class HairInit(object):
|
|
__instance = None
|
|
__first_init = False
|
|
|
|
def __new__(cls, gpu=True, use_enhance=False, infer_use_enhance=False, color_user_enhance=False):
|
|
if not cls.__instance:
|
|
cls.__instance = object.__new__(cls)
|
|
return cls.__instance
|
|
|
|
def __init__(self, gpu=True, use_enhance=False, infer_use_enhance=False, color_user_enhance=False):
|
|
if not self.__first_init:
|
|
worker_id = int(os.environ.get('APP_WORKER_ID', 1))
|
|
|
|
rand_max = 9527
|
|
self.gpu_index = (worker_id + rand_max) % train_gpu_nums
|
|
os.environ['CUDA_VISIBLE_DEVICES'] = str(self.gpu_index)
|
|
|
|
print('current worker id {} set the gpu id :{}'.format(worker_id, self.gpu_index))
|
|
device_id = self.gpu_index
|
|
|
|
|
|
self.get_landmark = Get_Landmark_mtcnn(gpu_id=device_id)
|
|
self.get_landmark_mtcnn = Get_Landmark_mtcnn(gpu_id=device_id)
|
|
self.face_recognition = MomocvFaceRecognitionServer(gpu_id=device_id)
|
|
|
|
self.hair_size = 768
|
|
self.color_output_size = 768
|
|
|
|
self.process_data = Process_Data(gpu, device_id)
|
|
self.process_data_infer = process_modules.Process_Data(gpu, device_id)
|
|
|
|
self.generator_hair = Generator_Hair(gpu, device_id)
|
|
self.hair_fusion = Generator_Fusion_Res(gpu, device_id)
|
|
self.use_enhance = use_enhance
|
|
|
|
self.face_enhance = FaceEnhancement(512, device_id)
|
|
self.change_haircolor = Change_Hair_Color(gpu, device_id)
|
|
self.logger_init = LogFactory.getLogger("init")
|
|
self.logger_process = LogFactory.getLogger("process")
|
|
self.logger_call = LogFactory.getLogger("call")
|
|
self.oss2 = OSS_object()
|
|
self.face_seg = FaceSeg(device_id)
|
|
self.chin_cls = chinClass(device_id)
|
|
model_path = "./weights/gender_models"
|
|
self.output_img_size = 128
|
|
if not os.path.exists(model_path):
|
|
print("GenderClassifyProcessor don't have model!")
|
|
self.gender_model = GenderClassifyProcessor(gpu_id=device_id)
|
|
print("Load model finish ... ")
|
|
self.baseColor_dir = os.path.join(config.get('default', "haircolorDir"),
|
|
config.get('default', "baseColor_ID"))
|
|
|
|
for i in range(2):
|
|
image = cv2.imread('data/front.jpg')
|
|
with torch.no_grad():
|
|
tmp_dir = config.get('default', "tmp_dir")
|
|
if not os.path.exists(tmp_dir):
|
|
os.makedirs(tmp_dir)
|
|
self.infer_hairstyle(image, 'data/template', tmp_dir, 'test.jpg')
|
|
|
|
self.effect_prepare_mask_fc32 = cv2.imread("./data/mask.png").astype(np.float32) / 255
|
|
# body process
|
|
self.person_processor = PersonProcessor_yolov5(gpu_id=device_id)
|
|
self.keypoints_processor = KeypointsProcessor(gpu_id=device_id)
|
|
self.human_keypoint = Human_Keypoints(gpu=True, device_id=device_id)
|
|
print("Load model finish ... ")
|
|
|
|
# infer init
|
|
# self.table_enlight = cv2.imread('./data/convert_enlight.png')
|
|
self.change_haircolor = Change_Hair_Color(gpu, device_id)
|
|
self.gender_classify = GenderClassifyProcessor(gpu_id=device_id)
|
|
|
|
HairInit.__first_init = True
|
|
print("[HAIR_INIT] All init done. ")
|
|
|
|
def infer_hairstyle(self, origin_img, hairstyle_dir, userinfo_dir, mask_newname):
|
|
config_path = os.path.join(hairstyle_dir, "config.json")
|
|
if not os.path.exists(config_path):
|
|
return None, None, 10001
|
|
config_info = json.load(open(config_path, "rb"))
|
|
gender = config_info["gender"]
|
|
ratio = int(config_info["ratio"])
|
|
# print("gender: ", gender, " ratio: ", ratio)
|
|
|
|
user_rgb_8uc3_orisize = origin_img.copy()
|
|
tt = time.time()
|
|
landmark1k_dir = osp.join(userinfo_dir, 'kpt_1k.txt')
|
|
if not osp.exists(landmark1k_dir):
|
|
with torch.no_grad():
|
|
landmarks_origin_img_1k = self.get_landmark.forward(origin_img)
|
|
if landmarks_origin_img_1k is None:
|
|
return None, None, 10001
|
|
np.savetxt(landmark1k_dir, landmarks_origin_img_1k)
|
|
else:
|
|
landmarks_origin_img_1k = np.loadtxt(landmark1k_dir)
|
|
|
|
user_bald_res_8uc3_orisize_dir = osp.join(userinfo_dir, 'bald_res_ori_raw.png')
|
|
# res_matting_8uc3_bald_orisize_dir = osp.join(userinfo_dir, 'res_matting_mask_ori.png')
|
|
user_baldseg_8uc3_orisize_dir = osp.join(userinfo_dir, 'bald_seg_ori.png')
|
|
user_baldseg_8uc3_768_dir = osp.join(userinfo_dir, 'user_baldseg_768.png')
|
|
user_bald_8uc3_768_dir = osp.join(userinfo_dir, 'bald_seg_768.png')
|
|
user_landmark_f1k2_768_dir = osp.join(userinfo_dir, 'landmark_f1k2_768.txt')
|
|
user_hairstyle_M_dir = osp.join(userinfo_dir, 'hairstyle_M.txt')
|
|
|
|
pre_list = [user_bald_res_8uc3_orisize_dir, user_baldseg_8uc3_orisize_dir, user_baldseg_8uc3_768_dir, user_bald_8uc3_768_dir,
|
|
user_landmark_f1k2_768_dir, user_hairstyle_M_dir]
|
|
|
|
condition_exist = True
|
|
for tmp_file in pre_list:
|
|
if not osp.exists(tmp_file):
|
|
condition_exist = False
|
|
if not condition_exist:
|
|
user_bald_res_8uc3_orisize, user_baldseg_8uc3_orisize, user_baldseg_8uc3_768, user_bald_8uc3_768, \
|
|
user_landmark_f1k2_768, user_hairstyle_M,_ = self.process_data.get_prepare_user_768_data(user_rgb_8uc3_orisize, landmarks_origin_img_1k, ratio=ratio)
|
|
cv2.imwrite(user_bald_res_8uc3_orisize_dir, user_bald_res_8uc3_orisize)
|
|
cv2.imwrite(user_baldseg_8uc3_orisize_dir, user_baldseg_8uc3_orisize)
|
|
cv2.imwrite(user_baldseg_8uc3_768_dir, user_baldseg_8uc3_768)
|
|
cv2.imwrite(user_bald_8uc3_768_dir, user_bald_8uc3_768)
|
|
np.savetxt(user_landmark_f1k2_768_dir, user_landmark_f1k2_768)
|
|
np.savetxt(user_hairstyle_M_dir, user_hairstyle_M)
|
|
else:
|
|
user_bald_res_8uc3_orisize = cv2.imread(user_bald_res_8uc3_orisize_dir)
|
|
user_baldseg_8uc3_orisize = cv2.imread(user_baldseg_8uc3_orisize_dir)
|
|
# user_baldseg_8uc3_768 = cv2.imread(user_baldseg_8uc3_768_dir)
|
|
# user_bald_8uc3_768 = cv2.imread(user_bald_8uc3_768_dir)
|
|
# user_landmark_f1k2_768 = np.loadtxt(user_landmark_f1k2_768_dir)
|
|
# user_hairstyle_M = np.loadtxt(user_hairstyle_M_dir)
|
|
if ratio == 0:
|
|
user_hairstyle_M = self.process_data.get_hair_M_boy_v1(landmarks_origin_img_1k)
|
|
elif ratio == 1:
|
|
user_hairstyle_M = self.process_data.get_hair_M_girl_v1(landmarks_origin_img_1k)
|
|
elif ratio == 2:
|
|
user_hairstyle_M = self.process_data.get_hair_M_girl_v2(landmarks_origin_img_1k)
|
|
else:
|
|
user_hairstyle_M = self.process_data.get_hair_M_girl_v1(landmarks_origin_img_1k)
|
|
user_landmark_f1k2_768 = landmark_processor.transform_points(landmarks_origin_img_1k, user_hairstyle_M)
|
|
user_bald_8uc3_768 = cv2.warpAffine(user_bald_res_8uc3_orisize, user_hairstyle_M,
|
|
(self.hair_size, self.hair_size))
|
|
user_baldseg_8uc3_768 = cv2.warpAffine(user_baldseg_8uc3_orisize, user_hairstyle_M,
|
|
(self.hair_size, self.hair_size))
|
|
print('condition cosst:', time.time() - tt )
|
|
# show_concat = np.concatenate((user_rgb_8uc3_orisize, user_bald_res_8uc3_orisize, user_baldseg_8uc3_orisize), axis=1)
|
|
# ratio = 1536. / max(show_concat.shape[:2])
|
|
# show_concat = cv2.resize(show_concat, (0, 0), fx=ratio, fy=ratio)
|
|
# cv2.imshow("show_concat", show_concat)
|
|
# cv2.waitKey()
|
|
|
|
another_pose_hair_img_path = os.path.join(hairstyle_dir, "input_another_pose_hair_image.npy")
|
|
if not os.path.exists(another_pose_hair_img_path):
|
|
return None, None, 10002
|
|
another_pose_hair_image = np.load(another_pose_hair_img_path)
|
|
# cv2.imshow('user_baldseg_8uc3_768', user_baldseg_8uc3_768)
|
|
# cv2.imshow('user_bald_8uc3_768', user_bald_8uc3_768)
|
|
|
|
t0 = time.time()
|
|
# 换发型
|
|
|
|
hair_gene_8uc3_768 = self.generator_hair.Generator_Hair_inference_use_pref(another_pose_hair_image,
|
|
user_baldseg_8uc3_768,
|
|
user_bald_8uc3_768,
|
|
user_landmark_f1k2_768, gender)
|
|
self.logger_process.info('Generator_Hair_inference_use_pref costs:{}'.format(time.time() - t0))
|
|
|
|
t1 = time.time()
|
|
|
|
hair_gene_fusion_8uc3_orisize, hair_gene_matte_8uc3_orisize = self.process_data.get_fusion_res_hairpaste(
|
|
user_bald_res_8uc3_orisize, hair_gene_8uc3_768, user_landmark_f1k2_768, user_hairstyle_M)
|
|
|
|
|
|
res_matting_mask_ori = osp.join(userinfo_dir, mask_newname)
|
|
res_matting_mask_ori_raw = osp.join(userinfo_dir, 'res_matting_mask_ori_raw.png')
|
|
if osp.exists(res_matting_mask_ori):
|
|
os.remove(res_matting_mask_ori)
|
|
if osp.exists(res_matting_mask_ori_raw):
|
|
os.remove(res_matting_mask_ori_raw)
|
|
|
|
self.logger_process.info('get_fusion_res_hairpaste costs:{}'.format(time.time() - t1))
|
|
|
|
user_res_8uc3_orisize = self.hair_fusion.inference(hair_gene_fusion_8uc3_orisize,
|
|
hair_gene_matte_8uc3_orisize,
|
|
landmarks_origin_img_1k,
|
|
user_baldseg_8uc3_orisize,
|
|
ratio)
|
|
hair_gene_matte_8uc3_orisize_cp = hair_gene_matte_8uc3_orisize.copy()
|
|
# show_concat = np.concatenate((hair_gene_fusion_8uc3_orisize, hair_gene_matte_8uc3_orisize, user_baldseg_8uc3_orisize, user_res_8uc3_orisize), axis=1)
|
|
# ratio = 1536. / max(show_concat.shape[:2])
|
|
# show_concat = cv2.resize(show_concat, (0, 0), fx=ratio, fy=ratio)
|
|
# cv2.imshow("mid_concat", show_concat)
|
|
# cv2.waitKey()
|
|
t2 = time.time()
|
|
|
|
user_res_8uc3_orisize_for_haircolor = user_res_8uc3_orisize.copy()
|
|
|
|
# if self.use_enhance:
|
|
user_res_8uc3_orisize_enhance = self.face_enhance.process(user_res_8uc3_orisize, landmarks_origin_img_1k)
|
|
t3 = time.time()
|
|
self.logger_process.info('use_enhance costs:{}'.format(t3 - t2))
|
|
|
|
user_bald_res_8uc3_orisize_enhance, user_baldseg_8uc3_orisize_enhance, user_baldseg_8uc3_768_enhance, user_bald_8uc3_768_enhance, \
|
|
user_landmark_f1k2_768_enhance, user_hairstyle_M_enhance, user_matting_8uc3_bald_orisize = self.process_data.get_prepare_user_768_data(
|
|
user_res_8uc3_orisize_enhance, landmarks_origin_img_1k, ratio=ratio)
|
|
|
|
t4 = time.time()
|
|
self.logger_process.info('gen bald costs:{}'.format(t4 - t3))
|
|
|
|
# 重新提取matting
|
|
_, hair_gene_matte_8uC0_orisize, _ = self.process_data.generator_matte.matte_inference(user_res_8uc3_orisize_enhance, landmarks_origin_img_1k)
|
|
t5 = time.time()
|
|
self.logger_process.info('gen matte_inference costs:{}'.format(t5 - t4))
|
|
hair_gene_matte_8uc3_orisize = np.repeat(hair_gene_matte_8uC0_orisize[:, :, np.newaxis], 3, axis=2)
|
|
hair_gene_matte_8uc3_orisize = cv2.blur(hair_gene_matte_8uc3_orisize, (3, 3))
|
|
|
|
# user_res_8uc3_orisize = user_bald_res_8uc3_orisize_enhance.astype(np.float32)/255. * (1. - hair_gene_matte_8uc3_orisize.astype(np.float32) / 255.) + user_res_8uc3_orisize_enhance.astype(np.float32)/255. * hair_gene_matte_8uc3_orisize.astype(np.float32) / 255.
|
|
# user_res_8uc3_orisize = (user_res_8uc3_orisize*255).astype(np.uint8)
|
|
|
|
user_res_fc32_orisize_enhance = user_res_8uc3_orisize_enhance.astype(np.float32)/255.
|
|
user_res_fc32_orisize = user_res_8uc3_orisize.astype(np.float32)/255.
|
|
hair_gene_matte_fc32_orisize = hair_gene_matte_8uc3_orisize.astype(np.float32)/255.
|
|
hair_gene_matte_fc32_orisize = cv2.GaussianBlur(hair_gene_matte_fc32_orisize, (11, 11), 0, 0)
|
|
user_res_fc32_orisize_enhance = user_res_fc32_orisize_enhance * hair_gene_matte_fc32_orisize + user_res_fc32_orisize * (1 - hair_gene_matte_fc32_orisize)
|
|
user_res_8uc3_orisize_enhance = (user_res_fc32_orisize_enhance * 255).astype(np.uint8)
|
|
|
|
# mid_show = np.concatenate((user_res_fc32_orisize_enhance, user_res_8uc3_orisize_enhance, hair_gene_matte_8uc3_orisize, user_res_8uc3_orisize), axis=1)
|
|
# ratio = 1536. / max(mid_show.shape[:2])
|
|
# mid_show = cv2.resize(mid_show, (0, 0), fx=ratio, fy=ratio)
|
|
# cv2.imshow("mid_show", mid_show)
|
|
# cv2.imshow("user_res_8uc3_orisize_enhance_fg", user_res_8uc3_orisize_enhance.astype(np.float32)/255. * hair_gene_matte_8uc3_orisize.astype(np.float32) / 255)
|
|
# cv2.waitKey()
|
|
|
|
user_bald_res_8uc3_orisize_for_fusion_dir = osp.join(userinfo_dir, 'bald_res_ori.png')
|
|
# if osp.exists(user_bald_res_8uc3_orisize_for_fusion_dir):
|
|
# os.remove(user_bald_res_8uc3_orisize_for_fusion_dir)
|
|
cv2.imwrite(user_bald_res_8uc3_orisize_for_fusion_dir, user_bald_res_8uc3_orisize_enhance)
|
|
# if osp.exists(res_matting_mask_ori_raw):
|
|
# os.remove(res_matting_mask_ori_raw)
|
|
cv2.imwrite(res_matting_mask_ori_raw, hair_gene_matte_8uc3_orisize)
|
|
|
|
# user_bald_res_8uc3_orisize_dir = osp.join(userinfo_dir, 'bald_res_ori.png')
|
|
# cv2.imwrite(user_bald_res_8uc3_orisize_dir, user_res_8uc3_orisize)
|
|
|
|
|
|
res_fix_img_mask_8uc4_orisize = np.concatenate((user_res_8uc3_orisize_enhance, hair_gene_matte_8uc3_orisize_cp[:, :, :1]), axis=2)
|
|
# cv2.imshow('')
|
|
cv2.imwrite(res_matting_mask_ori, res_fix_img_mask_8uc4_orisize)
|
|
print('costs:', time.time() - t5)
|
|
return user_res_8uc3_orisize_enhance, user_res_8uc3_orisize_for_haircolor, 0
|