#!/usr/bin/env python3 """接口2 female wave发型 分辨率对比测试 v3: 21图 × 5档 = 105次。 串行执行,记录耗时与成败,结果图按 图_档位 命名保存。支持断点续跑。 """ import base64, csv, json, os, time, io import requests from PIL import Image API = "http://127.0.0.1:8187/api/v1/hair/grow" TOKEN = "dev-shared-secret-2026" TIMEOUT = 600 OUT = "/home/ubuntu/hair/image/wave_test/out" PROGRESS = "/home/ubuntu/hair/image/wave_test/progress.json" IMG_DIR = "/home/ubuntu/hair/image" # 21张图:19张girl_img + asdf + qwer IMAGES = [] for f in sorted(os.listdir(os.path.join(IMG_DIR, "girl_img"))): if f.lower().endswith((".jpg", ".jpeg", ".png")): IMAGES.append((os.path.splitext(f)[0], os.path.join(IMG_DIR, "girl_img", f))) IMAGES.append(("asdf", os.path.join(IMG_DIR, "asdf.jpg"))) IMAGES.append(("qwer", os.path.join(IMG_DIR, "qwer.jpg"))) # wave = female hair_style 5 STYLE_IDX = 5 # 5档: 原图(0) / 1024 / 896 / 768 / 640 SIDES = [ ("origin", 0), ("s1024", 1024), ("s896", 896), ("s768", 768), ("s640", 640), ] def load_progress(): if os.path.exists(PROGRESS): try: return json.load(open(PROGRESS)) except Exception: pass return {"done": [], "results": []} def save_progress(prog): json.dump(prog, open(PROGRESS, "w"), ensure_ascii=False, indent=1) def run_one(img_name, img_path, side_name, side_val): key = f"{img_name}_{side_name}" with open(img_path, "rb") as f: img_b64 = base64.b64encode(f.read()).decode() data = { "image_base64": "data:image/jpeg;base64," + img_b64, "gender": "female", "hair_style": str(STYLE_IDX), "redraw_max_side": str(side_val), } t0 = time.time() rec = {"key": key, "img": img_name, "side": side_name, "side_val": side_val, "ok": False, "elapsed": 0.0, "err": "", "out_w": 0, "out_h": 0, "bytes": 0} try: r = requests.post(API, data=data, headers={"X-Internal-Token": TOKEN}, timeout=TIMEOUT) rec["elapsed"] = round(time.time() - t0, 1) d = r.json() if d.get("code") != 0: rec["err"] = f"code={d.get('code')} {d.get('message','')}"[:200] return rec results = (d.get("data") or {}).get("results") or [] if not results: rec["err"] = "空结果" return rec it = results[0] grown = it.get("grown_image_base64") if not grown: rec["err"] = "无生发图(重绘失败/OOM?)" return rec raw = base64.b64decode(grown) im = Image.open(io.BytesIO(raw)) rec["out_w"], rec["out_h"] = im.size out_path = f"{OUT}/{key}.jpg" with open(out_path, "wb") as fo: fo.write(raw) rec["ok"] = True rec["bytes"] = len(raw) except requests.exceptions.Timeout: rec["elapsed"] = round(time.time() - t0, 1) rec["err"] = f"超时(>{TIMEOUT}s)" except Exception as e: rec["elapsed"] = round(time.time() - t0, 1) rec["err"] = f"{type(e).__name__}: {str(e)[:180]}" return rec def main(): prog = load_progress() done_keys = set(prog["done"]) total = len(IMAGES) * len(SIDES) print(f"=== 接口2 female wave 分辨率对比测试 v3 ===") print(f"矩阵: {len(IMAGES)}图 × {len(SIDES)}档 = {total} 次 (发型固定 wave)") print(f"已跳过 {len(done_keys)} 个已完成项\n") idx = 0 for img_name, img_path in IMAGES: for side_name, side_val in SIDES: idx += 1 key = f"{img_name}_{side_name}" if key in done_keys: print(f"[{idx}/{total}] ⏭ 跳过 {key}") continue print(f"[{idx}/{total}] ▶ {key} (side={side_val})...", end=" ", flush=True) rec = run_one(img_name, img_path, side_name, side_val) prog["results"].append(rec) prog["done"].append(key) save_progress(prog) if rec["ok"]: print(f"✅ {rec['elapsed']}s {rec['out_w']}x{rec['out_h']} ({rec['bytes']}B)") else: print(f"❌ {rec['elapsed']}s {rec['err']}") time.sleep(2) print("\n" + "=" * 70) print("汇总") print("=" * 70) write_report(prog["results"]) print(f"\n结果图: {OUT}/") print(f"CSV: {OUT}/report.csv JSON: {OUT}/report.json") def write_report(results): with open(f"{OUT}/report.csv", "w", newline="") as f: w = csv.writer(f) w.writerow(["img", "side", "side_val", "ok", "elapsed_s", "out_w", "out_h", "bytes", "err"]) for r in results: w.writerow([r["img"], r["side"], r["side_val"], r["ok"], r["elapsed"], r["out_w"], r["out_h"], r["bytes"], r["err"]]) json.dump(results, open(f"{OUT}/report.json", "w"), ensure_ascii=False, indent=1) # 各档平均耗时 print("\n--- 各档平均耗时(秒)---") for side_name, side_val in SIDES: ts = [r["elapsed"] for r in results if r["side"] == side_name and r["ok"]] if ts: avg = sum(ts) / len(ts) print(f" {side_name:8s} (={side_val:>4}): 平均 {avg:5.1f}s [{min(ts):.1f}~{max(ts):.1f}] 成功 {len(ts)}") if __name__ == "__main__": main()