#coding:utf-8 import traceback from uuid import uuid4 import imghdr import torch from gevent import monkey monkey.patch_all() import base64 import os import random import shutil import time import json import os.path as osp import urllib.request import hashlib import cv2 from datetime import datetime import glob import numpy as np from core.hairstyle_model import HairStyle_Model from hairstyle_model_infer import HairStyle_Model_Infer from prepare_ref_hairstyle_data import prepare_single, prepare_single_color from utils.callback import recall from gen_super_image import webui_img2img, webui_img2img_diy, webui_super_res_img from utils import enhance_hair import configparser from common.logger import config from queue import Queue from concurrent.futures import ThreadPoolExecutor from common.callback import * from utils import call_hair_inter from utils import landmark_processor from change_color import process_infer, resize_pre_webui hairstyle_process = HairStyle_Model(gpu=True,use_enhance=True) hairstyle_process_infer = HairStyle_Model_Infer(gpu=True, use_enhance=False) user_img_save_dir = config.get('default', 'userDir') user_img_tmp_dir = config.get('default', 'tmp_dir') user_img_res_dir = config.get('default', 'res_dir') ref_user_dir = config.get('default', 'ref_user_dir') train_save_dir = config.get('default', 'train_dir') hair_template_material_dir = config.get('default', 'hair_template_material_dir') ref_color_dir = config.get('default', 'ref_color') ref_color_imgs_dir = config.get('default', 'ref_color_img') train_upload_dir = config.get('default', 'upload_train_dir') def download_img(img_url, userId=None, isfix=False, ismask=False): try: img_name = img_url.split("/")[-1] tmp_dir = osp.join(user_img_tmp_dir, img_name) # if osp.exists(tmp_dir): # os.remove(tmp_dir) print(img_url) download_success = False for i in range(3): hairstyle_process.oss2.download_img(img_url, tmp_dir) if osp.exists(tmp_dir) and osp.getsize(tmp_dir) > 0: download_success = True break if download_success: img_type = imghdr.what(tmp_dir) new_tmp_dir = tmp_dir[:tmp_dir.rfind(".")+1] + img_type shutil.move(tmp_dir, new_tmp_dir) print("save path", new_tmp_dir) else: return None, None return new_tmp_dir, None except Exception as e: print(e) return None,None def change_hairstyle_v2(): hairstyle_dir = config.get('default', 'hairstyleDir') user_dir = config.get('default', 'userDir') ref_img_dir = config.get('default', 'refImgDir') res_dir = config.get('default', 'res_dir') hair_template_material_dir = config.get('default', 'hair_template_material_dir') train_dir = config.get('default', 'train_dir') start_time0 = time.time() is_hr = False ret = { "state": -1, "msg": "fail", "result": "", "umd": "" } # get req input try: hair_id = "042abc5d-262f-473e-bd62-79ab0a2fdce6" task_id = hair_id + "_" + str(uuid4()) hair_material_dir = os.path.join(train_dir, hair_id) user_img_url = "https://ydapp-1317132355.cos.ap-beijing.myqcloud.com/hair_mz/images/tmp/20240923166694465407_2.jpg" userId = "47385" user_img_path, _ = download_img(user_img_url, userId) print("--------------download:", time.time() - start_time0) except Exception as e: print(e) try: # 获取用户图 user img user_img_name = user_img_path[user_img_path.rfind("/") + 1:] new_user_img_path = os.path.join(user_dir, user_img_name) print("new_user_img_path: ", new_user_img_path) if os.path.exists(new_user_img_path): os.remove(new_user_img_path) shutil.copy(user_img_path, new_user_img_path) # 获取发型图 ref hair img template_ref_hair_name = "" template_ref_hair_path = "" hair_material_save_dir = os.path.join(hair_template_material_dir, hair_id) # upload train dir train_img_save_dir = os.path.join(train_upload_dir, hair_id) hair_name_lists = os.listdir(train_img_save_dir) for hair_name in hair_name_lists: if "first##" in hair_name: template_ref_hair_name = hair_name template_ref_hair_path = os.path.join(train_img_save_dir, hair_name) break new_hair_ref_img_path = os.path.join(ref_img_dir, template_ref_hair_name) print("new_hair_ref_img_path: ", new_hair_ref_img_path) # hair_template_save_dir = os.path.join(hair_template_dir, hair_id) # pkl_process(hair_template_save_dir) long_flag = False long_txt_path = os.path.join(hair_material_dir, "long.txt") if os.path.exists(long_txt_path): long_flag = True print("long_flag:", long_flag) # 1. gen hair res img material start1 = time.time() material_save_path = os.path.join(hairstyle_dir, hair_id) print("first_hair_material_save_path :", material_save_path) if not os.path.exists(material_save_path): print("!!! gen material !!!!") shutil.copy(template_ref_hair_path, new_hair_ref_img_path) prepare_single(new_hair_ref_img_path, material_save_path, long_flag) print("--------------prepare_single:", time.time() - start1) # 2. swap user hair to ref img hair dst_path = os.path.join(res_dir, task_id + ".png") if not os.path.exists(os.path.dirname(dst_path)): os.makedirs(os.path.dirname(dst_path)) start2 = time.time() origin_img = cv2.imread(new_user_img_path) # resize origin user img user_scale = 1920 / max(origin_img.shape[0], origin_img.shape[1]) if user_scale < 1.0: origin_img = cv2.resize(origin_img, (0, 0), fx=user_scale, fy=user_scale, interpolation=cv2.INTER_LANCZOS4) ret_dict, status = hairstyle_process_infer.infer_hairstyle(origin_img, material_save_path, return_pt1k=True, use_enhance=False) img_res = ret_dict['user_res_8uc3_orisize'] origin_pt1k = ret_dict['landmarks_origin_img_1k'] user_bald_res_8uc3_orisize = ret_dict['user_bald_res_8uc3_orisize'] print("--------------hairstyle_process infer_hairstyle:", time.time() - start2) if status == 0: # 读取新图抠图 hair_matting_path = os.path.join(material_save_path, "hair_mask_2.png") new_matting = cv2.imread(hair_matting_path, cv2.IMREAD_GRAYSCALE) # 读取结果图 result_img = img_res user_orig_mask_path = os.path.join(material_save_path, "user_orig_mask.png") origin_matting = cv2.imread(user_orig_mask_path, cv2.IMREAD_GRAYSCALE) # 获取头发处理的局部区域图像 box_info = hairstyle_process.get_body_info(result_img) dst_size = (576, 768) box_w, box_h = box_info[2] - box_info[0], box_info[3] - box_info[1] scale = min(dst_size[1] / box_h, dst_size[0] / box_w) rotate_center = [(box_info[2] + box_info[0]) * 0.5, (box_info[3] + box_info[1]) * 0.5] M = cv2.getRotationMatrix2D(rotate_center, 0, scale) M[:, 2] += np.float32([dst_size[0] * 0.5, dst_size[1] * 0.5]) - np.float32(rotate_center) crop_result = landmark_processor.high_quality_warpAffine(result_img, M, dst_size) # matting_merge = np.concatenate([origin_matting[:, :, np.newaxis], new_matting[:, :, np.newaxis]], # axis=2) # matting_merge = np.max(matting_merge, axis=2) matting_merge = new_matting[:, :, np.newaxis] crop_matting = cv2.warpAffine(matting_merge, M, dst_size) crop_new_matting = cv2.warpAffine(new_matting, M, dst_size) mask = (crop_matting > 10).astype(np.float32) if not is_hr: mask_dilate = cv2.dilate(mask, np.ones((3, 9), np.uint8)) else: mask_dilate = cv2.dilate(mask, np.ones((6, 18), np.uint8)) pt1k = landmark_processor.transform_points(origin_pt1k, M) # face_mask = landmark_processor.draw_hull_mask(pt1k.astype(np.int32), # w=mask_dilate.shape[1], h=mask_dilate.shape[0], # is_gray=True).astype(np.float32) face_mask = landmark_processor.draw_half_mask(pt1k.astype(np.int32), w=mask_dilate.shape[1], h=mask_dilate.shape[0], is_gray=True).astype(np.float32) face_mask = cv2.erode(face_mask, np.ones((19, 19), np.uint8)) face_mask = cv2.blur(face_mask, (11, 11)) crop_new_matting_f32 = crop_new_matting.astype(np.float32) / 255 face_mask = np.clip(face_mask - crop_new_matting_f32, 0, 1) mask_dilate = np.clip(mask_dilate - face_mask, 0, 1) final_img = crop_result mask_dilate = np.clip(mask_dilate * 255, 0, 255).astype(np.uint8) # get gender config_json_path = os.path.join(material_save_path, "config.json") with open(config_json_path, "r") as f: config_json_content = json.load(f) in_gender = config_json_content["gender"] print("in_gender:", in_gender) images_dir = os.path.join(hair_material_dir, "images") txt_dir = os.path.join(images_dir, os.listdir(images_dir)[0]) txt_path = glob.glob(txt_dir + '/*.txt')[0] with open(txt_path, 'r') as f: p_tag = f.readline() if "titor hairstyle, faceless, no human, gray background, simple background" in p_tag: p_tag = p_tag[p_tag.find("simple background, ") + len("simple background, "):] else: p_tag = "" start4 = time.time() # cv2.imshow("final_img", final_img) # cv2.imshow("face_mask", face_mask) # cv2.imshow("mask_dilate", mask_dilate) denoising_strength = 0.6 sd_result = webui_img2img(img=final_img, mask_img=mask_dilate, in_gender=in_gender, task_id=task_id, hair_id=hair_id, lora_material_path=hair_material_dir, tag=p_tag, is_hr=is_hr, denoising_strength=denoising_strength) # cv2.imshow("final_img", cv2.resize(final_img, (768, 768), interpolation=cv2.INTER_AREA)) # cv2.imshow("mask_dilate", cv2.resize(mask_dilate, (768, 768), interpolation=cv2.INTER_AREA)) # cv2.waitKey(0) print("--------------webui_img2img:", time.time() - start4) start5 = time.time() origin_img_final = origin_img.copy() # restore the sd_result_small to origin_img M_inv = cv2.invertAffineTransform(M) cv2.warpAffine(sd_result, M_inv, (origin_img_final.shape[1], origin_img_final.shape[0]), dst=origin_img_final, borderMode=cv2.BORDER_TRANSPARENT, flags=cv2.INTER_LANCZOS4) # face_mask_origin = cv2.warpAffine(face_mask, M_inv, (origin_img_final.shape[1], origin_img_final.shape[0]))[ # :, :, np.newaxis] # origin_img_final = (origin_img_final * ( # 1 - face_mask_origin) + user_bald_res_8uc3_orisize * face_mask_origin).astype(np.uint8) cv2.imwrite(dst_path, origin_img_final) print("origin_img_final shape:", origin_img_final.shape) # cv2.imshow("origin_img_final", origin_img_final) # cv2.waitKey(0) # t_name_res = 'digital_cloth/' + task_id + ".png" # res_url = oss_2.upload_file(dst_path, t_name_res) # print('res url:', res_url) # print("--------------write:", time.time() - start5) ret_url = hairstyle_process.oss2.upload_file(dst_path, "hair_mz/images/hairstyle/{}/{}".format(hair_id, str(uuid4()) + '.jpg')) ret["msg"] = 'success' ret['state'] = 0 ret['result'] = ret_url print("---------------------------------- cost: ", time.time() - start_time0) print("status :", status) print("\n\n") print('\n\n\n\n') except Exception as e: print(e) ret["msg"] = str(traceback.print_exc()) ret["result"] = "" ret['state'] = -1 if __name__ == '__main__': change_hairstyle_v2()