Files
hair/local_test/bench2.py
xsl b61ea6f33b feat(接口5): 生发机制对齐接口2 + 统一ComfyUI重绘提示词
接口5改造(生发机制与接口2一致,按性别分流):
- generate_hairline_pngs: 生发图来源从 _grow_from_texture(局部inpaint)
  改为按性别分流——female走 generate_grow_results_swap(swapHair+Flux整帧重绘),
  male走 generate_grow_results(ComfyUI add_hair)
- 新增参数 redraw_max_side/unet_name/v2_defaults(female路径透传)
- 接口5 handler 加 flux_model/redraw_max_side Form参数
- 保留接口5独有输出: 3档叠图(middle/high/low)/中心点/face_measure
- 已验证: female日志出现"接口2女 管线降分辨率max_side=640"+swap+Flux;
  male走add_hair.json; generate_grow_image=false正确跳过生发

统一ComfyUI重绘提示词:
- "填充遮罩区域的头发,皮肤加一点磨皮,再加一点美颜" → "填充遮罩区域的头发"
- 覆盖: _REDRAW_PROMPT/_DEFAULT_PROMPT常量 + app.py各接口Form默认 +
  工作流JSON节点60(add_hair/0716add-hair-api/hair_repaint) + 全部测试页输入框
2026-07-27 23:36:15 +08:00

70 lines
2.6 KiB
Python

#!/usr/bin/env python3
"""Steps sweep + resolution test on the working 9B fp8-fast config."""
import io, time
import requests
import numpy as np
from PIL import Image, ImageFilter
import app as A
COMFY = "http://127.0.0.1:8188"
MODEL = "flux2.0/flux-2-klein-9b-fp8.safetensors"
DTYPE = "fp8_e4m3fn_fast"
def upload(scale=1.0):
image = Image.open("用来重绘.jpg").convert("RGB")
mask_img = Image.open("用来重绘.png").convert("RGBA")
if scale != 1.0:
w, h = image.size
w, h = int(w * scale) // 8 * 8, int(h * scale) // 8 * 8
image = image.resize((w, h), Image.LANCZOS)
mask_data = np.max(np.array(mask_img), axis=2)
m = Image.fromarray(mask_data, mode="L")
if m.size != image.size:
m = m.resize(image.size, Image.LANCZOS)
m = m.filter(ImageFilter.GaussianBlur(radius=4))
alpha = Image.eval(m, lambda x: 255 - x)
r, g, b = image.split()
rgba = Image.merge("RGBA", (r, g, b, alpha))
buf = io.BytesIO(); rgba.save(buf, format="PNG"); buf.seek(0)
up = requests.post(f"{COMFY}/upload/image",
files={"image": ("hair_input.png", buf, "image/png")}).json()
return up["name"], image.size
def run(fname, steps):
wf = A.build_workflow(fname, "填充遮罩区域的头发")
wf["16"]["inputs"]["unet_name"] = MODEL
wf["16"]["inputs"]["weight_dtype"] = DTYPE
wf["1"]["inputs"]["steps"] = steps
pid = requests.post(f"{COMFY}/prompt", json={"prompt": wf}).json()["prompt_id"]
dl = time.time() + 120
while time.time() < dl:
time.sleep(0.1)
h = requests.get(f"{COMFY}/history/{pid}").json()
if pid in h and "17" in h[pid].get("outputs", {}):
ts = {m[0]: m[1].get("timestamp") for m in h[pid]["status"]["messages"]}
return (ts["execution_success"] - ts["execution_start"]) / 1000.0
return float("nan")
print("=== 步数扫描 (9B fp8-fast, 原分辨率 1024x775) ===", flush=True)
fname, sz = upload(1.0)
run(fname, 6) # warmup
res = {}
for s in [2, 3, 4, 6, 8]:
t = min(run(fname, s), run(fname, s))
res[s] = t
print(f" steps={s}: {t:.2f}s", flush=True)
# derive per-step cost & fixed overhead via two points
per = (res[8] - res[2]) / 6
fixed = res[2] - per * 2
print(f" -> 每步 ~{per:.3f}s, 固定开销(VAE/编码/colormatch/加载) ~{fixed:.2f}s", flush=True)
print("\n=== 分辨率影响 (steps=4) ===", flush=True)
for scale in [1.0, 0.75, 0.6]:
fn, s2 = upload(scale)
run(fn, 4) # warmup
t = min(run(fn, 4), run(fn, 4))
print(f" {s2[0]}x{s2[1]} ({s2[0]*s2[1]/1e6:.2f}MP): {t:.2f}s", flush=True)