Files
change_hair/project/hair_service_sd/prepare_ref_hairstyle_data.py
T
xsl 443cfa298f 初始化:换发型/换发色/训练发型服务
包含:
- 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密钥已脱敏为环境变量,原文件备份在本地
2026-07-07 13:53:52 +08:00

350 lines
17 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import cv2
import numpy as np
import torch
import json
import time
import re
import os
from process_modules import Get_Landmark, Process_Data
from utils import landmark_processor
class GenderClassifyProcessor(object):
def __init__(self, gpu_id=0):
model_path = "./weights/gender_models"
self.output_img_size = 128
if not os.path.exists(model_path):
print("GenderClassifyProcessor don't have model!")
if gpu_id == 'cpu':
self.device = torch.device(gpu_id)
else:
self.device = torch.device('cuda:{0}'.format(gpu_id))
self.gender_model = cv2.dnn.readNetFromCaffe(os.path.join(model_path, "gender.prototxt"), os.path.join(model_path, "gender.caffemodel"))
def forward(self, img, landmark_137):
image_to_face_mat = landmark_processor.get_transform_mat_sex(landmark_137, self.output_img_size)
gender_img = cv2.warpAffine(img, image_to_face_mat, (self.output_img_size, self.output_img_size), cv2.INTER_LANCZOS4)
inpBlob = cv2.dnn.blobFromImage(gender_img, 1.0, (self.output_img_size, self.output_img_size), (0, 0, 0), swapRB=False,
crop=False)
self.gender_model.setInput(inpBlob)
output = self.gender_model.forward()
is_female = True
if output[0][0] > output[0][1]:
is_female = False
return is_female
class Prepare_Ref_HairStyle_Data(object):
def __init__(self, gpu, device_id):
if gpu and torch.cuda.is_available():
self.device = torch.device("cuda:%d" % device_id if device_id >= 0 else "cpu")
self.cuda = True
else:
self.device = torch.device("cpu")
self.cuda = False
self.process_data = Process_Data(gpu, device_id)
self.get_landmark = Get_Landmark(gpu_id=device_id)
self.gender_classify = GenderClassifyProcessor(gpu_id=device_id)
def get_prepare_ref_768_color_data(self, ref_rgb_8uc3_orisize):
ref_landmark_1k2_f_orisize = self.get_landmark.forward(ref_rgb_8uc3_orisize)
ref_rgb_8uc3_color_768, ref_matting_8uc3_color_768, ref_baldseg_8uc3_color_768, ref_landmark_f1k2_color_768 = \
self.process_data.get_prepare_ref_768_bald_data(ref_rgb_8uc3_orisize, ref_landmark_1k2_f_orisize)
return ref_rgb_8uc3_color_768, ref_matting_8uc3_color_768, ref_baldseg_8uc3_color_768, ref_landmark_f1k2_color_768
def get_prepare_ref_768_color_data_landmark1k(self, ref_rgb_8uc3_orisize,ref_landmark_1k2_f_orisize):
# ref_landmark_1k2_f_orisize = self.get_landmark.inference(ref_rgb_8uc3_orisize)
ref_rgb_8uc3_color_768, ref_matting_8uc3_color_768, ref_baldseg_8uc3_color_768, ref_landmark_f1k2_color_768 = \
self.process_data.get_prepare_ref_768_bald_data(ref_rgb_8uc3_orisize, ref_landmark_1k2_f_orisize)
return ref_rgb_8uc3_color_768, ref_matting_8uc3_color_768, ref_baldseg_8uc3_color_768, ref_landmark_f1k2_color_768
def calculate_hair_ratio_after_align(self, hair_mask, origin_landmark1k, img_size=768):
image_to_face_mat = landmark_processor.get_transform_mat_hair_ratio_v1(origin_landmark1k, 768, ratio=0.35, h_offset=0.32)
hair_mask_align = cv2.warpAffine(hair_mask, image_to_face_mat, (img_size, img_size))
hair_rect = cv2.boundingRect(hair_mask_align[:, :, :1])
hair_mask_ratio = hair_rect[2] * hair_rect[3] / (img_size * img_size)
return hair_mask_ratio
def check_female_hair_ratio(self, origin_img_8uc3, landmark_1k2_f_orisize):
ref_matte_fg_8uc3_orisize, ref_matte_pred_8uc1_orisize = self.process_data.generator_matte.matte_inference(origin_img_8uc3, landmark_1k2_f_orisize)
ref_matte_pred_8uc3_orisize = np.repeat(ref_matte_pred_8uc1_orisize[:, :, np.newaxis], 3, axis=2)
hairstyle_M = self.process_data.get_hair_M_girl_v1(landmark_1k2_f_orisize)
ref_rgb_8uc3_768 = cv2.warpAffine(ref_matte_pred_8uc3_orisize, hairstyle_M, (768, 768))
# cv2.imshow("ref_rgb_8uc3_768", ref_rgb_8uc3_768)
# cv2.waitKey()
edge_width = 5
if (ref_rgb_8uc3_768[-edge_width:, :, :]).max() > 0 or (ref_rgb_8uc3_768[:, -edge_width:, :]).max() > 0 or (ref_rgb_8uc3_768[:, :edge_width, :]).max() > 0:
return False
else:
return True
def get_prepare_ref_768_data(self, ref_rgb_8uc3_orisize, long_flag=False):
ref_landmark_1k2_f_orisize = self.get_landmark.forward(ref_rgb_8uc3_orisize)
# ref_matte_fg_8uc3_orisize, ref_matte_pred_8uc1_orisize = self.process_data.generator_matte.matte_inference(
# ref_rgb_8uc3_orisize, ref_landmark_1k2_f_orisize)
ref_landmark_137kpts_f_orisize = landmark_processor.pts_1k_to_137(ref_landmark_1k2_f_orisize)
gender_res = self.gender_classify.forward(ref_rgb_8uc3_orisize, ref_landmark_137kpts_f_orisize)
# hair_ratio = self.calculate_hair_ratio_after_align(ref_rgb_8uc3_orisize, ref_landmark_1k2_f_orisize)
# if hair_ratio > 0.3:
# ratio = 2
# else:
# if gender_res:
# ratio = 1
# else:
# ratio = 0
if not gender_res:
ratio = 0
else:
check_res = self.check_female_hair_ratio(ref_rgb_8uc3_orisize, ref_landmark_137kpts_f_orisize)
if check_res:
ratio = 1
else:
ratio = 2
if gender_res:
gender = "girl"
else:
gender = "boy"
print("gender: ", gender, " ratio: ", ratio)
# show_concat = np.concatenate((ref_rgb_8uc3_orisize, ref_matte_fg_8uc3_orisize), axis=1)
# resize_ratio = 1024. / max(show_concat.shape[:2])
# show_concat = cv2.resize(show_concat, (0, 0), fx=resize_ratio, fy=resize_ratio)
# print("gender_res: ", gender_res, " hair_ratio: ", hair_ratio)
# cv2.imshow("show_concat", show_concat)
# cv2.imshow("ref_matte_pred_8uc1_orisize", ref_matte_pred_8uc1_orisize)
# cv2.waitKey()
ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768, ref_landmark_f1k2_768 = \
self.process_data.get_prepare_ref_768_data(ref_rgb_8uc3_orisize, ref_landmark_1k2_f_orisize, ratio, long=long_flag)
# cv2.imshow("ref_rgb_8uc3_768", ref_rgb_8uc3_768)
# cv2.waitKey()
return ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768, ref_landmark_f1k2_768, gender, ratio
def Generator_reftensor(self, ref_rgb_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768,
ref_landmark_f1k2_768):
"""
input
图像尺寸基于: 人脸 512, 图像均为3通道
ref_rgb_8uc3_512: 参考图, size 512 uint8 0-255
ref_matting_8uc3_512:参考图 matting alpha 值, uint8 0-255
ref_baldseg_8uc3_512: 参考图 光头分割, uint8 0-255
ref_landmark_f1k2_512 参考图 关键点 1k*2 float32
output
input_another_pose_hair_image: 参考图 条件图, float32 0-255
"""
# 8 ********************** another pose hair_image **********************
another_pose_image = ref_rgb_8uc3_768.copy()
another_nohair_pose_mask = ref_baldseg_8uc3_768.copy()
# cv2.imshow("another_nohair_pose_mask", another_nohair_pose_mask)
# cv2.waitKey()
another_nohair_pose_mask[(np.abs(another_nohair_pose_mask - [0, 0, 255]) < 50).all(axis=2)] = [0, 255, 255]
another_nohair_pose_mask[(another_nohair_pose_mask == [0, 0, 0]).all(axis=2)] = [255, 255, 255]
another_pose_pts137 = landmark_processor.pts_1k_to_137(ref_landmark_f1k2_768).astype(np.int32)
cv2.fillPoly(another_nohair_pose_mask,
np.concatenate((another_pose_pts137[47:35:-1], another_pose_pts137[56:64],
[another_pose_pts137[48], another_pose_pts137[22]]))[
np.newaxis, :, :], (255, 255, 0))
cv2.fillPoly(another_nohair_pose_mask,
np.concatenate((another_pose_pts137[22:37], another_pose_pts137[56:47:-1]))[np.newaxis, :, :],
(255, 0, 128))
cv2.fillPoly(another_nohair_pose_mask, another_pose_pts137[88:104][np.newaxis, :, :], (255, 0, 255)) # eye
cv2.fillPoly(another_nohair_pose_mask, another_pose_pts137[105:121][np.newaxis, :, :], (255, 0, 255)) # eye
# Label nose
cv2.fillPoly(another_nohair_pose_mask, another_pose_pts137[64:79][np.newaxis, :, :], (255, 255, 255))
# Label eyebrow
cv2.fillPoly(another_nohair_pose_mask, another_pose_pts137[121:129][np.newaxis, :, :], (255, 128, 0))
cv2.fillPoly(another_nohair_pose_mask, another_pose_pts137[129:137][np.newaxis, :, :], (255, 128, 0))
another_pose_hair_alpha = ref_matting_8uc3_768.copy() / 255.
another_pose_hair_image = (another_pose_image * another_pose_hair_alpha + another_nohair_pose_mask * (
1 - another_pose_hair_alpha)).astype(np.uint8)
input_another_pose_hair_image = another_pose_hair_image.astype(np.float32) / 255
return input_another_pose_hair_image
def prepare_single(img_path=None, dst_dir=None, long_flag=False):
# img_path = "/home/yangchaojie/Desktop/hairstyle/test_data/1123/pics/test3.jpg"
# img_path = "./test_data/ref_imgs/female1.jpg"
# dst_dir = "./data/ref_hairstyle/tmp"
# img_path = "/home/szlc/Downloads/111/111/ddbc6cb6-766c-4f08-b1fd-22c6bab86432.jpg"
# dst_dir = "/home/szlc/Downloads/111/res"
if not os.path.exists(dst_dir):
os.makedirs(dst_dir)
prepare_data_process = Prepare_Ref_HairStyle_Data(True, device_id=0)
ref_rgb_8uc3_orisize = cv2.imread(img_path)
# debug = True
debug = False
version = "20221206"
start_time = time.time()
ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768, ref_landmark_f1k2_768, gender, ratio = \
prepare_data_process.get_prepare_ref_768_data(ref_rgb_8uc3_orisize, long_flag)
if not debug:
cv2.imwrite(os.path.join(dst_dir, "ref_rgb_8uc3_768.png"), ref_rgb_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_matting_fg_8uc3_768.png"), ref_matting_fg_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_matting_8uc3_768.png"), ref_matting_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_baldseg_8uc3_768.png"), ref_baldseg_8uc3_768)
np.savetxt(os.path.join(dst_dir, "ref_landmark_f1k2_768.txt"), ref_landmark_f1k2_768)
input_another_pose_hair_image = prepare_data_process.Generator_reftensor(ref_rgb_8uc3_768, ref_matting_8uc3_768,
ref_baldseg_8uc3_768,
ref_landmark_f1k2_768)
print("cost time: ", time.time() - start_time)
if debug:
print("gender: ", gender, " ratio: ", ratio)
show_concat = np.concatenate((ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768), axis=1)
# cv2.imshow("show_concat", cv2.resize(show_concat, (0, 0), fx=0.6, fy=0.6))
# cv2.imshow("input_another_pose_hair_image", input_another_pose_hair_image)
# cv2.waitKey()
if not debug:
np.save(os.path.join(dst_dir, "input_another_pose_hair_image.npy"), input_another_pose_hair_image)
config_dict = {}
config_dict['gender'] = gender
config_dict['version'] = version
config_dict['ratio'] = str(ratio)
with open(os.path.join(dst_dir, "config.json"), "w") as f:
json.dump(config_dict, f)
print("写入文件完成...", dst_dir)
def prepare_single_color(img_path=None, dst_dir=None, long_flag=False):
if not os.path.exists(dst_dir):
os.makedirs(dst_dir)
prepare_data_process = Prepare_Ref_HairStyle_Data(True, device_id=0)
ref_rgb_8uc3_orisize = cv2.imread(img_path)
# debug = True
debug = False
version = "20221206"
start_time = time.time()
ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768, ref_landmark_f1k2_768, gender, ratio = \
prepare_data_process.get_prepare_ref_768_data(ref_rgb_8uc3_orisize, long_flag)
if not debug:
cv2.imwrite(os.path.join(dst_dir, "ref_rgb_8uc3_color_768.png"), ref_rgb_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_matting_fg_8uc3_color_768.png"), ref_matting_fg_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_matting_8uc3_color_768.png"), ref_matting_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_baldseg_8uc3_color_768.png"), ref_baldseg_8uc3_768)
np.savetxt(os.path.join(dst_dir, "ref_landmark_f1k2_color_768.txt"), ref_landmark_f1k2_768)
input_another_pose_hair_image = prepare_data_process.Generator_reftensor(ref_rgb_8uc3_768, ref_matting_8uc3_768,
ref_baldseg_8uc3_768,
ref_landmark_f1k2_768)
print("cost time: ", time.time() - start_time)
if debug:
print("gender: ", gender, " ratio: ", ratio)
show_concat = np.concatenate((ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768), axis=1)
# cv2.imshow("show_concat", cv2.resize(show_concat, (0, 0), fx=0.6, fy=0.6))
# cv2.imshow("input_another_pose_hair_image", input_another_pose_hair_image)
# cv2.waitKey()
if not debug:
np.save(os.path.join(dst_dir, "input_another_pose_hair_image.npy"), input_another_pose_hair_image)
config_dict = {}
config_dict['gender'] = gender
config_dict['version'] = version
config_dict['ratio'] = str(ratio)
with open(os.path.join(dst_dir, "config.json"), "w") as f:
json.dump(config_dict, f)
print("写入文件完成...", dst_dir)
def prepare_multi():
img_dir = "/home/yangchaojie/Desktop/hairstyle/test_data/testdata_tj_0309/raw_pics"
save_dir = "/home/yangchaojie/Desktop/hairstyle/test_data/testdata_tj_0309/ref_haircolor"
# pattern = re.compile(r'^[^\.].*\.jpg$')
pattern = re.compile(r'^[^\.].*\.(jpg|JPG)$')
for dirpath, dirnames, filenames in os.walk(img_dir, followlinks=True):
for filename in filenames:
if not pattern.match(filename): continue
img_path = os.path.join(dirpath, filename)
dst_dir = os.path.join(save_dir, filename[:-4])
if not os.path.exists(dst_dir):
os.makedirs(dst_dir)
prepare_data_process = Prepare_Ref_HairStyle_Data(True, device_id=0)
ref_rgb_8uc3_orisize = cv2.imread(img_path)
# debug = True
debug = False
version = "20230309"
start_time = time.time()
ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768, ref_landmark_f1k2_768, gender, ratio = \
prepare_data_process.get_prepare_ref_768_data(ref_rgb_8uc3_orisize)
if not debug:
cv2.imwrite(os.path.join(dst_dir, "ref_rgb_8uc3_768.png"), ref_rgb_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_matting_fg_8uc3_768.png"), ref_matting_fg_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_matting_8uc3_768.png"), ref_matting_8uc3_768)
cv2.imwrite(os.path.join(dst_dir, "ref_baldseg_8uc3_768.png"), ref_baldseg_8uc3_768)
np.savetxt(os.path.join(dst_dir, "ref_landmark_f1k2_768.txt"), ref_landmark_f1k2_768)
input_another_pose_hair_image = prepare_data_process.Generator_reftensor(ref_rgb_8uc3_768, ref_matting_8uc3_768,
ref_baldseg_8uc3_768,
ref_landmark_f1k2_768)
print("cost time: ", time.time() - start_time)
if debug:
print("gender: ", gender, " ratio: ", ratio)
show_concat = np.concatenate((ref_rgb_8uc3_768, ref_matting_fg_8uc3_768, ref_matting_8uc3_768, ref_baldseg_8uc3_768), axis=1)
cv2.imshow("show_concat", cv2.resize(show_concat, (0, 0), fx=0.6, fy=0.6))
cv2.imshow("input_another_pose_hair_image", input_another_pose_hair_image)
cv2.waitKey()
if not debug:
np.save(os.path.join(dst_dir, "input_another_pose_hair_image.npy"), input_another_pose_hair_image)
config_dict = {}
config_dict['gender'] = gender
config_dict['version'] = version
config_dict['ratio'] = str(ratio)
with open(os.path.join(dst_dir, "config.json"), "w") as f:
json.dump(config_dict, f)
print("写入文件完成...", dst_dir)
if __name__ == '__main__':
prepare_single()
# prepare_multi()