Files
hair/benchmark_grow.py
T
2026-07-18 15:30:31 +08:00

451 lines
19 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
"""基准测试:19张图片 × 5种发际线,记录每步耗时和显存变化。
用法: python3 benchmark_grow.py
输出: benchmark_results.json + benchmark_report.html
"""
import json, os, time, base64, subprocess, re, glob, shutil
from datetime import datetime
from pathlib import Path
import urllib.request, urllib.error
API_URL = "http://127.0.0.1:8187/api/v1/hair/grow"
TOKEN = "dev-shared-secret-2026"
IMG_DIR = "/home/ubuntu/hair/image/girl_img"
OUT_DIR = "/home/ubuntu/hair/benchmark_out"
RESULT_JSON = os.path.join(OUT_DIR, "benchmark_results.json")
REPORT_HTML = os.path.join(OUT_DIR, "benchmark_report.html")
HAIRLINES = [
("1", "ellipse", "椭圆形"),
("2", "flower", "花瓣形"),
("3", "heart", "心形"),
("4", "straight", "直线形"),
("5", "wave", "波浪形"),
]
os.makedirs(OUT_DIR, exist_ok=True)
# ── 工具函数 ──────────────────────────────────────────────
def get_vram():
"""返回 (used_MB, free_MB)"""
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():
"""返回各进程显存占用 dict"""
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 read_worker_log_tail(n=50):
"""读取 worker.log 最后 n 行"""
log_path = "/home/ubuntu/hair/worker.log"
try:
result = subprocess.run(["tail", "-n", str(n), log_path],
capture_output=True, text=True, timeout=5)
return result.stdout
except Exception:
return ""
def call_api(image_path, hair_style_value):
"""调用接口2,返回 (json_dict, elapsed_sec, error_or_None)"""
import mimetypes
boundary = "----BenchmarkBoundary" + str(int(time.time()*1000))
filename = os.path.basename(image_path)
mime = mimetypes.guess_type(image_path)[0] or "image/jpeg"
with open(image_path, "rb") as f:
img_data = f.read()
body_parts = []
body_parts.append(f"--{boundary}\r\n".encode())
body_parts.append(f'Content-Disposition: form-data; name="image_file"; filename="{filename}"\r\n'.encode())
body_parts.append(f"Content-Type: {mime}\r\n\r\n".encode())
body_parts.append(img_data)
body_parts.append(f"\r\n--{boundary}\r\n".encode())
body_parts.append(b'Content-Disposition: form-data; name="gender"\r\n\r\n')
body_parts.append(b"female")
body_parts.append(f"\r\n--{boundary}\r\n".encode())
body_parts.append(b'Content-Disposition: form-data; name="hair_style"\r\n\r\n')
body_parts.append(hair_style_value.encode())
body_parts.append(f"\r\n--{boundary}\r\n".encode())
body_parts.append(b'Content-Disposition: form-data; name="use_mask"\r\n\r\n')
body_parts.append(b"0")
body_parts.append(f"\r\n--{boundary}\r\n".encode())
body_parts.append(b'Content-Disposition: form-data; name="prompt"\r\n\r\n')
body_parts.append(b"\xe5\xa1\xab\xe5\x85\x85\xe9\x81\xae\xe7\xbd\xa9\xe5\x8c\xba\xe5\x9f\x9f\xe7\x9a\x84\xe5\xa4\xb4\xe5\x8f\x91\xef\xbc\x8c\xe7\x9a\xae\xe8\x82\xa4\xe5\x8a\xa0\xe4\xb8\x80\xe7\x82\xb9\xe7\xa3\xa8\xe7\x9a\xae")
body_parts.append(f"\r\n--{boundary}--\r\n".encode())
body = b"".join(body_parts)
req = urllib.request.Request(API_URL, 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:
raw = resp.read()
elapsed = time.time() - t0
data = json.loads(raw)
return data, elapsed, None
except urllib.error.HTTPError as e:
elapsed = time.time() - t0
try:
data = json.loads(e.read())
except Exception:
data = {"error": str(e)}
return data, elapsed, f"HTTP {e.code}"
except Exception as e:
elapsed = time.time() - t0
return {}, elapsed, str(e)
def parse_step_timing(log_text):
"""从 worker.log 文本中解析步骤耗时"""
timing = {}
for line in log_text.splitlines():
# 匹配 "步骤1 遮罩完成 耗时=949ms"
m = re.search(r"步骤(\d+)\s+\S+\s+耗时=(\d+)ms", line)
if m:
timing[f"step{m.group(1)}_ms"] = int(m.group(2))
# 匹配 "人脸检出 px_per_cm=30.669 图尺寸=813x967"
if "人脸检出" in line:
m2 = re.search(r"px_per_cm=([\d.]+)", line)
if m2:
timing["px_per_cm"] = float(m2.group(1))
# 匹配 "头发分割完成"
if "头发分割完成" in line:
timing["hair_seg"] = True
# 匹配 "换发型图失败"
if "换发型图失败" in line or "换发型服务不可达" in line:
timing["swap_error"] = line.strip()[-100:]
# 匹配 "重绘" 相关
if "重绘" in line and ("完成" in line or "失败" in line):
timing["redraw_status"] = "完成" if "完成" in line else "失败"
return timing
# ── 主流程 ────────────────────────────────────────────────
def run_benchmark():
images = sorted(glob.glob(os.path.join(IMG_DIR, "girl*.jpg")))
print(f"找到 {len(images)} 张图片")
# 加载已有结果(支持断点续跑)
results = []
if os.path.exists(RESULT_JSON):
with open(RESULT_JSON) as f:
results = json.load(f)
print(f"已有 {len(results)} 条记录,继续未完成的测试")
total_calls = len(images) * len(HAIRLINES)
done = len(results)
print(f"总计 {total_calls} 次调用,已完成 {done},剩余 {total_calls - done}")
log_offset = 0
try:
log_offset = subprocess.check_output(["wc", "-l", "/home/ubuntu/hair/worker.log"],
text=True, timeout=5).split()[0]
log_offset = int(log_offset)
except Exception:
pass
for img_idx, img_path in enumerate(images):
img_name = os.path.basename(img_path)
# 跳过已完成的图片
img_results = [r for r in results if r["image"] == img_name]
if len(img_results) >= len(HAIRLINES):
print(f"[{img_idx+1}/{len(images)}] {img_name} 已完成,跳过")
continue
for hl_id, hl_key, hl_label in HAIRLINES:
# 跳过已完成的
existing = [r for r in results if r["image"] == img_name and r["hairline_id"] == hl_id]
if existing:
continue
print(f"\n[{img_idx+1}/{len(images)}] {img_name}{hl_label}({hl_id}) ...")
# 记录 worker.log 行数
try:
log_before = int(subprocess.check_output(
["wc", "-l", "/home/ubuntu/hair/worker.log"], text=True, timeout=5).split()[0])
except Exception:
log_before = 0
# VRAM before
vram_before, vram_free_before = get_vram()
procs_before = get_gpu_procs()
t_start = time.time()
# 调用API
api_result, elapsed, error = call_api(img_path, hl_id)
t_end = time.time()
# VRAM after
vram_after, vram_free_after = get_vram()
procs_after = get_gpu_procs()
# 读取新日志
try:
log_diff = subprocess.check_output(
["tail", "-n", "+{}".format(log_before + 1), "/home/ubuntu/hair/worker.log"],
text=True, timeout=5)
except Exception:
log_diff = ""
step_timing = parse_step_timing(log_diff)
# 提取结果图片
preview_b64 = ""
grown_b64 = ""
api_code = api_result.get("code", -1)
api_msg = api_result.get("message", "")
api_results = api_result.get("data", {}).get("results", [])
if api_results:
r0 = api_results[0]
preview_b64 = r0.get("image_base64", "")
grown_b64 = r0.get("grown_image_base64", "")
# 保存缩略图
thumb_dir = os.path.join(OUT_DIR, "thumbs")
os.makedirs(thumb_dir, exist_ok=True)
if grown_b64:
grown_bytes = base64.b64decode(grown_b64)
thumb_path = os.path.join(thumb_dir, f"{img_name}_hl{hl_id}_grown.jpg")
with open(thumb_path, "wb") as f:
f.write(grown_bytes)
if preview_b64:
preview_bytes = base64.b64decode(preview_b64)
thumb_path = os.path.join(thumb_dir, f"{img_name}_hl{hl_id}_preview.png")
with open(thumb_path, "wb") as f:
f.write(preview_bytes)
record = {
"image": img_name,
"image_idx": img_idx + 1,
"hairline_id": hl_id,
"hairline_key": hl_key,
"hairline_label": hl_label,
"timestamp": datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"total_time_s": round(elapsed, 2),
"api_code": api_code,
"api_msg": api_msg,
"error": error,
"has_preview": bool(preview_b64),
"has_grown": bool(grown_b64),
"preview_size": len(preview_b64),
"grown_size": len(grown_b64),
"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_after - vram_before,
"procs_before": procs_before,
"procs_after": procs_after,
"step_timing": step_timing,
}
results.append(record)
print(f" → HTTP code={api_code} time={elapsed:.1f}s vram={vram_before}{vram_after}MB "
f"preview={'✓' if preview_b64 else '✗'} grown={'✓' if grown_b64 else '✗'}")
# 保存中间结果
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}")
print(f"📊 总调用: {len(results)}/{total_calls}")
def generate_html(results):
"""生成HTML报告"""
# 统计
total = len(results)
success = sum(1 for r in results if r["has_grown"])
failed = total - success
times = [r["total_time_s"] for r in results if r["has_grown"]]
avg_time = sum(times) / len(times) if times else 0
max_time = max(times) if times else 0
min_time = min(times) if times else 0
# 按发际线类型分组统计
hl_stats = {}
for r in results:
if r["has_grown"]:
hl = r["hairline_label"]
if hl not in hl_stats:
hl_stats[hl] = {"count": 0, "times": []}
hl_stats[hl]["count"] += 1
hl_stats[hl]["times"].append(r["total_time_s"])
# 按图片分组
img_groups = {}
for r in results:
img = r["image"]
if img not in img_groups:
img_groups[img] = []
img_groups[img].append(r)
# 生成VRAM变化数据
vram_data = [(i, r["vram_after_mb"]) for i, r in enumerate(results)]
html = f"""<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<title>接口2 基准测试报告</title>
<style>
* {{ margin:0; padding:0; box-sizing:border-box; }}
body {{ font-family:-apple-system,BlinkMacSystemFont,'Segoe UI',Roboto,sans-serif; background:#1a1a2e; color:#e0e0e0; padding:20px; }}
h1 {{ text-align:center; margin-bottom:20px; color:#00d4ff; }}
.subtitle {{ text-align:center; color:#888; margin-bottom:30px; font-size:14px; }}
.summary {{ display:grid; grid-template-columns:repeat(auto-fit,minmax(180px,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:32px; 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:12px 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; }}
.img-cell {{ text-align:left; }}
.time-bar {{ display:inline-block; height:20px; background:linear-gradient(90deg,#0f3460,#00d4ff); border-radius:4px; vertical-align:middle; min-width:2px; }}
.hl-badge {{ display:inline-block; padding:2px 8px; border-radius:10px; font-size:11px; font-weight:600; }}
.hl-ellipse {{ background:#1b4332; color:#52b788; }}
.hl-flower {{ background:#3a0ca3; color:#c77dff; }}
.hl-heart {{ background:#6a040f; color:#ff6b6b; }}
.hl-straight {{ background:#0077b6; color:#90e0ef; }}
.hl-wave {{ background:#9d4edd; color:#e0aaff; }}
.ok {{ color:#0f0; }} .fail {{ color:#f44; }}
.vram-chart {{ margin:20px 0; }}
.vram-bars {{ display:flex; align-items:flex-end; height:120px; gap:2px; padding:10px; background:#0d1117; border-radius:8px; }}
.vram-bar {{ flex:1; background:linear-gradient(180deg,#00d4ff,#0f3460); border-radius:2px 2px 0 0; min-height:2px; position:relative; }}
.vram-bar:hover::after {{ content:attr(data-val) 'MB'; position:absolute; bottom:100%; left:50%; transform:translateX(-50%); background:#333; padding:2px 6px; border-radius:4px; font-size:10px; white-space:nowrap; }}
.section-title {{ font-size:18px; font-weight:600; margin:30px 0 15px; color:#00d4ff; border-bottom:1px solid #333; padding-bottom:10px; }}
.hl-stats {{ display:grid; grid-template-columns:repeat(5,1fr); gap:15px; margin-bottom:20px; }}
.hl-card {{ background:#16213e; border-radius:12px; padding:15px; text-align:center; }}
.hl-card .avg {{ font-size:24px; font-weight:700; color:#00d4ff; }}
.hl-card .minmax {{ font-size:11px; color:#888; margin-top:5px; }}
.thumb {{ max-width:100px; max-height:100px; border-radius:4px; cursor:pointer; }}
.thumb:hover {{ transform:scale(2); transition:transform 0.3s; }}
</style>
</head>
<body>
<h1>📊 接口2 基准测试报告</h1>
<div class="subtitle">生成时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')} {total} 次调用</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">{failed}</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">{min_time:.1f}s</div><div class="label">最快</div></div>
<div class="card time"><div class="num">{max_time:.1f}s</div><div class="label">最慢</div></div>
</div>
"""
# 按发际线类型统计
if hl_stats:
html += '<div class="section-title">按发际线类型统计</div><div class="hl-stats">'
for hl_label in ["椭圆形","花瓣形","心形","直线形","波浪形"]:
if hl_label in hl_stats:
s = hl_stats[hl_label]
times = s["times"]
avg = sum(times) / len(times)
mn, mx = min(times), max(times)
cls = hl_label
html += f'<div class="hl-card"><div class="hl-badge hl-{{cls}}">{hl_label}</div><div class="avg">{avg:.1f}s</div><div class="minmax">{mn:.1f}~{mx:.1f}s ({s["count"]}次)</div></div>'
else:
html += f'<div class="hl-card"><div class="hl-badge">{hl_label}</div><div class="avg">-</div><div class="minmax">未完成</div></div>'
html += '</div>'
# VRAM 变化图
if vram_data:
max_vram = max(v for _, v in vram_data if v > 0) or 1
html += '<div class="section-title">显存变化</div><div class="vram-chart"><div class="vram-bars">'
for i, (_, vram) in enumerate(vram_data):
if vram > 0:
h = int(vram / max_vram * 100)
html += f'<div class="vram-bar" style="height:{h}%" data-val="{vram}" title="第{i+1}次"></div>'
else:
html += f'<div class="vram-bar" style="height:0%" data-val="0"></div>'
html += '</div></div>'
# 详细结果表格
html += '<div class="section-title">详细结果</div><table><thead><tr>'
html += '<th>#</th><th>图片</th><th>发际线</th><th>总耗时</th><th>遮罩步骤</th>'
html += '<th>显存前</th><th>显存后</th><th>显存变化</th>'
html += '<th>预览图</th><th>生发图</th><th>状态</th>'
html += '</tr></thead><tbody>'
for i, r in enumerate(results):
hl_cls = r["hairline_key"]
status = '<span class="ok">✓ 成功</span>' if r["has_grown"] else f'<span class="fail">✗ {r.get("error","")}</span>'
step_ms = r.get("step_timing", {}).get("step1_ms", "")
step_str = f"{step_ms}ms" if step_ms else "-"
vram_d = r["vram_delta_mb"]
vram_d_str = f'<span style="color:{"#f44" if vram_d>0 else "#0f0"}">{"+" if vram_d>=0 else ""}{vram_d}</span>'
# 缩略图
thumb_grown = ""
if r["has_grown"]:
thumb_path = f"thumbs/{r['image']}_hl{r['hairline_id']}_grown.jpg"
if os.path.exists(os.path.join(OUT_DIR, thumb_path)):
thumb_grown = f'<img class="thumb" src="{thumb_path}">'
html += f"""<tr>
<td>{i+1}</td>
<td class="img-cell">{r['image']}</td>
<td><span class="hl-badge hl-{hl_cls}">{r['hairline_label']}</span></td>
<td><div class="time-bar" style="width:{min(r['total_time_s'],300)}px"></div> {r['total_time_s']:.1f}s</td>
<td>{step_str}</td>
<td>{r['vram_before_mb']}MB</td>
<td>{r['vram_after_mb']}MB</td>
<td>{vram_d_str}</td>
<td>{'✓' if r['has_preview'] else '✗'}</td>
<td>{thumb_grown if thumb_grown else ('✓' if r['has_grown'] else '✗')}</td>
<td>{status}</td>
</tr>"""
html += '</tbody></table>'
html += '</body></html>'
with open(REPORT_HTML, "w", encoding="utf-8") as f:
f.write(html)
if __name__ == "__main__":
run_benchmark()