save code
This commit is contained in:
+67
-10
@@ -4,11 +4,27 @@ import io
|
||||
import json
|
||||
import time
|
||||
import random
|
||||
import logging
|
||||
import traceback
|
||||
import requests
|
||||
from flask import Flask, request, jsonify, send_file
|
||||
import numpy as np
|
||||
from PIL import Image, ImageFilter
|
||||
|
||||
# 让 PIL 支持 iPhone 的 HEIC/HEIF 照片(浏览器 accept="image/*" 会允许选中它们)。
|
||||
try:
|
||||
from pillow_heif import register_heif_opener
|
||||
register_heif_opener()
|
||||
_HEIF_OK = True
|
||||
except Exception:
|
||||
_HEIF_OK = False
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
)
|
||||
log = logging.getLogger("hair")
|
||||
|
||||
app = Flask(__name__)
|
||||
COMFYUI_URL = "http://127.0.0.1:8188"
|
||||
|
||||
@@ -18,6 +34,7 @@ _CORS_HEADERS = {
|
||||
"Access-Control-Allow-Origin": "*",
|
||||
"Access-Control-Allow-Methods": "POST, OPTIONS",
|
||||
"Access-Control-Allow-Headers": "Content-Type",
|
||||
"Access-Control-Expose-Headers": "X-Generate-Time",
|
||||
}
|
||||
|
||||
|
||||
@@ -43,7 +60,7 @@ def build_workflow(image_filename, prompt_text, seed=None):
|
||||
# Loaders
|
||||
"16": {"class_type": "UNETLoader", "inputs": {
|
||||
"unet_name": "flux2.0/flux-2-klein-9b-fp8.safetensors",
|
||||
"weight_dtype": "fp8_e4m3fn"}},
|
||||
"weight_dtype": "fp8_e4m3fn_fast"}},
|
||||
"3": {"class_type": "VAELoader", "inputs": {
|
||||
"vae_name": "flux2-vae.safetensors"}},
|
||||
"61": {"class_type": "CLIPLoader", "inputs": {
|
||||
@@ -136,7 +153,7 @@ def build_workflow(image_filename, prompt_text, seed=None):
|
||||
"1": {"class_type": "BasicScheduler", "inputs": {
|
||||
"model": ["2", 0],
|
||||
"scheduler": "simple",
|
||||
"steps": 6,
|
||||
"steps": 4,
|
||||
"denoise": 1}},
|
||||
"20": {"class_type": "BasicGuider", "inputs": {
|
||||
"model": ["2", 0],
|
||||
@@ -181,20 +198,42 @@ def index():
|
||||
|
||||
@app.route("/api/generate", methods=["POST"])
|
||||
def generate():
|
||||
t_start = time.time()
|
||||
try:
|
||||
if "image" not in request.files or "mask" not in request.files:
|
||||
msg = f"缺少上传文件: files={list(request.files.keys())}"
|
||||
log.warning(msg)
|
||||
return jsonify({"error": msg}), 400
|
||||
image_file = request.files["image"]
|
||||
mask_file = request.files["mask"]
|
||||
prompt_text = request.form.get("prompt", "填充遮罩区域的头发,皮肤加一点磨皮")
|
||||
log.info(
|
||||
"收到请求: image=%s mask=%s prompt=%r",
|
||||
image_file.filename, mask_file.filename, prompt_text,
|
||||
)
|
||||
|
||||
# Load original image as RGB
|
||||
image = Image.open(image_file).convert("RGB")
|
||||
try:
|
||||
image = Image.open(image_file).convert("RGB")
|
||||
except Exception as e:
|
||||
log.error("无法解码人物图片 %s: %s", image_file.filename, e)
|
||||
hint = "" if _HEIF_OK else "(当前不支持 HEIC)"
|
||||
return jsonify({
|
||||
"error": f"无法识别人物图片格式{hint},请改用 JPG/PNG: {e}"
|
||||
}), 400
|
||||
|
||||
# Load mask and extract mask data from ALL channels (R, G, B, A)
|
||||
# This handles different mask formats:
|
||||
# - Red mask (R=255 where drawn): user-provided PNG
|
||||
# - White mask (R=G=B=255 where drawn): frontend canvas
|
||||
# - Alpha mask (A=255 where drawn): transparent brush
|
||||
mask_img = Image.open(mask_file).convert("RGBA")
|
||||
try:
|
||||
mask_img = Image.open(mask_file).convert("RGBA")
|
||||
except Exception as e:
|
||||
log.error("无法解码遮罩图片 %s: %s", mask_file.filename, e)
|
||||
return jsonify({
|
||||
"error": f"无法识别遮罩图片格式,请改用 JPG/PNG: {e}"
|
||||
}), 400
|
||||
mask_arr = np.array(mask_img)
|
||||
# Use max of all channels: 255 where any color/alpha is drawn, 0 where empty
|
||||
mask_data = np.max(mask_arr, axis=2) # (H, W) uint8
|
||||
@@ -229,6 +268,7 @@ def generate():
|
||||
)
|
||||
upload_data = upload_resp.json()
|
||||
if "name" not in upload_data:
|
||||
log.error("ComfyUI 上传图片失败: %s", upload_data)
|
||||
return jsonify({"error": f"Upload failed: {upload_data}"}), 500
|
||||
image_filename = upload_data["name"]
|
||||
|
||||
@@ -241,12 +281,16 @@ def generate():
|
||||
)
|
||||
prompt_data = prompt_resp.json()
|
||||
if "error" in prompt_data:
|
||||
log.error(
|
||||
"ComfyUI /prompt 校验失败: error=%s node_errors=%s",
|
||||
prompt_data.get("error"), prompt_data.get("node_errors"),
|
||||
)
|
||||
return jsonify({"error": json.dumps(prompt_data["error"], ensure_ascii=False)}), 500
|
||||
prompt_id = prompt_data["prompt_id"]
|
||||
|
||||
# Poll for completion (5 min timeout)
|
||||
for _ in range(150):
|
||||
time.sleep(2)
|
||||
# Poll for completion (5 min timeout, 0.1s interval)
|
||||
for _ in range(3000):
|
||||
time.sleep(0.1)
|
||||
history_resp = requests.get(
|
||||
f"{COMFYUI_URL}/history/{prompt_id}", timeout=10
|
||||
)
|
||||
@@ -254,7 +298,14 @@ def generate():
|
||||
if prompt_id in history_data:
|
||||
status = history_data[prompt_id].get("status", {})
|
||||
if status.get("status_str") == "error":
|
||||
return jsonify({"error": "Workflow execution failed"}), 500
|
||||
log.error(
|
||||
"ComfyUI 工作流执行失败: %s",
|
||||
json.dumps(status, ensure_ascii=False),
|
||||
)
|
||||
return jsonify({
|
||||
"error": "Workflow execution failed",
|
||||
"detail": status.get("messages", status),
|
||||
}), 500
|
||||
outputs = history_data[prompt_id].get("outputs", {})
|
||||
if "17" in outputs: # SaveImage node
|
||||
image_info = outputs["17"]["images"][0]
|
||||
@@ -266,16 +317,22 @@ def generate():
|
||||
params={"filename": filename, "subfolder": subfolder, "type": img_type},
|
||||
timeout=30,
|
||||
)
|
||||
return send_file(
|
||||
elapsed = time.time() - t_start
|
||||
log.info("重绘完成,服务端耗时 %.2fs", elapsed)
|
||||
resp = send_file(
|
||||
io.BytesIO(view_resp.content), mimetype="image/png"
|
||||
)
|
||||
resp.headers["X-Generate-Time"] = f"{elapsed:.2f}"
|
||||
return resp
|
||||
|
||||
return jsonify({"error": "Timeout: workflow did not complete in 5 minutes"}), 500
|
||||
|
||||
except requests.ConnectionError:
|
||||
log.error("无法连接 ComfyUI @ %s", COMFYUI_URL)
|
||||
return jsonify({"error": "Cannot connect to ComfyUI at " + COMFYUI_URL + ". Is it running?"}), 503
|
||||
except Exception as e:
|
||||
return jsonify({"error": str(e)}), 500
|
||||
log.error("生成失败,未捕获异常:\n%s", traceback.format_exc())
|
||||
return jsonify({"error": f"{type(e).__name__}: {e}"}), 500
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user