Files
change_hair/project/hair_service_sd/enhance_test.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

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()