硬件迁移:从 RTX 5090 (32GB) 迁移到 RTX 3090 (24GB) 主要改动: - 所有启动脚本和配置文件中的 /home/xsl/ 路径替换为 /home/ubuntu/ - 适配新的 24GB VRAM 环境
233 lines
8.7 KiB
Python
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()
|