初始化:换发型/换发色/训练发型服务
包含: - 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密钥已脱敏为环境变量,原文件备份在本地
This commit is contained in:
@@ -0,0 +1,515 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""区域生发服务(走换发型工作流,端口 8888)
|
||||
|
||||
工作流与换发型一致:infer_hairstyle_diy_jy → warpAffine → photo_service+LoRA → webui → 贴回。
|
||||
mask = origin_matting ∪ new_matting ∪ 手绘mask(不减刘海)。
|
||||
|
||||
启动: python app.py
|
||||
页面: http://<host>:8888
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import base64
|
||||
import traceback
|
||||
|
||||
# 把 hair_service_sd 加入 path,使其模块可被导入
|
||||
HAIR_SERVICE_DIR = "/home/xsl/change_hair/project/hair_service_sd"
|
||||
sys.path.insert(0, HAIR_SERVICE_DIR)
|
||||
# 切换到 hair_service_sd 目录,使 common.logger 能读到 config/configure.ini(相对路径)
|
||||
os.chdir(HAIR_SERVICE_DIR)
|
||||
os.environ.setdefault("HF_HUB_OFFLINE", "1")
|
||||
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
|
||||
os.environ.setdefault("CRYPTOGRAPHY_OPENSSL_NO_LEGACY", "1")
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from flask import Flask, request, jsonify, send_from_directory
|
||||
from gevent import pywsgi
|
||||
|
||||
PORT = 8888
|
||||
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
app = Flask(__name__, static_folder="static", static_url_path="/static")
|
||||
|
||||
# ===== 模型懒加载(首次请求时初始化,避免启动卡死)=====
|
||||
_hairstyle_process = None
|
||||
_landmark_processor = None
|
||||
|
||||
|
||||
def _get_models():
|
||||
"""懒加载换发型全套模型"""
|
||||
global _hairstyle_process, _landmark_processor
|
||||
if _hairstyle_process is None:
|
||||
print("[init] 加载换发型模型(首次,约30-90秒)...")
|
||||
t0 = time.time()
|
||||
import torch
|
||||
from core.hairstyle_model import HairStyle_Model
|
||||
from utils import landmark_processor
|
||||
_hairstyle_process = HairStyle_Model(gpu=True, use_enhance=True)
|
||||
_landmark_processor = landmark_processor
|
||||
print(f"[init] 模型加载完成,耗时 {time.time()-t0:.1f}s")
|
||||
return _hairstyle_process, _landmark_processor
|
||||
|
||||
|
||||
@app.route("/")
|
||||
def index():
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static"), "index.html")
|
||||
|
||||
|
||||
@app.route("/swap")
|
||||
def swap_page():
|
||||
"""换发型测试页(独立,无遮罩)"""
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static"), "swap.html")
|
||||
|
||||
|
||||
@app.route("/test_new")
|
||||
def test_new_page():
|
||||
"""新发型效果测试页(聚焦本次训练的30款发型)"""
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static"), "test_new.html")
|
||||
|
||||
|
||||
@app.route("/debug")
|
||||
def debug_page():
|
||||
"""换发型调试页(全参数可视化)"""
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static"), "debug.html")
|
||||
|
||||
|
||||
@app.route("/hairline")
|
||||
def hairline_page():
|
||||
"""发际线带重绘实验页"""
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static"), "hairline.html")
|
||||
|
||||
|
||||
@app.route("/manual")
|
||||
def manual_page():
|
||||
"""手绘mask重绘测试页"""
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static"), "manual.html")
|
||||
|
||||
|
||||
@app.route("/api/swap", methods=["POST"])
|
||||
def api_swap_proxy():
|
||||
"""代理换发型请求到 8801(避免前端跨域问题)"""
|
||||
import requests as req
|
||||
try:
|
||||
data = request.json
|
||||
resp = req.post("http://127.0.0.1:8801/api/swapHair/v1",
|
||||
json=data, timeout=600)
|
||||
return jsonify(resp.json())
|
||||
except Exception as e:
|
||||
return jsonify({"state": -1, "msg": f"换发型代理失败: {e}"}), 500
|
||||
|
||||
|
||||
@app.route("/api/swap_viz", methods=["POST"])
|
||||
def api_swap_viz():
|
||||
"""换发型(带可视化中间产物)
|
||||
|
||||
入参 JSON:
|
||||
img: 原图 base64
|
||||
hair_id: 发型ID
|
||||
is_hr: "true"/"false"
|
||||
返回:
|
||||
{state, result: 最终图b64, steps: [{title, desc, images:[{label, b64}]}]}
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
data = request.json
|
||||
img_b64 = data.get("img", "")
|
||||
hair_id = data.get("hair_id", "")
|
||||
is_hr = str(data.get("is_hr", "false")).lower() == "true"
|
||||
strict_mask = data.get("strict_mask", False)
|
||||
|
||||
if not img_b64 or not hair_id:
|
||||
return jsonify({"state": -1, "msg": "img 和 hair_id 不能为空"}), 400
|
||||
|
||||
img = _b64_to_ndarray(img_b64, color=True)
|
||||
if img is None:
|
||||
return jsonify({"state": -1, "msg": "img 解析失败"}), 400
|
||||
|
||||
print(f"[swap_viz] img={img.shape}, hair_id={hair_id}, is_hr={is_hr}, strict_mask={strict_mask}")
|
||||
hairstyle_process, landmark_processor = _get_models()
|
||||
|
||||
from hair_swap_viz import hair_swap_viz
|
||||
task_id = f"swapviz_{int(time.time()*1000)}"
|
||||
result, steps = hair_swap_viz(
|
||||
origin_img=img, hair_id=hair_id,
|
||||
hairstyle_process=hairstyle_process,
|
||||
landmark_processor=landmark_processor,
|
||||
task_id=task_id, is_hr=is_hr, strict_mask=strict_mask)
|
||||
|
||||
_, buf = cv2.imencode(".jpg", result, [cv2.IMWRITE_JPEG_QUALITY, 95])
|
||||
result_b64 = base64.b64encode(buf).decode("utf-8")
|
||||
print(f"[swap_viz] 完成,总耗时 {time.time()-t0:.1f}s,{len(steps)}个步骤")
|
||||
return jsonify({"state": 0, "result": result_b64, "steps": steps})
|
||||
|
||||
except Exception as e:
|
||||
print(f"[swap_viz] 失败: {e}")
|
||||
traceback.print_exc()
|
||||
return jsonify({"state": -1, "msg": f"换发型失败: {e}"}), 500
|
||||
|
||||
|
||||
@app.route("/api/swap_debug", methods=["POST"])
|
||||
def api_swap_debug():
|
||||
"""换发型调试接口(全参数可调 + 每步可视化)。
|
||||
|
||||
入参 JSON:
|
||||
img: 原图 base64
|
||||
hair_id: 发型ID
|
||||
# 流程开关
|
||||
cut_bang: bool, 是否减刘海圆(默认 true)
|
||||
strict_mask: bool, 严格按mask贴回(默认 false)
|
||||
seamless_blend: bool, 泊松融合消除接缝(默认 true,仅 strict_mask 时生效)
|
||||
# 尺寸/对齐
|
||||
is_hr: bool, 高清模式(默认 true)
|
||||
dilate_kernel: [x,y], mask膨胀核(默认 [6,18])
|
||||
# SD 推理(仅 denoising 可调)
|
||||
denoising_strength: float 0~1, 重绘强度(默认 0.6)
|
||||
# 贴回/融合
|
||||
blend_dilate: [x,y], strict贴回mask膨胀(默认 [5,5])
|
||||
seamless_dilate: [x,y], 泊松融合mask膨胀(默认 [9,9])
|
||||
返回:
|
||||
{state, result, steps, params}
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
d = request.json
|
||||
img_b64 = d.get("img", "")
|
||||
hair_id = d.get("hair_id", "")
|
||||
if not img_b64 or not hair_id:
|
||||
return jsonify({"state": -1, "msg": "img 和 hair_id 不能为空"}), 400
|
||||
|
||||
img = _b64_to_ndarray(img_b64, color=True)
|
||||
if img is None:
|
||||
return jsonify({"state": -1, "msg": "img 解析失败"}), 400
|
||||
|
||||
# 解析参数(带默认值 + 类型转换)
|
||||
def g(k, default): return d.get(k, default)
|
||||
dk = g("dilate_kernel", [6, 18])
|
||||
bd = g("blend_dilate", [5, 5])
|
||||
sd = g("seamless_dilate", [9, 9])
|
||||
|
||||
print(f"[swap_debug] hair_id={hair_id}, cut_bang={g('cut_bang',True)}, "
|
||||
f"strict_mask={g('strict_mask',False)}, seamless_blend={g('seamless_blend',True)}, "
|
||||
f"is_hr={g('is_hr',True)}, denoising={g('denoising_strength',0.6)}")
|
||||
|
||||
hairstyle_process, landmark_processor = _get_models()
|
||||
from hair_swap_debug import hair_swap_debug
|
||||
task_id = f"swapdbg_{int(time.time()*1000)}"
|
||||
result, steps, params = hair_swap_debug(
|
||||
origin_img=img, hair_id=hair_id,
|
||||
hairstyle_process=hairstyle_process,
|
||||
landmark_processor=landmark_processor,
|
||||
task_id=task_id,
|
||||
cut_bang=bool(g("cut_bang", True)),
|
||||
strict_mask=bool(g("strict_mask", False)),
|
||||
seamless_blend=bool(g("seamless_blend", True)),
|
||||
is_hr=bool(g("is_hr", True)),
|
||||
dilate_kernel=(int(dk[0]), int(dk[1])),
|
||||
denoising_strength=float(g("denoising_strength", 0.6)),
|
||||
blend_dilate=(int(bd[0]), int(bd[1])),
|
||||
seamless_dilate=(int(sd[0]), int(sd[1])),
|
||||
)
|
||||
|
||||
_, buf = cv2.imencode(".jpg", result, [cv2.IMWRITE_JPEG_QUALITY, 95])
|
||||
result_b64 = base64.b64encode(buf).decode("utf-8")
|
||||
print(f"[swap_debug] 完成,总耗时 {time.time()-t0:.1f}s")
|
||||
return jsonify({"state": 0, "result": result_b64, "steps": steps, "params": params})
|
||||
|
||||
except Exception as e:
|
||||
print(f"[swap_debug] 失败: {e}")
|
||||
traceback.print_exc()
|
||||
return jsonify({"state": -1, "msg": f"换发型失败: {e}"}), 500
|
||||
|
||||
|
||||
@app.route("/api/swap_hairline", methods=["POST"])
|
||||
def api_swap_hairline():
|
||||
"""发际线带重绘接口(实验性,流程同 swap_debug,仅步骤③改为边界带)。
|
||||
|
||||
比 swap_debug 多一个参数:
|
||||
method: str, 发际线mask方案,可选 boundary_band/mediapipe/landmark_1k/deeplab(默认 mediapipe)
|
||||
band_width: int, 边界带宽度(形态学核大小,默认15,带宽≈2*band_width)
|
||||
其余参数同 swap_debug。
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
d = request.json
|
||||
img_b64 = d.get("img", "")
|
||||
hair_id = d.get("hair_id", "")
|
||||
if not img_b64 or not hair_id:
|
||||
return jsonify({"state": -1, "msg": "img 和 hair_id 不能为空"}), 400
|
||||
|
||||
img = _b64_to_ndarray(img_b64, color=True)
|
||||
if img is None:
|
||||
return jsonify({"state": -1, "msg": "img 解析失败"}), 400
|
||||
|
||||
def g(k, default): return d.get(k, default)
|
||||
dk = g("dilate_kernel", [6, 18])
|
||||
bd = g("blend_dilate", [5, 5])
|
||||
sd = g("seamless_dilate", [9, 9])
|
||||
|
||||
print(f"[swap_hairline] hair_id={hair_id}, method={g('method','mediapipe')}, "
|
||||
f"band_width={g('band_width',15)}, denoising={g('denoising_strength',0.6)}")
|
||||
|
||||
hairstyle_process, landmark_processor = _get_models()
|
||||
from hair_swap_hairline import hair_swap_hairline
|
||||
task_id = f"hairline_{int(time.time()*1000)}"
|
||||
result, steps, params = hair_swap_hairline(
|
||||
origin_img=img, hair_id=hair_id,
|
||||
hairstyle_process=hairstyle_process,
|
||||
landmark_processor=landmark_processor,
|
||||
task_id=task_id,
|
||||
method=str(g("method", "mediapipe")),
|
||||
strict_mask=bool(g("strict_mask", False)),
|
||||
seamless_blend=bool(g("seamless_blend", True)),
|
||||
is_hr=bool(g("is_hr", True)),
|
||||
dilate_kernel=(int(dk[0]), int(dk[1])),
|
||||
denoising_strength=float(g("denoising_strength", 0.6)),
|
||||
blend_dilate=(int(bd[0]), int(bd[1])),
|
||||
seamless_dilate=(int(sd[0]), int(sd[1])),
|
||||
band_width=int(g("band_width", 15)),
|
||||
preview_only=bool(g("preview_only", False)),
|
||||
height_ratio=float(g("height_ratio", 0.432)),
|
||||
width_ratio=float(g("width_ratio", 0.144)),
|
||||
corner_ratio=float(g("corner_ratio", 0.25)),
|
||||
vertical_offset=float(g("vertical_offset", 0.0)),
|
||||
refiner_switch_at=float(g("refiner_switch_at", 0.5)),
|
||||
)
|
||||
|
||||
_, buf = cv2.imencode(".jpg", result, [cv2.IMWRITE_JPEG_QUALITY, 95])
|
||||
result_b64 = base64.b64encode(buf).decode("utf-8")
|
||||
print(f"[swap_hairline] 完成,总耗时 {time.time()-t0:.1f}s")
|
||||
return jsonify({"state": 0, "result": result_b64, "steps": steps, "params": params})
|
||||
|
||||
except Exception as e:
|
||||
print(f"[swap_hairline] 失败: {e}")
|
||||
traceback.print_exc()
|
||||
return jsonify({"state": -1, "msg": f"发际线带重绘失败: {e}"}), 500
|
||||
|
||||
|
||||
@app.route("/api/swap_manual", methods=["POST"])
|
||||
def api_swap_manual():
|
||||
"""手绘mask版换发型接口(重绘区域完全由用户手绘mask决定)。
|
||||
|
||||
入参 JSON:
|
||||
img: 原图 base64
|
||||
mask: 用户手绘mask base64(白色=重绘区,与img同分辨率)
|
||||
hair_id: 发型ID
|
||||
其余参数同 swap_debug(strict_mask/seamless_blend/is_hr/dilate_kernel/
|
||||
denoising_strength/blend_dilate/seamless_dilate)
|
||||
返回: {state, result, steps, params}
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
d = request.json
|
||||
img_b64 = d.get("img", "")
|
||||
mask_b64 = d.get("mask", "")
|
||||
hair_id = d.get("hair_id", "")
|
||||
if not img_b64 or not mask_b64:
|
||||
return jsonify({"state": -1, "msg": "img 和 mask 不能为空(请先手绘重绘区域)"}), 400
|
||||
if not hair_id:
|
||||
return jsonify({"state": -1, "msg": "hair_id 不能为空"}), 400
|
||||
|
||||
img = _b64_to_ndarray(img_b64, color=True)
|
||||
hand_mask = _b64_to_ndarray(mask_b64, color=False)
|
||||
if img is None:
|
||||
return jsonify({"state": -1, "msg": "img 解析失败"}), 400
|
||||
if hand_mask is None:
|
||||
return jsonify({"state": -1, "msg": "mask 解析失败"}), 400
|
||||
|
||||
def g(k, default): return d.get(k, default)
|
||||
dk = g("dilate_kernel", [6, 18])
|
||||
bd = g("blend_dilate", [5, 5])
|
||||
sd = g("seamless_dilate", [9, 9])
|
||||
|
||||
print(f"[swap_manual] hair_id={hair_id}, img={img.shape}, mask={hand_mask.shape}, "
|
||||
f"strict_mask={g('strict_mask',False)}, denoising={g('denoising_strength',0.6)}")
|
||||
|
||||
hairstyle_process, landmark_processor = _get_models()
|
||||
from hair_swap_manual import hair_swap_manual
|
||||
task_id = f"manual_{int(time.time()*1000)}"
|
||||
result, steps, params = hair_swap_manual(
|
||||
origin_img=img, hand_mask=hand_mask, hair_id=hair_id,
|
||||
hairstyle_process=hairstyle_process,
|
||||
landmark_processor=landmark_processor,
|
||||
task_id=task_id,
|
||||
strict_mask=bool(g("strict_mask", False)),
|
||||
seamless_blend=bool(g("seamless_blend", True)),
|
||||
is_hr=bool(g("is_hr", True)),
|
||||
dilate_kernel=(int(dk[0]), int(dk[1])),
|
||||
denoising_strength=float(g("denoising_strength", 0.6)),
|
||||
blend_dilate=(int(bd[0]), int(bd[1])),
|
||||
seamless_dilate=(int(sd[0]), int(sd[1])),
|
||||
feather_px=int(g("feather_px", 0)),
|
||||
enhance=bool(g("enhance", False)),
|
||||
enhance_denoising=float(g("enhance_denoising", 0.35)),
|
||||
)
|
||||
|
||||
_, buf = cv2.imencode(".jpg", result, [cv2.IMWRITE_JPEG_QUALITY, 95])
|
||||
result_b64 = base64.b64encode(buf).decode("utf-8")
|
||||
print(f"[swap_manual] 完成,总耗时 {time.time()-t0:.1f}s")
|
||||
return jsonify({"state": 0, "result": result_b64, "steps": steps, "params": params})
|
||||
|
||||
except Exception as e:
|
||||
print(f"[swap_manual] 失败: {e}")
|
||||
traceback.print_exc()
|
||||
return jsonify({"state": -1, "msg": f"手绘mask换发型失败: {e}"}), 500
|
||||
|
||||
|
||||
def _b64_to_ndarray(b64_str, color=True):
|
||||
if "," in b64_str and b64_str.startswith("data:"):
|
||||
b64_str = b64_str.split(",", 1)[1]
|
||||
data = base64.b64decode(b64_str)
|
||||
flag = cv2.IMREAD_COLOR if color else cv2.IMREAD_GRAYSCALE
|
||||
return cv2.imdecode(np.frombuffer(data, np.uint8), flag)
|
||||
|
||||
|
||||
@app.route("/api/hairstyles")
|
||||
def api_hairstyles():
|
||||
"""返回可用发型列表(含性别,用于前端下拉框分组)"""
|
||||
try:
|
||||
from common.logger import config
|
||||
hairstyle_dir = config.get('default', 'hairstyleDir')
|
||||
train_dir = config.get('default', 'train_dir')
|
||||
upload_dir = config.get('default', 'upload_train_dir')
|
||||
from hair_grow_swap import list_hairstyles
|
||||
styles = list_hairstyles(hairstyle_dir, train_dir, upload_dir)
|
||||
boy = [s for s in styles if s["gender"] == "boy"]
|
||||
girl = [s for s in styles if s["gender"] == "girl"]
|
||||
print(f"[hairstyles] 共 {len(styles)} 个可用发型 (boy={len(boy)}, girl={len(girl)})")
|
||||
return jsonify({"state": 0, "data": styles, "count": len(styles)})
|
||||
except Exception as e:
|
||||
traceback.print_exc()
|
||||
return jsonify({"state": -1, "msg": str(e)}), 500
|
||||
|
||||
|
||||
@app.route("/preview/<hair_id>")
|
||||
def preview_img(hair_id):
|
||||
"""直接返回预览图文件(比 base64 API 快,浏览器可缓存)"""
|
||||
effect_path = os.path.join(BASE_DIR, "static", "previews", f"{hair_id}.jpg")
|
||||
if os.path.exists(effect_path):
|
||||
return send_from_directory(os.path.join(BASE_DIR, "static", "previews"), f"{hair_id}.jpg")
|
||||
# 回退到 ref_rgb
|
||||
try:
|
||||
from common.logger import config
|
||||
hairstyle_dir = config.get('default', 'hairstyleDir')
|
||||
fallback = os.path.join(hairstyle_dir, hair_id, "ref_rgb_8uc3_768.png")
|
||||
if os.path.exists(fallback):
|
||||
return send_file_or_404(fallback)
|
||||
except Exception:
|
||||
pass
|
||||
return ("", 404)
|
||||
|
||||
|
||||
@app.route("/train_src/<hair_id>")
|
||||
def train_src_img(hair_id):
|
||||
"""返回发型的训练原图(hair_type_images/<hair_id>.jpg/.png)。
|
||||
用于测试页展示发型真实样子,而非套在标准脸上的效果图。
|
||||
"""
|
||||
src_dir = "/home/xsl/change_hair/hair_type_images"
|
||||
for ext in (".jpg", ".jpeg", ".png"):
|
||||
path = os.path.join(src_dir, hair_id + ext)
|
||||
if os.path.exists(path):
|
||||
return send_from_directory(src_dir, hair_id + ext)
|
||||
return ("", 404)
|
||||
|
||||
|
||||
@app.route("/api/hairstyle_preview/<hair_id>")
|
||||
def api_hairstyle_preview(hair_id):
|
||||
"""返回某发型的预览图。
|
||||
优先用生成的效果图 static/previews/<hair_id>.jpg(发型套在标准脸上的样子),
|
||||
没有则回退到 first## 原始上传图。
|
||||
"""
|
||||
try:
|
||||
# 1. 优先:生成的效果图(发型套在 boy/girl 标准脸上)
|
||||
preview = None
|
||||
effect_path = os.path.join(BASE_DIR, "static", "previews", f"{hair_id}.jpg")
|
||||
if os.path.exists(effect_path):
|
||||
preview = effect_path
|
||||
else:
|
||||
# 2. 回退:first## 原始上传图
|
||||
from common.logger import config
|
||||
upload_dir = config.get('default', 'upload_train_dir')
|
||||
save_dir = os.path.join(upload_dir, hair_id)
|
||||
if os.path.isdir(save_dir):
|
||||
for name in os.listdir(save_dir):
|
||||
if "first##" in name:
|
||||
preview = os.path.join(save_dir, name)
|
||||
break
|
||||
# 3. 再回退:ref_rgb
|
||||
if not preview or not os.path.exists(preview):
|
||||
hairstyle_dir = config.get('default', 'hairstyleDir')
|
||||
fallback = os.path.join(hairstyle_dir, hair_id, "ref_rgb_8uc3_768.png")
|
||||
if os.path.exists(fallback):
|
||||
preview = fallback
|
||||
if not preview:
|
||||
return jsonify({"state": -1, "msg": "无预览图"}), 404
|
||||
img = cv2.imread(preview)
|
||||
img = cv2.resize(img, (256, 256), interpolation=cv2.INTER_AREA)
|
||||
_, buf = cv2.imencode(".jpg", img, [cv2.IMWRITE_JPEG_QUALITY, 85])
|
||||
b64 = base64.b64encode(buf).decode()
|
||||
return jsonify({"state": 0, "preview": b64})
|
||||
except Exception as e:
|
||||
return jsonify({"state": -1, "msg": str(e)}), 500
|
||||
|
||||
|
||||
@app.route("/api/grow", methods=["POST"])
|
||||
def api_grow():
|
||||
"""生发接口(走换发型工作流)
|
||||
|
||||
入参 JSON:
|
||||
img: 原图 base64
|
||||
mask: 手绘遮罩 base64(白=生发区,与img同分辨率)
|
||||
hair_id: 选择的发型ID(必填)
|
||||
is_hr: 是否高清 "true"/"false",默认 "true"
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
data = request.json
|
||||
img_b64 = data.get("img", "")
|
||||
mask_b64 = data.get("mask", "")
|
||||
hair_id = data.get("hair_id", "")
|
||||
is_hr = str(data.get("is_hr", "true")).lower() == "true"
|
||||
|
||||
if not img_b64 or not mask_b64:
|
||||
return jsonify({"state": -1, "msg": "img 和 mask 不能为空"}), 400
|
||||
if not hair_id:
|
||||
return jsonify({"state": -1, "msg": "hair_id 不能为空(请先选择发型)"}), 400
|
||||
|
||||
img = _b64_to_ndarray(img_b64, color=True)
|
||||
mask = _b64_to_ndarray(mask_b64, color=False)
|
||||
if img is None:
|
||||
return jsonify({"state": -1, "msg": "img 解析失败"}), 400
|
||||
if mask is None:
|
||||
return jsonify({"state": -1, "msg": "mask 解析失败"}), 400
|
||||
|
||||
print(f"[grow] img={img.shape}, mask={mask.shape}, hair_id={hair_id}, is_hr={is_hr}")
|
||||
|
||||
# 加载模型(首次慢)
|
||||
hairstyle_process, landmark_processor = _get_models()
|
||||
|
||||
# 调用生发(走换发型工作流)
|
||||
from hair_grow_swap import hair_grow_swap
|
||||
task_id = f"grow_{int(time.time()*1000)}"
|
||||
result = hair_grow_swap(
|
||||
origin_img=img, hand_mask=mask, hair_id=hair_id,
|
||||
hairstyle_process=hairstyle_process,
|
||||
landmark_processor=landmark_processor,
|
||||
task_id=task_id, is_hr=is_hr)
|
||||
|
||||
_, buf = cv2.imencode(".jpg", result, [cv2.IMWRITE_JPEG_QUALITY, 95])
|
||||
result_b64 = base64.b64encode(buf).decode("utf-8")
|
||||
print(f"[grow] 完成,总耗时 {time.time()-t0:.1f}s")
|
||||
return jsonify({"state": 0, "result": result_b64})
|
||||
|
||||
except Exception as e:
|
||||
print(f"[grow] 失败: {e}")
|
||||
traceback.print_exc()
|
||||
return jsonify({"state": -1, "msg": f"生发失败: {e}"}), 500
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print(f"生发服务启动,端口 {PORT}")
|
||||
print(f"测试页面: http://0.0.0.0:{PORT}")
|
||||
print(f"注意:首次请求会加载换发型模型(约30-60秒)")
|
||||
server = pywsgi.WSGIServer(("0.0.0.0", PORT), app)
|
||||
server.serve_forever()
|
||||
Reference in New Issue
Block a user