硬件迁移:从 RTX 5090 (32GB) 迁移到 RTX 3090 (24GB) 主要改动: - 所有启动脚本和配置文件中的 /home/xsl/ 路径替换为 /home/ubuntu/ - 适配新的 24GB VRAM 环境
151 lines
6.2 KiB
Python
151 lines
6.2 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
"""批量并行训练多个发型
|
||
|
||
流程:
|
||
阶段A: 串行 step1(准备训练数据,共用GPU模型,避免重复加载)
|
||
阶段B: 并行 step2(LoRA训练,kohya独立进程,限制并发数)
|
||
阶段C: 启动服务后串行 step3(回调生成ref材质) + step4(模板) + step5(预览)
|
||
|
||
用法:
|
||
cd /home/ubuntu/change_hair/project/hair_service_sd
|
||
python /home/ubuntu/change_hair/train_hairstyles_parallel.py
|
||
"""
|
||
import os
|
||
import sys
|
||
import time
|
||
import json
|
||
import subprocess
|
||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||
|
||
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")
|
||
os.environ.setdefault("CUDA_VISIBLE_DEVICES", "0")
|
||
|
||
# ====== 配置:本次要训练的发型 ======
|
||
# (hair_id, input_dir, gender, template_img)
|
||
HAIRSTYLES = [
|
||
("chang_tuoyuan", "/home/ubuntu/change_hair/hair_type_images/chang_tuoyuan", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_tuoyuan/chang_tuoyuan.jpg"),
|
||
("chang_bolang", "/home/ubuntu/change_hair/hair_type_images/chang_bolang", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_bolang/chang_bolang.jpg"),
|
||
("chang_zhixian", "/home/ubuntu/change_hair/hair_type_images/chang_zhixian", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_zhixian/chang_zhixian.jpg"),
|
||
("chang_huaban", "/home/ubuntu/change_hair/hair_type_images/chang_huaban", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_huaban/chang_huaban.jpg"),
|
||
("chang_xinxing", "/home/ubuntu/change_hair/hair_type_images/chang_xinxing", "girl", "/home/ubuntu/change_hair/hair_type_images/chang_xinxing/chang_xinxing.jpg"),
|
||
]
|
||
PARALLEL = 2 # LoRA 训练并发数
|
||
|
||
# 导入训练模块(复用其 step 函数)
|
||
from train_hairstyle_full import (
|
||
step1_prepare_data, step2_train, step3_wait_and_callback,
|
||
step4_template_material, step5_preview, get_models
|
||
)
|
||
import train_hairstyle_full as TH
|
||
from common.logger import config
|
||
|
||
|
||
def log(msg):
|
||
print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True)
|
||
|
||
|
||
def main():
|
||
t0 = time.time()
|
||
ids = [h[0] for h in HAIRSTYLES]
|
||
log(f"批量训练 {len(ids)} 个发型: {ids} (并发={PARALLEL})")
|
||
|
||
# ============ 阶段A: 串行 step1(准备数据)============
|
||
log("=" * 60)
|
||
log("阶段A: 串行准备训练数据 (step1)")
|
||
log("=" * 60)
|
||
log(" 预加载模型...")
|
||
get_models() # 预加载,5个发型共用
|
||
step1_ok, step1_fail = [], []
|
||
for hair_id, input_dir, gender, _ in HAIRSTYLES:
|
||
# 清理可能的历史训练数据,确保干净
|
||
train_dir = config.get('default', 'train_dir')
|
||
out_dir = os.path.join(train_dir, hair_id, "images", "1_hairstyle")
|
||
os.makedirs(out_dir, exist_ok=True)
|
||
os.makedirs(os.path.join(train_dir, hair_id, "model"), exist_ok=True)
|
||
|
||
log(f" --- step1: {hair_id} ---")
|
||
try:
|
||
ok = step1_prepare_data(hair_id, input_dir, gender)
|
||
if ok:
|
||
step1_ok.append(hair_id)
|
||
log(f" ✓ {hair_id} 数据准备完成")
|
||
else:
|
||
step1_fail.append(hair_id)
|
||
log(f" ✗ {hair_id} 数据准备失败")
|
||
except Exception as e:
|
||
step1_fail.append(hair_id)
|
||
log(f" ✗ {hair_id} 数据准备异常: {e}")
|
||
log(f"阶段A完成: 成功 {len(step1_ok)}/{len(ids)},失败 {step1_fail}")
|
||
if not step1_ok:
|
||
log("全部 step1 失败,退出"); return
|
||
|
||
# 释放 step1 用的 GPU 模型,给 step2 训练腾显存
|
||
log("释放 step1 模型显存...")
|
||
import torch
|
||
try:
|
||
for m in [TH._detector, TH._aligner, TH._matte]:
|
||
if m is not None and hasattr(m, 'model'):
|
||
pass
|
||
torch.cuda.empty_cache()
|
||
except Exception:
|
||
pass
|
||
|
||
# ============ 阶段B: 并行 step2(LoRA 训练)============
|
||
log("=" * 60)
|
||
log(f"阶段B: 并行 LoRA 训练 (step2, 并发={PARALLEL})")
|
||
log("=" * 60)
|
||
step2_results = {} # hair_id -> bool
|
||
with ThreadPoolExecutor(max_workers=PARALLEL) as ex:
|
||
futures = {ex.submit(step2_train, hid): hid for hid in step1_ok}
|
||
for fut in as_completed(futures):
|
||
hid = futures[fut]
|
||
try:
|
||
ok = fut.result()
|
||
step2_results[hid] = ok
|
||
log(f" {'✓' if ok else '✗'} {hid} step2 (训练启动)={'成功' if ok else '失败'}")
|
||
except Exception as e:
|
||
step2_results[hid] = False
|
||
log(f" ✗ {hid} step2 异常: {e}")
|
||
|
||
# 等待所有 LoRA 训练真正完成(step2 只是启动,训练在 photo_service 后台跑)
|
||
log("等待所有 LoRA 训练完成(检查权重文件)...")
|
||
train_dir = config.get('default', 'train_dir')
|
||
pending = [hid for hid in step1_ok if step2_results.get(hid)]
|
||
completed = set()
|
||
deadline = time.time() + 3600 # 最多等1小时
|
||
while pending and time.time() < deadline:
|
||
still = []
|
||
for hid in pending:
|
||
lora = os.path.join(train_dir, hid, "model", "hairstyle_hd_lora.safetensors")
|
||
if os.path.exists(lora) and os.path.getsize(lora) > 100000000:
|
||
completed.add(hid)
|
||
log(f" ✓ {hid} LoRA 训练完成 ({os.path.getsize(lora)//1024//1024}MB)")
|
||
else:
|
||
still.append(hid)
|
||
pending = still
|
||
if pending:
|
||
log(f" ...等待中: {pending}")
|
||
time.sleep(30)
|
||
if pending:
|
||
log(f"⚠️ 超时未完成: {pending}")
|
||
log(f"阶段B完成: LoRA 训练完成 {len(completed)}/{len(step1_ok)}")
|
||
|
||
# 写一个完成清单文件,供阶段C脚本读取
|
||
done_file = "/home/ubuntu/change_hair/project/logs/train_batch_done.json"
|
||
with open(done_file, "w") as f:
|
||
json.dump({"completed": sorted(completed), "all": ids}, f)
|
||
log(f"已写入完成清单: {done_file}")
|
||
log(f"\n批量训练阶段A+B总耗时 {time.time()-t0:.0f}s")
|
||
log(f"完成训练的发型: {sorted(completed)}")
|
||
log("接下来需启动服务跑阶段C(step3/4/5 生成材质和预览)")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|