包含: - 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密钥已脱敏为环境变量,原文件备份在本地
160 lines
5.6 KiB
Python
160 lines
5.6 KiB
Python
import os
|
|
import sys
|
|
from webui_im2im import ControlnetRequestImg2Img
|
|
import numpy as np
|
|
import base64
|
|
import cv2
|
|
import os,sys
|
|
from gevent import pywsgi, monkey
|
|
|
|
monkey.patch_all()
|
|
|
|
# 将当前工作目录切换到当前目录
|
|
project_dir = os.path.dirname(os.path.abspath(__file__))
|
|
os.chdir(project_dir)
|
|
sys.path.append(project_dir)
|
|
|
|
from flask import Flask, request, jsonify
|
|
import global_variable as global_var
|
|
import json
|
|
|
|
app = Flask(__name__)
|
|
|
|
|
|
@app.route('/template/list', methods=['POST'])
|
|
def get_template_list():
|
|
import os
|
|
import json
|
|
try:
|
|
# 读取data/template.json文件,并返回
|
|
with open('data/template.json', 'r') as f:
|
|
template_list = json.load(f)
|
|
ret_dict = dict(code=0, message='success', data=template_list)
|
|
# 返回结果作为 JSON 响应
|
|
return jsonify(ret_dict)
|
|
|
|
except Exception as e:
|
|
ret_dict = dict(code=-1, message=str(e), data=[])
|
|
return jsonify(ret_dict)
|
|
|
|
|
|
@app.route('/user/list', methods=['POST'])
|
|
def get_user_list():
|
|
try:
|
|
# 获取用户文件夹路径
|
|
user_list_dir = os.path.join(global_var.service_data_dir, 'user_data')
|
|
|
|
ret_user_list = []
|
|
|
|
# 遍历查找user_list_dir目录下所有的cfg.json文件
|
|
for root, dirs, files in os.walk(user_list_dir):
|
|
for file in files:
|
|
if file != 'cfg.json': continue
|
|
json_path = os.path.join(root, file)
|
|
lora_path = os.path.join(root, 'lora.safetensors')
|
|
if not os.path.exists(lora_path): continue
|
|
# 读取cfg.json文件
|
|
with open(json_path, 'r') as f:
|
|
user_info = json.load(f)
|
|
user_dict = dict(user_id=user_info['user_id'], face_img_url=user_info['face_img_url'])
|
|
ret_user_list.append(user_dict)
|
|
|
|
ret_dict = dict(code=0, message='success', data=ret_user_list)
|
|
return jsonify(ret_dict)
|
|
|
|
except Exception as e:
|
|
ret_dict = dict(code=-1, message='error', data=[])
|
|
return jsonify(ret_dict)
|
|
|
|
|
|
@app.route('/user/generate', methods=['POST'])
|
|
def generate_photo():
|
|
try:
|
|
# 获取请求参数
|
|
request_data = request.get_json()
|
|
#判断'user_id'和'base_img'是否在请求参数中
|
|
assert 'user_id' in request_data and 'base_img' in request_data, 'user_id and base_img is required'
|
|
user_id = request_data['user_id']
|
|
base_img_b64 = request_data['base_img']
|
|
|
|
user_path = os.path.join(global_var.service_data_dir, 'user_data', user_id)
|
|
assert os.path.isdir(user_path), 'user_id not exist'
|
|
user_lora = os.path.join(user_path, 'lora.safetensors')
|
|
usr_config_path = os.path.join(user_path, 'cfg.json')
|
|
assert os.path.exists(user_lora) and os.path.exists(usr_config_path), 'lora file or config not exist'
|
|
|
|
# 读取用户配置文件
|
|
with open(usr_config_path, 'r') as f:
|
|
user_info = json.load(f)
|
|
lora_md5 = user_info['lora_md5']
|
|
dst_lora_path = os.path.join(global_var.webui_lora_dir, lora_md5 + '.safetensors')
|
|
if not os.path.exists(dst_lora_path):
|
|
os.system(f'cp {user_lora} {dst_lora_path}')
|
|
|
|
#将模板图像转化为numpy数组
|
|
prompt = f'<lora:{lora_md5}:0.8>,easyphoto_face, easyphoto, 1person,face,suit'
|
|
neg_prompt = '(worst quality:2),(low quality:2),(normal quality:2),lowres,watermark'
|
|
|
|
image_array = np.frombuffer(base64.b64decode(base_img_b64), np.uint8)
|
|
base_img = cv2.imdecode(image_array, cv2.IMREAD_COLOR)
|
|
# base_img = cv2.resize(base_img, (512, 512))
|
|
# cv2.imshow('image', base_img)
|
|
# cv2.waitKey()
|
|
|
|
# 生成图片
|
|
control_net = ControlnetRequestImg2Img(prompt, neg_prompt)
|
|
control_net.build_body(dst_width=base_img.shape[1], dst_height=base_img.shape[0], cfg_scale=3.5, base_img=base_img)
|
|
output = control_net.send_request()
|
|
generate_photo = output['images'][0]
|
|
|
|
# 清理硬盘空间
|
|
os.remove(dst_lora_path)
|
|
|
|
# # 将生成的图片转化为base64编码
|
|
# retval, bytes = cv2.imencode('.png', generate_photo)
|
|
# generate_photo = base64.b64encode(bytes).decode('utf-8')
|
|
return jsonify(dict(code=0, message='success', generate_photo_b64=generate_photo))
|
|
|
|
except Exception as e:
|
|
ret_dict = dict(code=-1, message=str(e), data=[])
|
|
return jsonify(ret_dict)
|
|
|
|
|
|
def webd_service():
|
|
# 用于启动webd的后台服务
|
|
current_file_dir = os.path.dirname(os.path.abspath(__file__))
|
|
webd_path = os.path.join(current_file_dir, 'webd', 'webd')
|
|
|
|
print('webd server started...')
|
|
cmd = f"{webd_path} -w {global_var.service_data_dir} -g rlT -l 10219"
|
|
os.system(cmd)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
# 服务启动的数据目录
|
|
global_var.service_data_dir = sys.argv[1]
|
|
|
|
# 本地webui的lora存储目录
|
|
global_var.webui_lora_dir = sys.argv[2]
|
|
|
|
# webui_server_port
|
|
global_var.webui_server_port = int(sys.argv[3])
|
|
|
|
# server_port
|
|
global_var.server_port = int(sys.argv[4])
|
|
|
|
# 检查webui_lora_dir目录是否存在
|
|
assert os.path.isdir(global_var.webui_lora_dir), 'webui_lora_dir should be a directory'
|
|
|
|
# 检查service_data_dir目录是否存在
|
|
if not os.path.exists(global_var.service_data_dir):
|
|
os.makedirs(global_var.service_data_dir, exist_ok=True)
|
|
else:
|
|
assert os.path.isdir(global_var.service_data_dir), 'service_data_dir should be a directory'
|
|
|
|
# 启动服务
|
|
# app.run(debug=False, port=global_var.server_port, host='0.0.0.0')
|
|
|
|
server = pywsgi.WSGIServer(('0.0.0.0', global_var.server_port), app) # test port
|
|
server.serve_forever()
|