#!/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/xsl/change_hair/project/hair_service_sd python /home/xsl/change_hair/batch_train_hairstyles.py --src /home/xsl/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/xsl/change_hair/hair_type_images" WORK_DIR = "/home/xsl/change_hair/data/batch_train_inputs" LOG_FILE = "/home/xsl/change_hair/data/batch_train_log.txt" TRAIN_SCRIPT = "/home/xsl/change_hair/train_hairstyle_full.py" HAIR_SERVICE_DIR = "/home/xsl/change_hair/project/hair_service_sd" LOCK_FILE = "/home/xsl/change_hair/data/batch_train.pid" DONE_FILE = "/home/xsl/change_hair/data/batch_train_done.txt" # 已完成发型清单(断点续跑) LORA_DIR = "/home/xsl/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/xsl/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/xsl/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 ") log(f"详细日志: {LOG_FILE}") if __name__ == "__main__": main()