Files
hair/benchmark_rotation.py
T
UbuntuandCursor 87ff2c15d0 chore: 精简生发提示词,去掉磨皮/美颜要求
统一改为「填充遮罩区域的头发」,涉及后端默认值、ComfyUI 工作流 JSON、
测试页、benchmark 脚本、local_test。

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-07-29 15:57:57 +08:00

272 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""轮换测试:接口1 → 接口2女 → 接口2男 → 接口3,共5轮。
监控每次调用的耗时和显存变化,特别关注模型是否被换出显存。
"""
import json, os, time, base64, subprocess, urllib.request, urllib.error, mimetypes
from datetime import datetime
from pathlib import Path
API_BASE = "http://127.0.0.1:8187"
TOKEN = "dev-shared-secret-2026"
GIRL_IMG = "/home/ubuntu/hair/image/girl_img/girl13.jpg"
OUT_DIR = "/home/ubuntu/hair/benchmark_out/rotation"
RESULT_JSON = os.path.join(OUT_DIR, "rotation_results.json")
REPORT_HTML = os.path.join(OUT_DIR, "rotation_report.html")
ROUNDS = 5
os.makedirs(OUT_DIR, exist_ok=True)
def get_vram():
try:
out = subprocess.check_output(
["nvidia-smi", "--query-gpu=memory.used,memory.free",
"--format=csv,noheader,nounits"], text=True, timeout=5).strip()
used, free = out.split(",")
return int(used.strip()), int(free.strip())
except Exception:
return -1, -1
def get_gpu_procs():
try:
out = subprocess.check_output(
["nvidia-smi", "--query-compute-apps=pid,used_memory",
"--format=csv,noheader,nounits"], text=True, timeout=5).strip()
procs = {}
for line in out.splitlines():
parts = line.split(",")
if len(parts) >= 2:
procs[parts[0].strip()] = int(parts[1].strip())
return procs
except Exception:
return {}
def _multipart(fields, files=None):
boundary = "----RotTest" + str(int(time.time() * 1000))
parts = []
for k, v in fields.items():
parts.append(f"--{boundary}\r\n".encode())
parts.append(f'Content-Disposition: form-data; name="{k}"\r\n\r\n'.encode())
parts.append(str(v).encode())
parts.append(b"\r\n")
if files:
for field_name, (filename, data, mime) in files.items():
parts.append(f"--{boundary}\r\n".encode())
parts.append(f'Content-Disposition: form-data; name="{field_name}"; filename="{filename}"\r\n'.encode())
parts.append(f"Content-Type: {mime}\r\n\r\n".encode())
parts.append(data)
parts.append(b"\r\n")
parts.append(f"--{boundary}--\r\n".encode())
body = b"".join(parts)
return body, boundary
def call_api1(image_path):
"""接口1:四庭七眼测量"""
with open(image_path, "rb") as f:
img_data = f.read()
body, boundary = _multipart(
{"image_url": ""},
{"image_file": (os.path.basename(image_path), img_data, "image/jpeg")}
)
req = urllib.request.Request(f"{API_BASE}/api/v1/face/measure", data=body, method="POST")
req.add_header("Content-Type", f"multipart/form-data; boundary={boundary}")
req.add_header("X-Internal-Token", TOKEN)
t0 = time.time()
try:
with urllib.request.urlopen(req, timeout=120) as resp:
data = json.loads(resp.read())
return data, time.time() - t0, None
except Exception as e:
return {}, time.time() - t0, str(e)
def call_api2(image_path, gender, hair_style="1"):
"""接口2C端生发"""
with open(image_path, "rb") as f:
img_data = f.read()
fields = {"gender": gender, "hair_style": hair_style, "use_mask": "0",
"prompt": "填充遮罩区域的头发"}
body, boundary = _multipart(fields, {"image_file": (os.path.basename(image_path), img_data, "image/jpeg")})
req = urllib.request.Request(f"{API_BASE}/api/v1/hair/grow", data=body, method="POST")
req.add_header("Content-Type", f"multipart/form-data; boundary={boundary}")
req.add_header("X-Internal-Token", TOKEN)
t0 = time.time()
try:
with urllib.request.urlopen(req, timeout=600) as resp:
data = json.loads(resp.read())
return data, time.time() - t0, None
except Exception as e:
return {}, time.time() - t0, str(e)
def call_api3(image_path):
"""接口3B端生发(use_mask=False,直接送图)"""
with open(image_path, "rb") as f:
img_data = f.read()
fields = {"use_mask": "true", "prompt": "填充遮罩区域的头发"}
body, boundary = _multipart(fields, {"marked_image_file": (os.path.basename(image_path), img_data, "image/jpeg")})
req = urllib.request.Request(f"{API_BASE}/api/v1/hair/grow-b", data=body, method="POST")
req.add_header("Content-Type", f"multipart/form-data; boundary={boundary}")
req.add_header("X-Internal-Token", TOKEN)
t0 = time.time()
try:
with urllib.request.urlopen(req, timeout=600) as resp:
data = json.loads(resp.read())
return data, time.time() - t0, None
except Exception as e:
return {}, time.time() - t0, str(e)
def run_test():
results = []
tests = [
("iface1", "接口1-测量", lambda img: call_api1(img)),
("iface2f", "接口2-女-椭圆", lambda img: call_api2(img, "female", "1")),
("iface2m", "接口2-男-椭圆", lambda img: call_api2(img, "male", "1")),
("iface3", "接口3-B端生发", lambda img: call_api3(img)),
]
total = len(tests) * ROUNDS
print(f"轮换测试:{len(tests)} 个接口 × {ROUNDS} 轮 = {total} 次调用")
print(f"测试图片:{GIRL_IMG}\n")
for rnd in range(1, ROUNDS + 1):
print(f"\n{'='*60}")
print(f"第 {rnd}/{ROUNDS} 轮")
print(f"{'='*60}")
for tid, tname, fn in tests:
print(f"\n[{rnd}/{ROUNDS}] {tname} ...")
vram_before, vram_free_before = get_vram()
procs_before = get_gpu_procs()
data, elapsed, error = fn(GIRL_IMG)
vram_after, vram_free_after = get_vram()
procs_after = get_gpu_procs()
# 判断成功
code = data.get("code", -1) if data else -1
ok = (code == 0)
# VRAM 变化
vram_delta = vram_after - vram_before
procs_changed = False
for pid in set(list(procs_before.keys()) + list(procs_after.keys())):
before_mb = procs_before.get(pid, 0)
after_mb = procs_after.get(pid, 0)
if abs(after_mb - before_mb) > 10:
procs_changed = True
break
record = {
"round": rnd,
"test_id": tid,
"test_name": tname,
"timestamp": datetime.now().strftime("%H:%M:%S"),
"elapsed_s": round(elapsed, 2),
"success": ok,
"code": code,
"error": error,
"vram_before_mb": vram_before,
"vram_after_mb": vram_after,
"vram_free_before_mb": vram_free_before,
"vram_free_after_mb": vram_free_after,
"vram_delta_mb": vram_delta,
"procs_before": procs_before,
"procs_after": procs_after,
"procs_changed": procs_changed,
}
results.append(record)
status = "✓" if ok else "✗"
warn = " ⚠模型换出!" if procs_changed else ""
print(f" → {status} code={code} time={elapsed:.1f}s "
f"vram={vram_before}{vram_after}MB (Δ{vram_delta:+d}){warn}")
with open(RESULT_JSON, "w") as f:
json.dump(results, f, indent=2, ensure_ascii=False)
generate_html(results)
generate_html(results)
print(f"\n✅ 完成!{RESULT_JSON}")
print(f"📄 报告:{REPORT_HTML}")
def generate_html(results):
total = len(results)
success = sum(1 for r in results if r["success"])
times = [r["elapsed_s"] for r in results if r["success"]]
avg_time = sum(times) / len(times) if times else 0
# 按接口分组
by_test = {}
for r in results:
tid = r["test_id"]
if tid not in by_test:
by_test[tid] = {"name": r["test_name"], "times": [], "vram_deltas": [], "procs_changed": 0}
if r["success"]:
by_test[tid]["times"].append(r["elapsed_s"])
by_test[tid]["vram_deltas"].append(r["vram_delta_mb"])
if r["procs_changed"]:
by_test[tid]["procs_changed"] += 1
html = f"""<!DOCTYPE html>
<html lang="zh-CN">
<head><meta charset="UTF-8"><title>接口轮换测试报告</title>
<style>
*{{margin:0;padding:0;box-sizing:border-box}}
body{{font-family:sans-serif;background:#1a1a2e;color:#e0e0e0;padding:20px}}
h1{{text-align:center;color:#00d4ff;margin-bottom:10px}}
.subtitle{{text-align:center;color:#888;margin-bottom:30px;font-size:14px}}
.summary{{display:grid;grid-template-columns:repeat(auto-fit,minmax(160px,1fr));gap:15px;margin-bottom:30px}}
.card{{background:#16213e;border-radius:12px;padding:20px;text-align:center;border:1px solid #333}}
.card .num{{font-size:28px;font-weight:700}}
.card .label{{font-size:12px;color:#888;margin-top:5px}}
.card.ok .num{{color:#0f0}} .card.err .num{{color:#f44}} .card.time .num{{color:#00d4ff}}
table{{width:100%;border-collapse:collapse;margin-bottom:30px;background:#16213e;border-radius:12px;overflow:hidden}}
th{{background:#0f3460;padding:10px 8px;text-align:center;font-size:13px;color:#fff}}
td{{padding:8px;text-align:center;border-bottom:1px solid #222;font-size:13px}}
tr:hover{{background:#1a1a3e}}
.section-title{{font-size:18px;font-weight:600;margin:30px 0 15px;color:#00d4ff;border-bottom:1px solid #333;padding-bottom:10px}}
.ok{{color:#0f0}} .fail{{color:#f44}} .warn{{color:#ff0}}
.time-bar{{display:inline-block;height:18px;background:linear-gradient(90deg,#0f3460,#00d4ff);border-radius:3px;vertical-align:middle;min-width:2px}}
</style></head><body>
<h1>📊 接口轮换测试报告</h1>
<div class="subtitle">{datetime.now().strftime('%Y-%m-%d %H:%M:%S')} {total} 次调用 {ROUNDS} 轮</div>
<div class="summary">
<div class="card ok"><div class="num">{success}</div><div class="label">成功</div></div>
<div class="card err"><div class="num">{total-success}</div><div class="label">失败</div></div>
<div class="card time"><div class="num">{avg_time:.1f}s</div><div class="label">平均耗时</div></div>
<div class="card time"><div class="num">{sum(t['procs_changed'] for t in by_test.values())}</div><div class="label">显存抖动次数</div></div>
</div>
<div class="section-title">按接口分组统计</div>
<table><thead><tr><th>接口</th><th>调用次数</th><th>平均耗时</th><th>最快</th><th>最慢</th><th>显存变化</th><th>模型换出</th></tr></thead><tbody>
"""
for tid, info in by_test.items():
times = info["times"]
avg = sum(times) / len(times) if times else 0
mn = min(times) if times else 0
mx = max(times) if times else 0
warn = f'<span class="warn">⚠ {info["procs_changed"]}次</span>' if info["procs_changed"] else "无"
html += f'<tr><td>{info["name"]}</td><td>{len(times)}</td><td>{avg:.1f}s</td><td>{mn:.1f}s</td><td>{mx:.1f}s</td><td>Δ{info["vram_deltas"][-1]:+d}MB</td><td>{warn}</td></tr>'
html += '</tbody></table><div class="section-title">详细结果</div><table><thead><tr><th>轮次</th><th>接口</th><th>耗时</th><th>显存前</th><th>显存后</th><th>Δ</th><th>模型换出</th><th>状态</th></tr></thead><tbody>'
for r in results:
status = '<span class="ok">✓</span>' if r["success"] else f'<span class="fail">✗ {str(r.get("error",""))[:30]}</span>'
warn = '<span class="warn">⚠</span>' if r["procs_changed"] else ""
html += f'<tr><td>{r["round"]}</td><td>{r["test_name"]}</td><td><div class="time-bar" style="width:{min(r["elapsed_s"]*3,200)}px"></div> {r["elapsed_s"]:.1f}s</td><td>{r["vram_before_mb"]}MB</td><td>{r["vram_after_mb"]}MB</td><td>{r["vram_delta_mb"]:+d}</td><td>{warn}</td><td>{status}</td></tr>'
html += '</tbody></table></body></html>'
with open(REPORT_HTML, "w", encoding="utf-8") as f:
f.write(html)
if __name__ == "__main__":
run_test()