97 lines
3.7 KiB
Python
97 lines
3.7 KiB
Python
#!/usr/bin/env python3
|
|
# -*- coding: utf-8 -*-
|
|
"""阶段C:为已训练的发型生成 ref材质 + 模板材质 + 预览图
|
|
|
|
前置:LoRA 已训练完成,webui/photo_service/hair_service_sd 已启动。
|
|
|
|
对每个发型依次:
|
|
step3: 调 hair_service_sd trainCallBack 回调生成 ref 材质
|
|
step4: 生成模板材质
|
|
step5: 生成预览图
|
|
"""
|
|
import os
|
|
import sys
|
|
import time
|
|
import json
|
|
import shutil
|
|
|
|
HAIR_SERVICE_DIR = "/home/ubuntu/change_hair/project/hair_service_sd"
|
|
os.chdir(HAIR_SERVICE_DIR)
|
|
sys.path.insert(0, HAIR_SERVICE_DIR)
|
|
os.environ.setdefault("HF_HUB_OFFLINE", "1")
|
|
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
|
|
os.environ.setdefault("CRYPTOGRAPHY_OPENSSL_NO_LEGACY", "1")
|
|
|
|
import requests as req
|
|
|
|
# (hair_id, gender, template_img)
|
|
HAIRSTYLES = [
|
|
("chang_tuoyuan", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_tuoyuan/chang_tuoyuan.jpg"),
|
|
("chang_bolang", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_bolang/chang_bolang.jpg"),
|
|
("chang_zhixian", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_zhixian/chang_zhixian.jpg"),
|
|
("chang_huaban", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_huaban/chang_huaban.jpg"),
|
|
("chang_xinxing", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_xinxing/chang_xinxing.jpg"),
|
|
]
|
|
|
|
HAIR_CALLBACK = "http://127.0.0.1:8801/api/hair/trainCallBack"
|
|
UPLOAD_TRAIN_DIR = "/home/ubuntu/change_hair/project/data/upload_train_imgs"
|
|
|
|
# 导入 step4/step5 函数(会触发模型加载)
|
|
from train_hairstyle_full import step4_template_material, step5_preview
|
|
from common.logger import config
|
|
|
|
|
|
def log(msg):
|
|
print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True)
|
|
|
|
|
|
def step3_callback(hair_id, template_img):
|
|
"""调 trainCallBack 生成 ref 材质。需要先把模板图放到 upload_train_imgs"""
|
|
log(f" --- step3: {hair_id} 生成 ref 材质 ---")
|
|
# 准备 upload_train_imgs/<hair_id>/first##<hair_id>.jpg
|
|
src_dir = os.path.join(UPLOAD_TRAIN_DIR, hair_id)
|
|
os.makedirs(src_dir, exist_ok=True)
|
|
dst = os.path.join(src_dir, f"first##{hair_id}.jpg")
|
|
shutil.copy(template_img, dst)
|
|
|
|
try:
|
|
r = req.post(HAIR_CALLBACK, json={
|
|
"task_id": f"train_{hair_id}", "hair_id": hair_id,
|
|
"state": 0, "msg": "头发lora训练成功", "is_tj": "1"
|
|
}, timeout=300)
|
|
d = r.json()
|
|
ref_dir = os.path.join(config.get('default', 'hairstyleDir'), hair_id)
|
|
ok = d.get("state") == 0 and os.path.exists(os.path.join(ref_dir, "ref_rgb_8uc3_768.png"))
|
|
log(f" {'✓' if ok else '✗'} {hair_id} ref材质 {'成功' if ok else '失败: ' + str(d.get('msg'))}")
|
|
return ok
|
|
except Exception as e:
|
|
log(f" ✗ {hair_id} ref材质异常: {e}")
|
|
return False
|
|
|
|
|
|
def main():
|
|
t0 = time.time()
|
|
log(f"阶段C: 为 {len(HAIRSTYLES)} 个发型生成材质和预览")
|
|
|
|
results = {}
|
|
for hair_id, gender, template_img in HAIRSTYLES:
|
|
log(f"\n{'='*50}")
|
|
log(f"处理: {hair_id}")
|
|
log(f"{'='*50}")
|
|
ok3 = step3_callback(hair_id, template_img)
|
|
ok4 = step4_template_material(hair_id, template_img)
|
|
ok5 = step5_preview(hair_id, gender)
|
|
results[hair_id] = (ok3, ok4, ok5)
|
|
log(f" {hair_id}: ref={'✓' if ok3 else '✗'} 模板={'✓' if ok4 else '✗'} 预览={'✓' if ok5 else '✗'}")
|
|
|
|
log(f"\n{'='*60}")
|
|
log(f"阶段C完成,耗时 {time.time()-t0:.0f}s")
|
|
log(f"{'='*60}")
|
|
for hair_id, (ok3, ok4, ok5) in results.items():
|
|
status = "✓全部成功" if (ok3 and ok4 and ok5) else f"⚠️ ref={'✓' if ok3 else '✗'} 模板={'✓' if ok4 else '✗'} 预览={'✓' if ok5 else '✗'}"
|
|
log(f" {hair_id}: {status}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|