完善部署并训练5个新发型 + 换发型集成文档
部署修复: - torch.load 增加 weights_only=False patch,兼容 PyTorch 2.6+ 加载旧权重 - OSS 改为懒加载,本地用 output_format=base64 无需配凭证即可启动 - 补全被 gitignore 误排除的必需代码:core/models/layers/data、models/layers/data、keypoints/lib - webui 训练命令 --xformers 改 --sdpa(修复 xformers 无 CUDA 支持报错) 功能调整: - hair_grow_service 端口改 8899、preview 路由修复(send_file) - list_hairstyles 增加发型白名单,测试页只展示当前5个发型 新增脚本: - train_lora_parallel.py:直接调 kohya 并行训练 LoRA(绕过 photo_service 串行限制) - train_hairstyles_parallel.py / train_batch_stepC.py:批量训练辅助脚本 - scripts/sync_data_to_server.sh:大文件断点续传到云服务器 文档: - docs/换发型集成文档.md:换发型完整流程、服务架构、资源依赖、训练方法、集成步骤
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
#!/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/<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()
|
||||
Reference in New Issue
Block a user