Files
change_hair/batch_train_hairstyles.py
T
xsl 6f34e8876c feat: 适配 RTX 3090 (24GB) 环境优化
硬件迁移:从 RTX 5090 (32GB) 迁移到 RTX 3090 (24GB)

主要改动:
- 所有启动脚本和配置文件中的 /home/xsl/ 路径替换为 /home/ubuntu/
- 适配新的 24GB VRAM 环境
2026-07-18 19:00:15 +08:00

233 lines
8.7 KiB
Python

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""批量训练发型 LoRA
对 hair_type_images/ 下每张图(一种发型):
1. 拷贝到独立输入目录(train_hairstyle_full.py 要求 --input 是目录)
2. 调用 train_hairstyle_full.py 完整流程:数据增强→训练→材质→女生预览
3. 串行执行,记录每个发型成功/失败
用法:
cd /home/ubuntu/change_hair/project/hair_service_sd
python /home/ubuntu/change_hair/batch_train_hairstyles.py --src /home/ubuntu/change_hair/hair_type_images --gender girl [--only 圆-心形] [--start-from 心形-心形]
注意:
- 中文 hair_id 直接用文件名(去扩展名)作为 ID
- 每个发型约 30-40 分钟,30 个串行约 15-20 小时
- 失败的发型记录到失败清单,最后可单独重跑
"""
import os
import sys
import time
import shutil
import argparse
import subprocess
from datetime import datetime
SRC_DEFAULT = "/home/ubuntu/change_hair/hair_type_images"
WORK_DIR = "/home/ubuntu/change_hair/data/batch_train_inputs"
LOG_FILE = "/home/ubuntu/change_hair/data/batch_train_log.txt"
TRAIN_SCRIPT = "/home/ubuntu/change_hair/train_hairstyle_full.py"
HAIR_SERVICE_DIR = "/home/ubuntu/change_hair/project/hair_service_sd"
LOCK_FILE = "/home/ubuntu/change_hair/data/batch_train.pid"
DONE_FILE = "/home/ubuntu/change_hair/data/batch_train_done.txt" # 已完成发型清单(断点续跑)
LORA_DIR = "/home/ubuntu/change_hair/data/train_material" # LoRA 输出根目录
def acquire_lock():
"""PID 锁,防止重复启动多个实例互相干扰"""
import psutil
if os.path.exists(LOCK_FILE):
try:
old_pid = int(open(LOCK_FILE).read().strip())
if psutil.pid_exists(old_pid):
# 检查是不是真的本脚本进程
try:
proc = psutil.Process(old_pid)
if "batch_train_hairstyles" in " ".join(proc.cmdline()):
print(f"✗ 已有批量训练实例在运行 (PID={old_pid}),请先停止它再启动。")
print(f" 锁文件: {LOCK_FILE}")
sys.exit(2)
except Exception:
pass
except (ValueError, OSError):
pass
os.makedirs(os.path.dirname(LOCK_FILE), exist_ok=True)
open(LOCK_FILE, "w").write(str(os.getpid()))
def release_lock():
try:
os.remove(LOCK_FILE)
except OSError:
pass
def load_done():
"""读取已完成的发型集合"""
if not os.path.exists(DONE_FILE):
return set()
return set(l.strip() for l in open(DONE_FILE, encoding="utf-8") if l.strip())
def mark_done(hair_id):
with open(DONE_FILE, "a", encoding="utf-8") as f:
f.write(hair_id + "\n")
def log(msg):
line = f"[{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}] {msg}"
print(line, flush=True)
with open(LOG_FILE, "a", encoding="utf-8") as f:
f.write(line + "\n")
def prepare_input_dir(hair_id, src_img, work_dir):
"""为单个发型创建独立输入目录,里面只放这一张图"""
d = os.path.join(work_dir, hair_id)
if os.path.exists(d):
shutil.rmtree(d)
os.makedirs(d, exist_ok=True)
# 统一用 png 拷入,避免扩展名问题
dst = os.path.join(d, hair_id + ".png")
shutil.copy(src_img, dst)
return d, dst
def run_one(hair_id, src_img, gender, py, done_set, force=False):
"""跑单个发型的完整训练流程,返回 (ok, elapsed, skipped)"""
# 断点续跑:LoRA 已存在视为完成
lora_path = os.path.join(LORA_DIR, hair_id, "model", "hairstyle_hd_lora.safetensors")
if not force and (hair_id in done_set or
(os.path.exists(lora_path) and os.path.getsize(lora_path) > 100000000)):
log(f" ⊙ {hair_id} 已完成(LoRA 存在),跳过")
return True, 0, True
t0 = time.time()
log(f"==== 开始训练: {hair_id} ====")
# 步骤1:准备输入目录
try:
in_dir, tpl_img = prepare_input_dir(hair_id, src_img, WORK_DIR)
log(f" 输入目录就绪: {in_dir}")
except Exception as e:
log(f" ✗ 输入目录准备失败: {e}")
return False, time.time() - t0, False
# 步骤2-5:调用 train_hairstyle_full.py
cmd = [
py, TRAIN_SCRIPT,
"--hair-id", hair_id,
"--input", in_dir,
"--gender", gender,
"--template-img", tpl_img,
]
log(f" 调用: {' '.join(cmd[:2])} ... --hair-id {hair_id}")
try:
# 子进程输出实时写到日志文件
proc_log = os.path.join("/home/ubuntu/change_hair/data", f"subprocess_{hair_id}.log")
with open(proc_log, "w", encoding="utf-8") as f:
r = subprocess.run(
cmd, cwd=HAIR_SERVICE_DIR,
stdout=f, stderr=subprocess.STDOUT,
timeout=5400, # 单发型最长 90 分钟
)
elapsed = time.time() - t0
if r.returncode == 0:
log(f" ✓ 训练成功: {hair_id} (耗时 {elapsed/60:.1f} 分钟)")
mark_done(hair_id)
return True, elapsed, False
else:
log(f" ✗ 训练失败(returncode={r.returncode}): {hair_id}, 详见 {proc_log}")
return False, elapsed, False
except subprocess.TimeoutExpired:
log(f" ✗ 训练超时(>90分钟): {hair_id}")
return False, time.time() - t0, False
except Exception as e:
log(f" ✗ 异常: {hair_id}: {e}")
return False, time.time() - t0, False
def main():
ap = argparse.ArgumentParser(description="批量训练发型")
ap.add_argument("--src", default=SRC_DEFAULT, help="原图目录")
ap.add_argument("--gender", default="girl", choices=["boy", "girl"])
ap.add_argument("--py", default="/home/ubuntu/miniconda3/envs/my_hair/bin/python",
help="python 解释器")
ap.add_argument("--only", default=None, help="只训练指定 hair_id(调试用)")
ap.add_argument("--start-from", default=None,
help="从指定 hair_id 开始(含),跳过之前的(断点续跑)")
ap.add_argument("--force", action="store_true",
help="强制重训(忽略已完成清单和已存在的 LoRA)")
args = ap.parse_args()
acquire_lock()
done_set = load_done() if not args.force else set()
# 收集发型清单(按文件名排序)
files = []
for f in sorted(os.listdir(args.src)):
if f.lower().endswith((".jpg", ".jpeg", ".png")):
files.append(f)
if not files:
log(f"✗ 源目录无图片: {args.src}")
release_lock(); sys.exit(1)
# 构造 (hair_id, src_img)
tasks = []
for f in files:
hair_id = os.path.splitext(f)[0]
tasks.append((hair_id, os.path.join(args.src, f)))
# 过滤
if args.only:
tasks = [t for t in tasks if t[0] == args.only]
if not tasks:
log(f"✗ --only {args.only} 未匹配到发型"); release_lock(); sys.exit(1)
if args.start_from:
idx = next((i for i, t in enumerate(tasks) if t[0] == args.start_from), 0)
tasks = tasks[idx:]
log(f"\n{'#'*60}")
log(f"# 批量训练启动 (PID={os.getpid()})")
log(f"# 共 {len(tasks)} 个发型 | gender={args.gender} | force={args.force}")
if done_set:
log(f"# 已完成 {len(done_set)} 个(将跳过)")
log(f"# 预计总耗时约 {len(tasks)*35} 分钟")
log(f"{'#'*60}\n")
os.makedirs(WORK_DIR, exist_ok=True)
os.makedirs(os.path.dirname(LOG_FILE), exist_ok=True)
results = [] # (hair_id, ok, elapsed, skipped)
t_total = time.time()
try:
for i, (hair_id, src) in enumerate(tasks, 1):
log(f"\n[进度 {i}/{len(tasks)}] === {hair_id} ===")
ok, el, skipped = run_one(hair_id, src, args.gender, args.py, done_set, args.force)
results.append((hair_id, ok, el, skipped))
finally:
release_lock()
# 汇总
log(f"\n{'#'*60}")
log(f"# 批量训练完成汇总")
log(f"{'#'*60}")
succ = [r for r in results if r[1]]
fail = [r for r in results if not r[1]]
skip = [r for r in results if r[3]]
total_time = time.time() - t_total
for r in results:
hid, ok, el = r[0], r[1], r[2]
flag = "SKIP" if r[3] else ("OK" if ok else "FAIL")
log(f" [{flag}] {hid:20s} ({el/60:.1f} 分钟)")
log(f"\n成功 {len(succ)}/{len(results)}(含跳过 {len(skip)}),失败 {len(fail)},总耗时 {total_time/3600:.1f} 小时")
if fail:
log(f"失败清单(可重跑): {' '.join(f[0] for f in fail)}")
log(f"重跑命令: python {__file__} --gender {args.gender} --only <hair_id>")
log(f"详细日志: {LOG_FILE}")
if __name__ == "__main__":
main()