#!/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/xsl/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/xsl/change_hair/hair_type_images/chang_tuoyuan/chang_tuoyuan.jpg"), ("chang_bolang", "girl", "/home/xsl/change_hair/hair_type_images/chang_bolang/chang_bolang.jpg"), ("chang_zhixian", "girl", "/home/xsl/change_hair/hair_type_images/chang_zhixian/chang_zhixian.jpg"), ("chang_huaban", "girl", "/home/xsl/change_hair/hair_type_images/chang_huaban/chang_huaban.jpg"), ("chang_xinxing", "girl", "/home/xsl/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/xsl/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//first##.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()