Files
change_hair/gen_hairstyle_previews.py
xsl 9b7324cd26 chore: 适配本机 xsl 路径(RTX 5090 机器)
将所有绝对路径从 /home/ubuntu 改回 /home/xsl,覆盖启动脚本、
configure.ini 及训练/测试脚本。纯路径替换,无功能性改动。

本分支(xsl5090)用于本机(RTX 5090)运行;master 维持 /home/ubuntu
供远程 ubuntu 机器使用。
2026-07-18 16:26:16 +08:00

123 lines
4.2 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""批量生成发型预览图
用 boy.png 给男款发型、girl.png 给女款发型跑换发型服务,
生成每个发型套在标准脸型上的效果图,存为缩略图供前端展示。
用法: python gen_hairstyle_previews.py [--limit N] (limit用于测试)
"""
import os
import sys
import json
import time
import base64
import requests
HAIR_SERVICE_DIR = "/home/xsl/change_hair/project/hair_service_sd"
sys.path.insert(0, HAIR_SERVICE_DIR)
os.chdir(HAIR_SERVICE_DIR)
BOY_IMG = "/home/xsl/change_hair/images/boy.png"
GIRL_IMG = "/home/xsl/change_hair/images/girl.png"
HAIRSTYLE_DIR = "/home/xsl/change_hair/project/data/ref_hairstyle"
TRAIN_DIR = "/home/xsl/change_hair/data/train_material"
PREVIEW_DIR = "/home/xsl/change_hair/hair_grow_service/static/previews" # 预览图存储目录
SWAP_API = "http://127.0.0.1:8801/api/swapHair/v1"
def load_hairstyles():
"""返回 {hair_id: gender}"""
styles = {}
for hid in os.listdir(HAIRSTYLE_DIR):
cfg = os.path.join(HAIRSTYLE_DIR, hid, "config.json")
lora = os.path.join(TRAIN_DIR, hid, "model", "hairstyle_hd_lora.safetensors")
if not (os.path.exists(cfg) and os.path.exists(lora)):
continue
g = json.load(open(cfg)).get("gender", "?")
styles[hid] = g
return styles
def gen_one(hair_id, gender, img_path):
"""给一个发型生成预览图,返回是否成功"""
out_path = os.path.join(PREVIEW_DIR, f"{hair_id}.jpg")
if os.path.exists(out_path):
return True, "已存在跳过"
with open(img_path, "rb") as f:
img_b64 = "data:image/png;base64," + base64.b64encode(f.read()).decode()
payload = {
"hair_id": hair_id,
"task_id": f"preview_{hair_id}",
"is_hr": "false", # 非高清,快一点
"user_img_path": img_b64,
"output_format": "base64"
}
try:
r = requests.post(SWAP_API, json=payload, timeout=180)
d = r.json()
if d.get("state") == 0 and d.get("data"):
import cv2
import numpy as np
img = cv2.imdecode(np.frombuffer(base64.b64decode(d["data"]), np.uint8), cv2.IMREAD_COLOR)
# 缩小到 300x400 当缩略图
if img is not None:
h, w = img.shape[:2]
scale = min(300 / w, 400 / h)
img = cv2.resize(img, (int(w * scale), int(h * scale)), interpolation=cv2.INTER_AREA)
cv2.imwrite(out_path, img, [cv2.IMWRITE_JPEG_QUALITY, 85])
return True, "生成成功"
return False, f"state={d.get('state')} msg={d.get('msg','')}"
except Exception as e:
return False, str(e)
def main():
import cv2
limit = None
if "--limit" in sys.argv:
limit = int(sys.argv[sys.argv.index("--limit") + 1])
os.makedirs(PREVIEW_DIR, exist_ok=True)
styles = load_hairstyles()
print(f"共 {len(styles)} 个发型需要生成预览图")
# 按性别选基准图
boy_styles = [(h, g) for h, g in styles.items() if g == "boy"]
girl_styles = [(h, g) for h, g in styles.items() if g == "girl"]
print(f" boy: {len(boy_styles)} (用 boy.png)")
print(f" girl: {len(girl_styles)} (用 girl.png)")
tasks = boy_styles + girl_styles
if limit:
tasks = tasks[:limit]
print(f" [测试模式] 只处理前 {limit} 个")
ok, fail, skip = 0, 0, 0
t0 = time.time()
for i, (hair_id, gender) in enumerate(tasks):
img_path = BOY_IMG if gender == "boy" else GIRL_IMG
out_path = os.path.join(PREVIEW_DIR, f"{hair_id}.jpg")
if os.path.exists(out_path):
skip += 1
continue
elapsed = time.time() - t0
eta = elapsed / max(ok + fail, 1) * (len(tasks) - i) if (ok + fail) > 0 else 0
print(f"[{i+1}/{len(tasks)}] {hair_id}({gender})...", end=" ", flush=True)
success, msg = gen_one(hair_id, gender, img_path)
if success:
ok += 1
print(f"✓ ({msg})")
else:
fail += 1
print(f"✗ ({msg})")
print(f"\n=== 完成: 成功{ok} 失败{fail} 跳过{skip} 总耗时{time.time()-t0:.0f}s ===")
print(f"预览图目录: {PREVIEW_DIR}")
if __name__ == "__main__":
main()