包含: - 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密钥已脱敏为环境变量,原文件备份在本地
115 lines
3.2 KiB
Python
115 lines
3.2 KiB
Python
#coding:utf-8
|
|
import torch
|
|
from uuid import uuid4
|
|
import base64
|
|
import os
|
|
import random
|
|
import shutil
|
|
import requests
|
|
import time
|
|
import json
|
|
import os.path as osp
|
|
from gen_super_image import webui_img2img
|
|
from gen_super_image import webui_img2img_diy
|
|
import urllib.request
|
|
import hashlib
|
|
import cv2
|
|
from datetime import datetime
|
|
import numpy as np
|
|
from core.hairstyle_model import HairStyle_Model
|
|
import configparser
|
|
from common.logger import config
|
|
|
|
hairstyle_process = HairStyle_Model(gpu=True,use_enhance=True)
|
|
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')
|
|
|
|
|
|
|
|
def download_img(img_url, userId, isfix=False, ismask=False):
|
|
try:
|
|
img_name = osp.basename(img_url)
|
|
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:
|
|
file_r = open(tmp_dir, 'rb')
|
|
md5_img = hashlib.md5(file_r.read()).hexdigest()
|
|
if isfix:
|
|
target_dir = osp.join(user_img_save_dir, userId, md5_img)
|
|
else:
|
|
target_dir = osp.join(user_img_save_dir, userId, md5_img)
|
|
os.makedirs(target_dir, exist_ok=True)
|
|
if ismask:
|
|
dst_file = osp.join(target_dir, md5_img + '.png')
|
|
else:
|
|
dst_file = osp.join(target_dir, md5_img + '.jpg')
|
|
shutil.move(tmp_dir, dst_file)
|
|
else:
|
|
return None, None
|
|
return dst_file, md5_img
|
|
except Exception as e:
|
|
print(e)
|
|
return None,None
|
|
|
|
def download_img_new(img_url, userId, isfix=False, ismask=False):
|
|
try:
|
|
img_name = osp.basename(img_url)
|
|
tmp_dir = osp.join(user_img_tmp_dir, img_name)
|
|
# if osp.exists(tmp_dir):
|
|
# os.remove(tmp_dir)
|
|
print(img_url)
|
|
|
|
r = requests.get(img_url)
|
|
# 写入图片
|
|
with open(tmp_dir, "wb") as f:
|
|
f.write(r.content)
|
|
|
|
|
|
return tmp_dir, None
|
|
except Exception as e:
|
|
print(e)
|
|
return None,None
|
|
|
|
|
|
def hair_enhance():
|
|
try:
|
|
# print('\n + hairEnhance input :', input)
|
|
img_path = "/home/data/hair/data/tmp/diy/1723090924587-1725452391390.jpg"
|
|
mask_path = "/home/data/hair/data/1/f1ae43b3-4788-45a6-8867-49f1458b3a08_mask.png"
|
|
# req_id = input['req_id']
|
|
gender = "girl"
|
|
# user_id = input['user_id']
|
|
|
|
# 发型图
|
|
final_img = cv2.imread(img_path)
|
|
|
|
# mask图
|
|
mask_dilate = cv2.imread(mask_path)
|
|
|
|
# 获取性别
|
|
in_gender = gender
|
|
sd_result = webui_img2img_diy(img=final_img, mask_img=mask_dilate, in_gender=in_gender, task_id="", tag="")
|
|
sd_save_path = "/home/data/hair/data/1/sd_res.png"
|
|
cv2.imwrite(sd_save_path, sd_result)
|
|
# cv2.imshow("sd_result", sd_result)
|
|
# cv2.waitKey(0)
|
|
|
|
|
|
except Exception as e:
|
|
print(e)
|
|
|
|
|
|
|
|
|
|
if __name__ == '__main__':
|
|
hair_diy()
|