初始化:换发型/换发色/训练发型服务
包含: - hair_service_sd: 主服务(换发型/换发色/生发,端口8801) - photo_service: LoRA调度+训练(端口32678) - hair_grow_service: 调试测试页(端口8888,含4个测试页) - 批量训练脚本(batch_train_hairstyles.py) - 发际线mask自动识别(hairline_mask.py,4种方案) - 手绘mask换发型(hair_swap_manual.py) - 文档:README.md + LARGE_FILES.md + docs/ 大文件(模型权重200G、训练数据123G)已排除,见 LARGE_FILES.md OSS/COS密钥已脱敏为环境变量,原文件备份在本地
This commit is contained in:
@@ -0,0 +1,232 @@
|
||||
#!/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 <hair_id>")
|
||||
log(f"详细日志: {LOG_FILE}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user