Files
change_hair/train_batch_stepC.py
T
2026-07-17 22:13:56 +08:00

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()