Files

2225 lines
126 KiB
Python
Raw Permalink Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""
脸部 LoRA 一站式控制台(Gradio Web UI
=======================================
所有功能集中在一个浏览器界面(localhost:7860):
Tab 1 素材自动处理 : auto 流水线(分类/裁剪/打标草稿/超量筛选/淘汰)
Tab 2 素材审核 : check 红黄绿判定报告
Tab 3 智能选图 : pick 多样化推荐
Tab 4 打标 : 照片墙 + 类型模板 + 发型 + 保存 + 整理训练集
Tab 5 融合工具 : merge_fixed(原图 + ComfyUI 修改图 无缝合并)
Tab 6 训练 : 启动训练 + 样本图浏览
用法:python caption_gui.py (浏览器打开 http://127.0.0.1:7860
"""
import contextlib
import glob
import io
import json
import tempfile
import time
import os
import subprocess
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
import face_checker as fc
import merge_fixed as mf
import gradio as gr
TOOLS_DIR = Path(__file__).resolve().parent
PROJ_DIR = TOOLS_DIR.parent
# 版本号:每次功能修改后递增,方便确认 GUI 是否已重启生效(界面+终端都会显示)
VERSION = "0.18.45"
# pythonw(无控制台)兼容:stdout/stderr 为 None 时重定向到 devnull,避免 print 崩溃
if sys.stdout is None:
sys.stdout = open(os.devnull, "w", encoding="utf-8")
if sys.stderr is None:
sys.stderr = open(os.devnull, "w", encoding="utf-8")
PHOTOS = PROJ_DIR / "photos"
TRAIN_DATASET = PROJ_DIR / "output" / "train_dataset"
CHECKPOINT_DIR = PROJ_DIR / "output" / "checkpoints"
SAMPLE_DIR = CHECKPOINT_DIR / "sample"
TRIGGER_DEFAULT = "lm_face_v1"
DIR_CFG = TOOLS_DIR / "gui_dir.txt"
STATE_FILE = TOOLS_DIR / "gui_state.json"
# ③打标 caption 快照:随素材目录走({素材目录}/.caption_backup/captions.json),
# 不放在工具目录——复制/移动素材时备份跟随,多个素材目录互不干扰
def clean_path(p):
"""清理路径:去掉首尾空白和引号(Ctrl+Shift+C 复制的路径带双引号)。非字符串原样返回"""
if not isinstance(p, str):
return p
return p.strip().strip('"').strip("'").strip()
def load_state():
if STATE_FILE.exists():
try:
return json.loads(STATE_FILE.read_text(encoding="utf-8"))
except Exception:
return {}
return {}
def save_state(key, value):
state = load_state()
# 只有字符串(路径/文本)需要去引号清理;数值/布尔原样保存
if isinstance(value, str):
value = clean_path(value)
state[key] = value
STATE_FILE.write_text(json.dumps(state, ensure_ascii=False, indent=1), encoding="utf-8")
def state_val(key, default):
return load_state().get(key, default)
def state_clamp(key, default, lo, hi):
"""读取状态并钳制到 [lo, hi],防止旧版本保存的超范围值导致 gradio 报错"""
v = load_state().get(key, default)
try:
v = float(v)
except (TypeError, ValueError):
return default
return min(hi, max(lo, v))
STYLES = {
"正脸特写": "photorealistic portrait, front view, neutral expression, soft natural lighting, plain background",
"45度半侧脸": "three-quarter view portrait, head slightly turned, neutral expression, natural lighting, plain background",
"半身照": "half body shot, standing, casual clothing, natural lighting, simple background",
"全身照": "full body shot, standing, casual clothing, outdoor scene, natural lighting",
"生活照": "candid lifestyle photo, natural scene, relaxed pose, natural lighting",
"艺术照": "studio portrait, professional photography, soft lighting, clean background",
}
COLORS = {"合格": "#16a34a", "警告": "#d97706", "不合格": "#dc2626"}
def get_photos_dir():
if DIR_CFG.exists():
p = Path(DIR_CFG.read_text(encoding="utf-8").strip())
if p.exists():
return p
return PHOTOS
def thumb_path(photos_dir, fname):
t = photos_dir / ".thumbs" / f"{Path(fname).stem}.jpg"
return str(t) if t.exists() else str(photos_dir / fname)
def run_with_log(func, *args, **kwargs):
"""执行函数并捕获 stdout 输出,返回 (日志文本, 返回值)"""
buf = io.StringIO()
with contextlib.redirect_stdout(buf):
result = func(*args, **kwargs)
return buf.getvalue(), result
# ================= Tab 1: 素材自动处理 =================
def build_auto_tab(vlm_backend):
with gr.Tab("① 素材自动处理"):
gr.Markdown("**auto 流水线**:丢一堆照片 → 自动质量评分/构图分类/裁剪/打标草稿/超量筛选/淘汰")
with gr.Row():
in_dir = gr.Textbox(label="输入照片目录", value=state_val("auto_in", str(PHOTOS)), scale=2)
out_dir = gr.Textbox(label="输出目录", value=state_val("auto_out", str(PROJ_DIR / "photos_auto")), scale=2)
with gr.Row():
trigger = gr.Textbox(label="触发词", value=state_val("auto_trigger", TRIGGER_DEFAULT), scale=1)
limit = gr.Slider(0, 40, value=int(state_clamp("auto_limit", 20, 0, 40)), step=1, label="目标数量(0=不筛选)", scale=1)
with gr.Row():
use_ollama = gr.Checkbox(value=bool(state_val("auto_use_ollama", False)),
label="启用 VLM 自动描述(使用顶部全局 VLM 后端:omlx=小果30B / ollama=本地8B", scale=2)
run_btn = gr.Button("🚀 运行 auto 流水线", variant="primary")
stats = gr.HTML("<span style='color:#666'>就绪</span>")
log = gr.Textbox(label="处理日志", lines=18, interactive=False, elem_id="auto_log",
autoscroll=True, # 配合 Blocks js 里的 tail-follow:内容更新时自动滚到底(用户上翻则暂停)
elem_classes=["log-tail"])
def do_auto(i, o, tr, lim, use_vlm, backend):
"""运行 auto 流水线(实时流式输出日志)。返回 generator:每帧 yield (stats_html, 累计日志)"""
import datetime as _dt
from PIL import Image
i, o = clean_path(i), clean_path(o)
save_state("auto_in", i)
save_state("auto_out", o)
save_state("auto_trigger", tr)
save_state("auto_limit", lim)
save_state("auto_use_ollama", use_vlm)
i, o = Path(i), Path(o)
if not i.exists():
yield '<span style="color:#dc2626">❌ 输入目录不存在</span>', ""
return
person_det = fc.PersonDetector()
det = fc.FaceDetector()
if lim > 0:
n1 = max(1, round(lim * 0.5))
n2 = max(1, round(lim * 0.3))
quota = {"特写": n1, "半身": n2, "全身": max(1, lim - n1 - n2)}
else:
quota = None
_t0 = _dt.datetime.now()
ts = _t0.strftime("%Y%m%d_%H%M%S")
# 日志实时落盘 + 流式返回:用可读的 pipe 捕获 stdout,后台线程读
runlog = PROJ_DIR / "temp" / "auto_runs"
runlog.mkdir(parents=True, exist_ok=True)
logfile = runlog / f"auto_{ts}.log"
header = (f"=== auto 运行 {ts} | GUI v{VERSION} | 后端={backend} | VLM启用={use_vlm} | "
f"开始={_t0.strftime('%H:%M:%S')}\n"
f"input={i} output={o} trigger={tr} limit={lim} "
f"quota={{特写:{quota['特写'] if quota else '不限'}, 半身:{quota['半身'] if quota else '不限'}, "
f"全身:{quota['全身'] if quota else '不限'}}} ===\n")
logfile.write_text(header, encoding="utf-8")
import threading
result_box = {}
def _worker():
import contextlib as _cl
class _LineWriter:
"""行缓冲 writerprint 每行立即落盘(redirect_stdout 到文件默认块缓冲,轮询读不到)"""
def __init__(self, path):
self._f = open(path, "a", encoding="utf-8")
def write(self, s):
self._f.write(s)
self._f.flush()
return len(s)
def flush(self):
self._f.flush()
w = _LineWriter(logfile)
try:
with _cl.redirect_stdout(w):
fc.auto_process(str(i), str(o), det, person_det,
trigger=tr, quota=quota, limit=lim > 0,
use_omlx=use_vlm, backend=backend)
result_box["ok"] = True
except Exception as e:
w.write(f"\n[异常] {repr(e)}\n")
result_box["err"] = repr(e)
result_box["ok"] = False
finally:
w.flush()
w._f.close()
t = threading.Thread(target=_worker, daemon=True)
t.start()
# 主循环:轮询 logfile 尾部,增量 yieldGradio generator 实时刷新)
_seen = len(header)
while t.is_alive():
time.sleep(0.4)
txt = logfile.read_text(encoding="utf-8", errors="replace")
if len(txt) > _seen:
_seen = len(txt)
tail = txt[_seen - 400:].strip().splitlines()
last = tail[-1].strip() if tail else ""
yield (f'<span style="color:#d97706">⏳ 处理中… {last}</span>', txt)
time.sleep(0.3)
txt = logfile.read_text(encoding="utf-8", errors="replace")
_t1 = _dt.datetime.now()
_dur = str(_t1 - _t0).split(".")[0]
# 落盘补全头部(耗时)
try:
with open(logfile, "r", encoding="utf-8") as rf:
c = rf.read()
c = c.replace(f"开始={_t0.strftime('%H:%M:%S')}",
f"开始={_t0.strftime('%H:%M:%S')} 结束={_t1.strftime('%H:%M:%S')} 总耗时={_dur}")
logfile.write_text(c, encoding="utf-8")
except Exception:
pass
if result_box.get("ok"):
# 参数链式传递:输出目录自动成为 ② 换图 的默认目录
DIR_CFG.write_text(str(o), encoding="utf-8")
yield (f'<span style="color:#16a34a">✅ 完成,输出到 {o}(后端 {backend},总耗时 {_dur}</span>'
f'<br>➡️ 下一步:到 <b>② 候选换图</b> 调整素材 + 审改 caption(目录已自动填好,F5 刷新生效)', txt)
else:
yield f'<span style="color:#dc2626">❌ auto 运行失败: {result_box.get("err", "未知错误")}</span>', txt
run_btn.click(do_auto, [in_dir, out_dir, trigger, limit, use_ollama, vlm_backend], [stats, log])
# ================= Tab 2: 素材审核 =================
def build_check_tab():
with gr.Tab("🔍 审核(可选)"):
gr.Markdown("**check(可选工具,非主线)**:红黄绿判定(分辨率/清晰度/人脸/遮挡/重复),生成 HTML 报告。"
"主线流程 ② auto 已内置质量审核,一般不需要单独用。")
with gr.Row():
c_dir = gr.Textbox(label="照片目录", value=state_val("check_dir", str(PHOTOS)), scale=3)
c_btn = gr.Button("🔍 审核", variant="primary")
c_stats = gr.HTML("<span style='color:#666'>就绪</span>")
c_log = gr.Textbox(label="结果", lines=15, interactive=False)
def do_check(d):
d = clean_path(d)
save_state("check_dir", d)
d = Path(d)
if not d.exists():
return '<span style="color:#dc2626">❌ 目录不存在</span>', ""
log_txt, results = run_with_log(fc.analyze_images, str(d), fc.FaceDetector(), verbose=False)
counts = {s: sum(1 for r in results if r["status"] == s) for s in fc.STATUS.values()}
report = Path(d) / "素材审核报告.html"
fc.build_html(results, report)
summary = (f"✅ 合格 {counts['合格']} · 🟡 警告 {counts['警告']} · ❌ 不合格 {counts['不合格']} · "
f"共 {len(results)} 张<br>报告: <a href='file:///{report}'>打开审核报告</a>")
detail = "\n".join(f"[{r['status']}] {r['file']} | {r['reasons'][1]} | {r['reasons'][2][:60]}"
for r in results[:30])
return f'<span style="color:#16a34a">{summary}</span>', detail
c_btn.click(do_check, [c_dir], [c_stats, c_log])
# ================= Tab 3: 智能选图 =================
def build_pick_tab():
with gr.Tab("🎯 选图(可选)"):
gr.Markdown("**pick(可选工具,非主线)**:从大量照片自动挑出多样化组合。"
"主线流程 ② auto 已内置智能筛选(角度保底+去重),一般不需要单独用。")
with gr.Row():
p_dir = gr.Textbox(label="照片目录", value=state_val("pick_dir", str(PHOTOS)), scale=3)
p_count = gr.Slider(5, 40, value=int(state_clamp("pick_count", 20, 5, 40)), step=1, label="选多少张", scale=1)
p_btn = gr.Button("🎯 智能选图", variant="primary")
p_stats = gr.HTML("<span style='color:#666'>就绪</span>")
p_log = gr.Textbox(label="推荐清单", lines=15, interactive=False)
def do_pick(d, n):
d = clean_path(d)
save_state("pick_dir", d)
save_state("pick_count", n)
d = Path(d)
if not d.exists():
return '<span style="color:#dc2626">❌ 目录不存在</span>', ""
log_txt, picked = run_with_log(fc.pick_diverse, fc.analyze_images(str(d), fc.FaceDetector(), verbose=False), n)
detail = "\n".join(f"[{r.get('compose','')}] {r['file']} ({r['status']})" for r in picked)
return (f'<span style="color:#16a34a">✅ 推荐 {len(picked)} 张(目标 {n}</span>', detail)
p_btn.click(do_pick, [p_dir, p_count], [p_stats, p_log])
# ================= Tab 4: 打标 =================
def gen_caption(trigger, style, hair, custom):
if custom.strip():
return custom.strip()
parts = [trigger.strip() or TRIGGER_DEFAULT]
if hair.strip():
parts.append(hair.strip())
parts.append(STYLES[style])
return ", ".join(parts)
def _collect_label_images(base):
"""收集打标目录的图片:优先按 auto 输出的 特写/半身/全身 子目录(有序分组),否则按扁平目录。
返回 [(img_path, kind)]kind 用于分组标签。不再跑人脸检测(素材已过质量关),启动秒开。"""
base = Path(base)
imgs = []
subdirs = [k for k in ("特写", "半身", "全身") if (base / k).is_dir()]
if subdirs:
for kind in subdirs:
for p in sorted((base / kind).glob("*.jpg")):
if ".bak" in p.name or p.name.startswith("换出_"):
continue
imgs.append((p, kind))
else:
for p in sorted(base.iterdir()):
if p.suffix.lower() in {".jpg", ".jpeg", ".png", ".webp"} and p.is_file():
imgs.append((p, "照片"))
return imgs
def build_label_tab(demo, vlm_backend):
with gr.Tab("③ 打标"):
gr.Markdown("**③ 打标**:输入目录 → 点「🔄 加载此目录」**立即重载**(不用 F5)→ 审改 caption "
"(单图可点「🔄 重新生成 caption」,用页面顶部全局 VLM 后端)"
"→ 💾 保存 → 📦 一键整理训练集 → ➡️ 去 **④ 训练**")
with gr.Row():
dir_tb = gr.Textbox(value=str(get_photos_dir()), label="打标目录(auto 输出目录或扁平目录)", scale=3)
dir_btn = gr.Button("🔄 加载此目录", variant="primary", scale=1)
@gr.render(inputs=dir_tb, triggers=[dir_btn.click, demo.load])
def render_cards(dir_value):
dir_value = clean_path(dir_value)
p = Path(dir_value) if dir_value else None
if not p or not p.exists():
gr.HTML('<span style="color:#dc2626">❌ 目录不存在,请检查路径</span>')
return
DIR_CFG.write_text(str(p), encoding="utf-8")
images = _collect_label_images(p)
n_backup = _backup_label_captions(p)
status_html = gr.HTML(
f'<span style="color:#666">已加载 {len(images)} 张({p})· 已备份 {n_backup} 个 caption</span>'
f'<span style="color:#93c5fd;font-size:11px;margin-left:8px">'
f'(误改可用「♻️ 重置所有修改」回退到加载时状态)</span>')
# 打标关键字统计(按类别,含水平/垂直角度)
kw_html = gr.HTML(_keyword_stats_selected(p), elem_id="label_kw_stats")
# 拖拽批量替换的通道(CSS 隐藏但保留 DOM——Gradio5 render 下 visible=False 不渲染 DOMJS 找不到)
kw_old_tb = gr.Textbox(visible=True, elem_id="kw_old_tb", label="")
kw_new_tb = gr.Textbox(visible=True, elem_id="kw_new_tb", label="")
kw_apply_btn = gr.Button(visible=True, elem_id="kw_apply_btn")
# 双击词条的编辑面板(比 prompt 舒服:内联编辑 + 空格拆分)
with gr.Row():
kw_edit_tb = gr.Textbox(
label="✏️ 编辑词条(双击统计项自动填入;多个 tag 用逗号分隔,中英文逗号均可 → 自动拆分)",
placeholder="例如:黑蕾丝抹胸裙, 粉百褶裙, 白手套, 粉高跟鞋", scale=3,
elem_id="kw_edit_tb")
kw_edit_btn = gr.Button("✅ 应用修改", variant="primary", scale=1)
kw_del_btn = gr.Button("🗑️ 删除此词条", variant="stop", scale=1)
kw_edit_status = gr.HTML('')
# 拥有该词条的图缩略图集合(双击词条后显示,确认影响范围)
# 用 HTML 自绘网格:固定小尺寸缩略图 + contain 完整显示 + flex 换行(Gradio Gallery 布局不可控)
kw_gallery = gr.HTML(
'<div style="color:#9ca3af;font-size:12px;padding:4px 0">'
'🖼️ 拥有该词条的图(双击词条后自动显示)</div>',
elem_id="kw_gallery")
with gr.Row():
opt_btn = gr.Button("🧠 一键 LLM 整理全部 Caption(拆短句+合并同义词)", variant="primary", scale=2)
save_btn = gr.Button("💾 保存所有 Caption", variant="primary", scale=1)
reset_btn = gr.Button("♻️ 重置所有修改(回退到加载时备份)", scale=1)
prepare_btn = gr.Button("📦 一键整理训练集 → 训练目录", scale=1)
cards = []
for ci, (img_path, kind) in enumerate(images):
txt_path = img_path.with_suffix(".txt")
existing = txt_path.read_text(encoding="utf-8").strip() if txt_path.exists() else ""
# elem_id 供缩略图点击跳转定位(img_card_{ci} / cap_tb_{ci}
with gr.Group(elem_id=f"img_card_{ci}"):
with gr.Row():
with gr.Column(scale=1, min_width=180):
gr.Image(value=str(img_path), type="filepath", height=170, container=False)
with gr.Column(scale=2):
gr.HTML(f"<div style='font-size:13px;color:#333;'><b>{kind}/{img_path.name}</b></div>")
cap_tb = gr.Textbox(label="Caption(可直接修改)", lines=3, value=existing,
elem_id=f"cap_tb_{ci}")
regen_btn = gr.Button("🔄 重新生成 caption(用顶部所选 VLM", size="sm")
cards.append((img_path, cap_tb))
def _make_regen(img_path=img_path, kind=kind, cap_tb=cap_tb):
def regen(backend):
"""单图重新生成 caption(统一走 fc.vlm_caption)。失败不覆盖原 caption。"""
try:
caption, angle_en = fc.vlm_caption(str(img_path), kind, backend=backend,
trigger=TRIGGER_DEFAULT)
if not caption:
return gr.update(), f'<span style="color:#d97706">⚠️ {img_path.name}: VLM 角度判定失败({backend} 未响应?),caption 未变</span>'
return caption, f'<span style="color:#16a34a">✅ {img_path.name} 已重新生成({backend}{angle_en}),确认后点 💾 保存</span>'
except Exception as e:
return gr.update(), f'<span style="color:#dc2626">❌ {img_path.name} 重新生成失败: {e}caption 未变)</span>'
return regen
def _make_blur_save(img_path=img_path, cap_tb=cap_tb):
def blur_save(cap):
"""Caption 框失焦 → 自动写 txt(编辑即固化文件,统计实时反映),不更新备份。
用户说:直接改 Caption 框里的文字,结束编辑时自动保存到 txt 并刷新统计。"""
try:
img_path.with_suffix(".txt").write_text((cap or "").strip(), encoding="utf-8")
return (f'<span style="color:#16a34a">💾 {img_path.name} caption 已自动保存</span>',
_keyword_stats_selected(p))
except Exception as e:
return f'<span style="color:#dc2626">❌ {img_path.name} 自动保存失败: {e}</span>', _keyword_stats_selected(p)
return blur_save
cap_tb.blur(_make_blur_save(), inputs=cap_tb, outputs=[status_html, kw_html])
regen_btn.click(_make_regen(), inputs=vlm_backend, outputs=[cap_tb, status_html])
def opt_all(*caps):
"""一键 LLM 整理所有 caption(批量一起提交,合并同义词/同类项)。
返回每个 cap_tb 的新值 + 状态。"""
cur = [c or "" for c in caps]
# 有内容的才提交
valid = [c for c in cur if c.strip()]
if not valid:
return (*[gr.update() for _ in cur],
'<span style="color:#d97706">⚠️ 没有可整理的 caption(先加载目录/输入 caption</span>')
status = (f'<span style="color:#d97706">⏳ 正在用 LLM 整理 {len(valid)} 个 caption(合并同义词/拆短句)…</span>')
# 先显示"处理中"(Gradio 同步调用会阻塞,简化:直接调,完成后更新)
mapping = _llm_optimize_captions(valid, backend=state_val("vlm_backend", "omlx-32b"))
outs = []
if not mapping:
for c in cur:
outs.append(gr.update())
return (*outs,
f'<span style="color:#dc2626">❌ LLM 整理失败(后端未响应?)——请检查 VLM 后端,或手动编辑</span>')
changed = 0
for c in cur:
# 空串保护:LLM 返回空时保留原值,不覆盖成空 caption
if c in mapping and mapping[c] and mapping[c] != c:
outs.append(mapping[c])
changed += 1
else:
outs.append(c)
msg = (f'<span style="color:#16a34a">✅ 已整理 {len(mapping)} 个 caption{changed} 个有变化)——'
f'结果已填入下方,确认后点 💾 保存(统计实时更新)</span>')
return (*outs, msg, _keyword_stats_from_texts(outs))
def save_all(*caps):
n = 0
for (img_path, _), cap in zip(cards, caps):
img_path.with_suffix(".txt").write_text(cap.strip(), encoding="utf-8")
n += 1
# 落盘后同步更新备份:备份=最近保存状态,「重置」只撤销未保存的修改,不会抹掉已保存内容
_backup_label_captions(p)
msg = (f'<span style="color:#16a34a">✅ 已保存 {n} 个 caption(原地保存到各图片旁)</span>'
f'<br>➡️ 下一步:点「📦 一键整理训练集」')
# 统计口径与加载一致:从磁盘全量读
return msg, _keyword_stats_selected(p)
def prepare():
train_ds = TRAIN_DATASET
cfg = _train_cfg()
if cfg.get("train_dataset"):
train_ds = Path(cfg["train_dataset"])
train_ds.mkdir(parents=True, exist_ok=True)
for f in train_ds.glob("*"):
if f.is_file():
f.unlink()
copied, skipped = 0, []
from PIL import Image as _Img
for i, (img_path, cap_tb) in enumerate(cards, 1):
txt_path = img_path.with_suffix(".txt")
if not txt_path.exists():
skipped.append(img_path.name)
continue
stem = f"img_{i:03d}"
_Img.open(img_path).convert("RGB").save(train_ds / f"{stem}.jpg", quality=95)
(train_ds / f"{stem}.txt").write_text(txt_path.read_text(encoding="utf-8").strip(), encoding="utf-8")
copied += 1
# 写回 train_config.json:保证 ④ 训练的目录默认值与此处一致(单一事实来源)
cfg["train_dataset"] = str(train_ds)
TRAIN_CONFIG.write_text(json.dumps(cfg, ensure_ascii=False, indent=1), encoding="utf-8")
msg = (f'<span style="color:#16a34a">✅ 已整理 {copied} 张到 {train_ds}</span>'
f'<br>➡️ 下一步:到 <b>④ 训练</b> Tab(目录已自动同步),先点「🔄 重新预缓存」,再启动训练')
if skipped:
msg += f'<br><span style="color:#d97706">⚠️ {len(skipped)} 张缺 caption{", ".join(skipped[:5])}),先保存再整理</span>'
return msg
opt_btn.click(opt_all, inputs=[c[1] for c in cards], outputs=[c[1] for c in cards] + [status_html, kw_html])
save_btn.click(save_all, inputs=[c[1] for c in cards], outputs=[status_html, kw_html])
def apply_kw_replace(old_word, new_word, *caps):
"""批量替换:把所有 caption 中的 old_word 替换成 new_word(写 txt 供统计,不更新备份)。
拖拽/双击触发的统一入口(旧文本 → 新文本)。确认后点 💾 保存才固化到备份锚点。"""
old_word = (old_word or "").strip()
new_word = (new_word or "").strip()
if not old_word or not new_word or old_word == new_word:
return (*caps, _keyword_stats_selected(p),
'<span style="color:#d97706">⚠️ 未执行:源/目标文本不能为空或相同</span>')
n = 0
new_caps = []
for (img_path, _), cap in zip(cards, caps):
if _kw_has(cap, old_word): # 词条级:只替换独立词条,不误伤「黑色长发」
cap = _kw_replace_all(cap, old_word, new_word)
# 编辑即写 txt(统计依赖文件),但不更新备份——只有 💾 保存才更新备份锚点
img_path.with_suffix(".txt").write_text(cap.strip(), encoding="utf-8")
n += 1
new_caps.append(cap)
msg = (f'<span style="color:#16a34a">✅ 已替换 {n} 个 caption'
f'「{old_word}」→「{new_word}」(词条级,已写入 txt;点 💾 保存固化备份锚点)</span>')
# 统计口径与加载一致:只统计选中图(与界面卡片同一口径)
return (*new_caps, _keyword_stats_selected(p), msg)
kw_apply_btn.click(apply_kw_replace,
inputs=[kw_old_tb, kw_new_tb] + [c[1] for c in cards],
outputs=[c[1] for c in cards] + [kw_html, status_html])
def apply_kw_edit(old_word, new_raw, *caps):
"""编辑面板:把 caption 中的 old_word 替换成编辑后的内容。
新文本用空格分隔多个词时 → 自动拆成多个 tag(, 连接),写回 txt,刷新统计。"""
import re as _re
old_word = (old_word or "").strip()
new_raw = (new_raw or "").strip()
if not old_word or not new_raw:
return (*[gr.update() for _ in caps], _keyword_stats_selected(p),
'<span style="color:#d97706">⚠️ 未执行:词条不能为空</span>')
# 逗号分隔(中英文均可)→ 拆成多个 tag;单 tag 原样(前后空白清掉)
# 不用空格分隔:英文词条(Hello Kitty / front view)本身含空格,空格有歧义
_tags = [x.strip() for x in _re.split(r"[,]", new_raw) if x.strip()]
new_word = ", ".join(_tags)
n = 0
new_caps = []
for (img_path, _), cap in zip(cards, caps):
if _kw_has(cap, old_word): # 词条级:只替换独立词条,不误伤「黑色长发」
cap = _kw_replace_all(cap, old_word, new_word)
img_path.with_suffix(".txt").write_text(cap.strip(), encoding="utf-8")
n += 1
new_caps.append(cap)
msg = (f'<span style="color:#16a34a">✅ 已应用:<b>{old_word}</b> → <b>{new_word}</b>{n} 个 caption</span>'
f'<br>空格已拆分多个 tag;已写 txt,可继续编辑或 💾 保存固化')
return (*new_caps, _keyword_stats_selected(p), msg)
kw_edit_btn.click(apply_kw_edit,
inputs=[kw_old_tb, kw_edit_tb] + [c[1] for c in cards],
outputs=[c[1] for c in cards] + [kw_html, kw_edit_status])
def delete_kw_word(old_word, *caps):
"""删除词条:从所有 caption 中移除该词条(写 txt,不更新备份——保存前可重置恢复)。
支持:独立段直接删 / +或、连接段内删 / 段删空后清理逗号。"""
import re as _re
old_word = (old_word or "").strip()
if not old_word:
return (*[gr.update() for _ in caps], _keyword_stats_selected(p), '',
'<span style="color:#d97706">⚠️ 未执行:先双击统计里的词条选择要删除的词</span>')
n, new_caps = 0, []
for (img_path, _), cap in zip(cards, caps):
if not _kw_has(cap, old_word): # 词条级判断(避免「长发」误删「黑色长发」)
new_caps.append(cap)
continue
# 按逗号分段,段内支持 + 、 空格 分隔
kept_segs = []
for seg in cap.split(","):
seg = seg.strip()
if not seg:
continue
# 段内拆分子项(+、/、、 都是并列分隔)
parts = _re.split(r"[+/、]", seg)
parts = [p.strip() for p in parts if p.strip()]
# 删除完全匹配的词条(子项级)
parts = [p for p in parts if p != old_word]
if not parts:
continue # 该段被删空
# 段内保留子项按原分隔符重组
if "+" in seg:
kept_segs.append("+".join(parts))
elif "、" in seg:
kept_segs.append("、".join(parts))
else:
kept_segs.append(" ".join(parts) if len(parts) > 1 else parts[0])
new_cap = ", ".join(kept_segs)
# 清理多余逗号/空格(如 "a,, b" → "a, b"
new_cap = _re.sub(r",\s*,", ",", new_cap).strip(" ,")
if new_cap != cap:
img_path.with_suffix(".txt").write_text(new_cap.strip(), encoding="utf-8")
n += 1
new_caps.append(new_cap)
msg = (f'<span style="color:#16a34a">✅ 已删除词条 <b>{old_word}</b>{n} 个 caption 受影响)</span>'
f'<br>已写 txt;可继续编辑或 💾 保存固化(保存前可 ♻️ 重置恢复)')
return (*new_caps, _keyword_stats_selected(p), '', msg)
kw_del_btn.click(delete_kw_word,
inputs=[kw_old_tb] + [c[1] for c in cards],
outputs=[c[1] for c in cards] + [kw_html, kw_edit_tb, kw_edit_status])
def show_kw_images(old_word):
"""双击词条(kw_old_tb 被 JS 填入)→ 找出 caption 含该词的选中图,渲染缩略图网格。
固定小缩略图(宽96/高128+ contain 完整显示 + flex 换行,不裁剪不滚动。"""
import html as _html
old_word = (old_word or "").strip()
if not old_word:
return ('<div style="color:#9ca3af;font-size:12px;padding:4px 0">'
'🖼️ 拥有该词条的图(双击词条后自动显示)</div>')
hits = [] # (图路径, 卡片索引)
for ci, (img_path, _kind) in enumerate(cards):
txt = img_path.with_suffix(".txt")
try:
cap_txt = txt.read_text(encoding="utf-8") if txt.exists() else ""
# 词条级:只显示独立词条 == old_word 的图(「长发」不含「黑色长发」的图)
if _kw_has(cap_txt, old_word):
hits.append((str(img_path), ci))
except Exception:
continue
if not hits:
return (f'<div style="color:#d97706;font-size:12px;padding:4px 0">'
f'🖼️ 没有图含「{_html.escape(old_word)}」</div>')
items = []
for p, ci in hits:
items.append(
f'<div style="text-align:center;margin:4px">'
f'<img src="/gradio_api/file/{_html.escape(p.replace(chr(92), "/"))}" '
f'style="width:96px;height:128px;object-fit:contain;border-radius:6px;'
f'border:1px solid #374151;background:#1f2937;cursor:pointer" '
f'title="点击跳到下方编辑该图的 Caption" '
f'onclick="scrollToKwCard({ci})">'
f'<div style="font-size:10px;color:#9ca3af;margin-top:2px;max-width:96px;'
f'overflow:hidden;text-overflow:ellipsis;white-space:nowrap">'
f'{_html.escape(Path(p).name)}</div></div>')
inner = "".join(items)
return (f'<div style="color:#e5e7eb;font-size:12px;padding:4px 0">'
f'🖼️ 拥有「{_html.escape(old_word)}」的图({len(hits)} 张)</div>'
f'<div style="display:flex;flex-wrap:wrap;gap:4px;padding:4px 0">{inner}</div>')
# kw_old_tb 被 JS(双击/拖拽)填值后 change → 刷新缩略图集合
kw_old_tb.change(show_kw_images, inputs=kw_old_tb, outputs=kw_gallery)
def reset_all(*caps):
"""用最近落盘状态(加载/保存/替换时的快照)恢复所有 caption(写文件 + 刷新文本框 + 刷新统计)。
只撤销未保存的修改,已保存的内容不会丢。备份缺失的保持当前值不变。"""
bp = _label_backup_path(p)
if not bp.exists():
return (*[gr.update() for _ in caps], _keyword_stats_selected(p),
'<span style="color:#d97706">⚠️ 未找到备份(请先重新加载目录)</span>')
try:
snap = json.loads(bp.read_text(encoding="utf-8"))
except Exception as e:
return (*[gr.update() for _ in caps], _keyword_stats_selected(p),
f'<span style="color:#dc2626">❌ 备份读取失败: {e}</span>')
new_caps, restored = [], 0
for (img_path, _), cur in zip(cards, caps):
key = str(img_path.with_suffix(".txt"))
if key in snap:
new_caps.append(snap[key])
img_path.with_suffix(".txt").write_text(snap[key], encoding="utf-8")
restored += 1
else:
new_caps.append(cur)
msg = (f'<span style="color:#16a34a">✅ 已恢复到最近保存/加载状态:{restored} 个 caption(已写文件)</span>'
f'<br>提示:未保存的修改已撤销;再次「🔄 加载此目录」或「💾 保存」会更新备份锚点')
# 统计口径与加载一致:只统计选中图(与界面卡片同一口径)
return (*new_caps, _keyword_stats_selected(p), msg)
reset_btn.click(reset_all, inputs=[c[1] for c in cards],
outputs=[c[1] for c in cards] + [kw_html, status_html])
prepare_btn.click(prepare, outputs=status_html)
# ================= Tab: 候选换图 =================
_SWAP_DET = [None]
def _swap_detector():
"""FaceDetector 懒加载单例(YuNet 初始化耗时,避免每次新建)"""
if _SWAP_DET[0] is None:
_SWAP_DET[0] = fc.FaceDetector()
return _SWAP_DET[0]
_ANGLE_H_OPTS = ["front view", "three-quarter view", "side view"]
_ANGLE_V_OPTS = ["平视", "high angle view", "low angle view"]
_EXPR_OPTS = ["微笑", "露齿笑", "大笑", "中性", "严肃", "惊讶", "其他"]
_COMP_OPTS = ["特写", "半身", "全身"]
_COMP_TEMPLATE = {"特写": "photorealistic portrait", "半身": "half body shot", "全身": "full body shot"}
_ANGLE_ZH = {"front view": "正面", "three-quarter view": "前侧", "side view": "侧面",
"high angle view": "俯拍", "low angle view": "仰拍"}
def _parse_caption(cap):
"""从 caption 解析结构化字段。返回 dict(h, v, comp, expr, desc_tail)"""
cap = cap or ""
h, v, comp, expr, tail = "front view", "平视", "特写", "", ""
# 角度:提取 水平/垂直 角度词
angs = [a for a in _ANGLE_H_OPTS + ["high angle view", "low angle view"] if a in cap]
h = angs[0] if angs else "front view"
v = "平视"
for a in angs:
if a in ("high angle view", "low angle view"):
v = a
h = next((x for x in angs if x in _ANGLE_H_OPTS), "front view")
# 构图模板
for c, tpl in _COMP_TEMPLATE.items():
if tpl in cap:
comp = c
break
# desc:模板之后的部分
parts = [p.strip() for p in cap.split(",")]
# parts[0]=trigger, parts[1]=角度+模板(可能含逗号), parts[2:]=desc
# 找 desc 起点:跳过 trigger 和 角度/模板部分
idx = 0
for i, p in enumerate(parts):
if i > 0 and not any(a in p for a in _ANGLE_H_OPTS + ["high angle view", "low angle view"]) \
and not any(t in p for t in _COMP_TEMPLATE.values()):
idx = i
break
desc_parts = parts[idx:]
tail = ", ".join(desc_parts) if desc_parts else ""
# 表情:从段里提取「标准表情词」本身(不是整段——整段可能是"微笑,淡妆口红"等长句,喂给下拉框会报
# "Value is not in the list of choices");匹配不到给默认"微笑"(下拉框 choices 必须命中)
expr = "微笑"
if desc_parts:
for dp in desc_parts:
hit = next((k for k in _EXPR_OPTS if k in dp), None)
if hit:
expr = hit
break
return {"h": h, "v": v, "comp": comp, "expr": expr, "tail": tail}
def _build_caption(h, v, comp, expr, tail, trigger="lm_face_v1"):
"""按结构化字段重建 caption(保留 desc 尾部;表情用关键词替换定位)"""
angle = h if v == "平视" else f"{h}, {v}"
tpl = _COMP_TEMPLATE.get(comp, "photorealistic portrait")
if not tail:
return f"{trigger}, {angle} {tpl}"
tparts = [p.strip() for p in tail.split(",")]
# 替换表情元素(首个含表情关键词的段)
replaced = False
for i, tp in enumerate(tparts):
if any(k in tp for k in _EXPR_OPTS):
tparts[i] = expr
replaced = True
break
if not replaced:
if comp == "特写":
tparts.insert(0, expr)
else:
tparts.insert(1, expr) # natural lighting 之后
return f"{trigger}, {angle} {tpl}, {', '.join(tparts)}"
def _stats_summary(base_dir):
"""统计已选素材的 水平角度/垂直角度/表情/构图 分布(不含重复的"已选N张"——status 行已有)"""
from collections import Counter
selected, _ = _swap_collect(base_dir)
h_c, v_c, e_c, comp_c = Counter(), Counter(), Counter(), Counter()
for p, _lbl in selected:
txt = Path(p).with_suffix(".txt")
cap = txt.read_text(encoding="utf-8").strip() if txt.exists() else ""
d = _parse_caption(cap)
h_c[_ANGLE_ZH.get(d["h"], d["h"])] += 1
v_c[d["v"]] += 1
e_c[d["expr"] or "?"] += 1
comp_c[d["comp"]] += 1
def fmt(c):
return " ".join(f"{k}:{n}" for k, n in c.most_common())
return (f"<div style='font-size:14px;color:#111;line-height:1.8;font-weight:500;"
f"background:#f0f4f8;border:1px solid #c8d4e0;border-radius:6px;padding:8px 12px;'>"
f"水平: {fmt(h_c)}<br>垂直: {fmt(v_c)}<br>表情: {fmt(e_c)}<br>构图: {fmt(comp_c)}</div>")
def _stats_advice(base_dir, total=None):
"""总数建议(唯一保留在下方的大提示):总数 过高/偏少"""
if total is None:
selected, _ = _swap_collect(base_dir)
total = len(selected)
SUGGEST_MAX_TOTAL = 30
if total > SUGGEST_MAX_TOTAL:
t = (f"图片总数 {total},高于建议上限 {SUGGEST_MAX_TOTAL}"
f"(建议精简到 ≤{SUGGEST_MAX_TOTAL},否则训练时间长且易过拟合)")
return (f'<div style="font-size:12px;color:#ffd666;background:rgba(255,214,102,0.12);'
f'border:1px solid #a16207;border-radius:6px;padding:6px 10px;line-height:1.7">⚠️ {t}</div>')
if total < 15:
t = f"图片总数 {total},偏少(建议 ≥15,尤其补 大笑/俯拍/仰拍 等稀缺素材)"
return (f'<div style="font-size:12px;color:#ffd666;background:rgba(255,214,102,0.12);'
f'border:1px solid #a16207;border-radius:6px;padding:6px 10px;line-height:1.7">⚠️ {t}</div>')
return '<span style="font-size:12px;color:#4ade80">✅ 素材分布健康,可整理训练集</span>'
def _field_warn_html(warns):
"""字段级紧凑警告:[(名称, 数量, 界限, 不够/过多)] → 行内警告 HTML(空则返回 """""
parts = []
for k, n, bound, kind in warns:
if kind == "不够" and n < bound:
parts.append(f"⚠️{k} {n} 不够(≥{bound})")
elif kind == "过多" and n > bound:
parts.append(f"⚠️{k} {n} 过多(≤{bound})")
if not parts:
return ""
return f'<span class="field-warn">{"".join(parts)}</span>'
def _stats_fields(base_dir):
"""返回 4 个字段的分布小字 HTML(各字段后带自己的 不够/过多 警告,同一行)+ 总数建议"""
from collections import Counter
selected, _ = _swap_collect(base_dir)
total = len(selected)
h_c, v_c, e_c, comp_c = Counter(), Counter(), Counter(), Counter()
for p, _lbl in selected:
txt = Path(p).with_suffix(".txt")
cap = txt.read_text(encoding="utf-8").strip() if txt.exists() else ""
d = _parse_caption(cap)
h_c[_ANGLE_ZH.get(d["h"], d["h"])] += 1
v_c[d["v"]] += 1
e_c[d["expr"] or "?"] += 1
comp_c[d["comp"]] += 1
def fmt(c, keys):
# 标签:亮灰小字;数字:亮色加粗稍大(适配深色背景)
return ' '.join(
f'<span style="font-size:11px;color:#9ca3af">{k}:</span>'
f'<span style="font-size:14px;color:#f3f4f6;font-weight:700">{c.get(k,0)}</span>'
for k in keys)
S = '<span class="field-stats">'
E = '</span>'
h_w = _field_warn_html([("正面", h_c.get("正面", 0), 18, "过多"),
("前侧", h_c.get("前侧", 0), 2, "不够"),
("侧面", h_c.get("侧面", 0), 1, "不够")])
v_w = _field_warn_html([("俯拍", v_c.get("high angle view", 0), 1, "不够"),
("仰拍", v_c.get("low angle view", 0), 1, "不够")])
e_w = _field_warn_html([("大笑", e_c.get("大笑", 0), 1, "不够"),
("露齿笑", e_c.get("露齿笑", 0), 2, "不够"),
("严肃", e_c.get("严肃", 0), 1, "不够"),
("惊讶", e_c.get("惊讶", 0), 1, "不够")])
comp_w = _field_warn_html([("特写", comp_c.get("特写", 0), 20, "过多"),
("半身", comp_c.get("半身", 0), 3, "不够"),
("全身", comp_c.get("全身", 0), 1, "不够")])
return (f"{S}{fmt(h_c,['正面','前侧','侧面'])}{E}{h_w}",
f"{S}{fmt(v_c,['平视','high angle view','low angle view'])}{E}{v_w}",
f"{S}{fmt(e_c,['微笑','露齿笑','大笑','中性','严肃','惊讶'])}{E}{e_w}",
f"{S}{fmt(comp_c,['特写','半身','全身'])}{E}{comp_w}",
_stats_advice(base_dir, total))
def _swap_collect(base_dir):
"""收集 auto 输出目录的 已选素材(特写/半身/全身)+ 未选中候选"""
base = Path(base_dir)
selected, cands = [], []
for kind in ("特写", "半身", "全身"):
d = base / kind
if d.exists():
for p in sorted(d.glob("*.jpg")):
if ".bak" in p.name or p.name.startswith("换出_"):
continue
selected.append((str(p), f"{kind}/{p.name}"))
ud = base / "未选中"
if ud.exists():
for p in sorted(ud.iterdir()):
if p.suffix.lower() in {".jpg", ".jpeg", ".png", ".webp"} and p.is_file():
cands.append((str(p), p.name))
return selected, cands
def _bust(items):
"""显示层 cache-bust:把图复制到带时间戳的临时目录,路径每次不同 → 浏览器强制重拉。
操作层(sp/cp State)仍用真实路径,此处只影响画廊显示。"""
import shutil as _sh
ts_dir = Path(tempfile.gettempdir()) / "swap_display" / str(int(time.time() * 1000))
ts_dir.mkdir(parents=True, exist_ok=True)
out = []
for p, label in items:
try:
src = Path(p)
dst = ts_dir / src.name
_sh.copy2(src, dst)
out.append((str(dst), label))
except Exception:
out.append((p, label))
return out
# 打标中应排除的固定模板词/结构性词(不算"可选取的测试 prompt 关键词"
_FIXED_CAP_WORDS = {
"lm_face_v1",
# 角度完整词不在这里过滤(应参与统计归入角度分类);只过滤拆词碎片
"front", "three-quarter", "side", "high", "low", "angle", "view", "degree",
# 构图模板(完整 + 碎片)
"photorealistic portrait", "half body shot", "full body shot",
"photorealistic", "portrait", "half", "body", "shot", "full",
# 光线
"soft natural lighting", "natural lighting", "bright lighting", "studio lighting",
"soft", "natural", "lighting", "bright", "studio",
# 表情模板词(neutral 是默认模板词,其余表情是可选取特征)
"neutral expression", "neutral", "expression",
}
def _llm_optimize_captions(captions, backend="omlx-32b", timeout=180):
"""一键 LLM 整理所有 caption(批量一起提交,让 LLM 合并同类项/同义词):
- 长句拆成逗号分隔的短 tag
- 合并同义词("长发高马尾" 和 "高马尾长发" 统一)
- 保留触发词、角度、构图等结构
返回 {原caption: 优化后caption};失败返回 None。
"""
import json as _json
import urllib.request
if not captions:
return {}
# 构造批量整理请求:所有 caption 编号,LLM 返回同名编号的优化结果
lines = []
for i, c in enumerate(captions, 1):
lines.append(f"[{i}] {c}")
prompt = (
"你是 LoRA 训练打标整理器。下面是一批人物照片的 caption,每行格式 `[编号] 内容`。\n"
"请统一优化每个 caption\n"
"1. 把长句描述拆成逗号分隔的短 tag(如 '黑色齐刘海高马尾长发粉色蝴蝶结发饰' → '齐刘海, 高马尾, 长发, 蝴蝶结发饰')\n"
"2. 合并同义词/同类项(如 '长发高马尾' 和 '高马尾长发' 统一为一种写法;'粉色的蝴蝶结' 和 '粉色蝴蝶结' 统一)\n"
"3. 保留开头的触发词 lm_face_v1、角度(front view 等)、构图(photorealistic portrait 等)不动\n"
"4. 每个词段独立、简洁、不重复,逗号分隔,不要多余解释\n"
"5. 不要输出任何思考过程/解释/开场白,直接按编号输出结果行\n"
"严格按同样编号输出每行:`[编号] 优化后的caption`,保持顺序和数量一致。\n\n"
+ "\n".join(lines)
)
try:
# 按后端发请求(sensenova 需 Beareromlx 无 authollama 用 /api/chat 格式)
from face_checker import _resolve_model, SENSENOVA_KEY, OLLAMA_MODEL
url, model = _resolve_model(backend)
payload = {
"model": model,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": 4096,
}
headers = {"Content-Type": "application/json"}
if backend == "sensenova":
headers["Authorization"] = f"Bearer {SENSENOVA_KEY}"
req = urllib.request.Request(url, data=_json.dumps(payload).encode(),
headers=headers, method="POST")
with urllib.request.urlopen(req, timeout=timeout) as resp:
result = _json.loads(resp.read())
msg = result["choices"][0]["message"]
text = (msg.get("content") or "").strip()
if not text:
# sensenova 等 thinking 模型可能把内容放 reasoning
text = (msg.get("reasoning") or "").strip()
# 去掉思考过程前缀
text = re.sub(r"^Thinking\s*Process:?\s*", "", text, flags=re.I)
# 解析 [编号] 行
import re
mapping = {}
for m in re.finditer(r"\[(\d+)\]\s*(.+)", text):
idx = int(m.group(1))
cap = m.group(2).strip()
if 1 <= idx <= len(captions):
mapping[captions[idx - 1]] = cap
return mapping if mapping else None
except Exception:
return None
# 打标关键字分类词典(关键词 → 类别)。类别顺序即展示顺序。匹配规则:完整词或包含词。
# 每类一组 (关键词, 匹配方式),匹配方式: "=" 精确 / "in" 包含
_KEYWORD_CATS = [
("表情", [("微笑", "="), ("露齿笑", "="), ("大笑", "="), ("中性", "="), ("严肃", "="), ("惊讶", "=")]),
("水平角度", [("front view", "="), ("three-quarter view", "="), ("side view", "="),
("正面", "="), ("前侧", "="), ("侧面", "=")]),
("垂直角度", [("high angle view", "="), ("low angle view", "="),
("平视", "="), ("俯拍", "="), ("仰拍", "=")]),
("构图", [("photorealistic portrait", "="), ("half body shot", "="), ("full body shot", "="),
("特写", "="), ("半身", "="), ("全身", "=")]),
("妆", [("素颜", "="), ("浓妆口红", "="), ("浓妆", "="), ("淡妆", "="), ("素妆", "="), ("妆", "in")]),
("发型", [("马尾", "in"), ("直发", "in"), ("卷发", "in"), ("长发", "in"), ("短发", "in"),
("束发", "in"), ("盘发", "in"), ("扎发", "in"), ("齐刘海", "in"), ("刘海", "in"),
("丸子头", "in"), ("发箍", "in"), ("高发髻", "in")]),
("配饰", [("鸭舌帽", "in"), ("墨镜", "in"), ("发饰", "in"), ("发夹", "in"), ("发卡", "in"),
("钻石项链", "in"), ("四叶草项链", "in"), ("项链", "in"), ("蝴蝶结", "in"),
("丝带", "in"), ("耳环", "in")]),
("姿势", [("坐姿", "in"), ("坐", "="), ("蹲姿", "in"), ("站立", "="), ("站", "="),
("双手扶肩", "in"), ("双手比划", "in"), ("举玩偶", "in"), ("叉腰", "in"),
("提裙", "in"), ("双臂张开", "in"), ("比耶", "in"), ("牵手", "in"), ("抱", "in")]),
("服装", [("裙", "in"), ("礼服", "in"), ("上衣", "in"), ("外套", "in"), ("牛仔裤", "in"),
("短裤", "in"), ("内搭", "in"), ("套装", "in"), ("衬衫", "in"), ("T恤", "in"),
("T恤", "in"), ("抹胸", "in"), ("蕾丝", "in"), ("手套", "in"), ("高跟鞋", "in"),
("鞋", "in"), ("帽", "in")]),
("光线/背景", [("自然光", "in"), ("室内光", "in"), ("明亮光", "in"), ("柔和光", "in"),
("灯光", "in"), ("房间", "in"), ("草地", "in"), ("花园", "in"), ("楼梯", "in"),
("墙面", "in"), ("沙发", "in"), ("玩偶", "in"), ("旋转木马", "in"),
("凯蒂猫", "in"), ("Hello Kitty", "in"), ("画框", "in"), ("相框", "in")]),
]
def _kw_has(cap, word):
"""词条级判断:caption 中是否存在独立词条 == word(按 , 、 + / 拆分,与统计口径一致)。
「长发」只匹配独立词条「长发」,不会误伤「黑色长发」。"""
import re as _re
word = (word or "").strip()
if not word:
return False
cap = (cap or "").replace("Hello Kitty", "HelloKitty").replace("hello kitty", "HelloKitty")
word = word.replace("Hello Kitty", "HelloKitty")
for seg in _re.split(r"[,]", cap):
for p in _re.split(r"[+/、]", seg):
if p.strip() == word:
return True
return False
def _kw_replace_all(cap, old_word, new_word):
"""词条级替换:caption 中所有独立词条 == old_word 替换为 new_word。
保留段内分隔符(+ / 、)和原始结构;不匹配的词条原样保留。"""
import re as _re
old_word = (old_word or "").strip()
new_word = new_word or ""
# 归一保护形式:caption 与目标词都按统计分词口径处理(Hello Kitty → HelloKitty
cap_norm = (cap or "").replace("Hello Kitty", "HelloKitty").replace("hello kitty", "HelloKitty")
old_norm = old_word.replace("Hello Kitty", "HelloKitty")
out_segs = []
for seg in _re.split(r"[,]", cap_norm):
if not seg.strip():
out_segs.append(seg)
continue
lead = seg[: len(seg) - len(seg.lstrip())] # 段前导空格(保留格式 ", xxx"
parts = _re.split(r"([+/、])", seg.strip()) # 保留分隔符
rebuilt = ""
for p in parts:
if p in ("+", "/", "、"): # 分隔符原样
rebuilt += p
elif p.strip() == old_norm: # 词条级精确命中
rebuilt += new_word
else:
rebuilt += p
out_segs.append(lead + rebuilt)
# 还原保护形式,保持与 caption 原文一致的 "Hello Kitty" 写法
return ",".join(out_segs).replace("HelloKitty", "Hello Kitty")
def _categorize_keyword(w):
"""把关键词归入类别,返回类别名或 None(未分类)。"""
wl = w.lower()
for cat, rules in _KEYWORD_CATS:
for kw, mode in rules:
if mode == "=" and wl == kw.lower():
return cat
if mode == "in" and kw.lower() in wl:
return cat
return None
def _kw_chip(w, n, hl):
"""统计词条 chip:可拖拽/可双击,data-word 存原词(dataset 自动解码 HTML 实体)"""
import html as _html
esc = _html.escape(w, quote=True)
color = "#fbbf24;font-weight:700" if hl else "#e5e7eb"
return (f'<span class="kw-item" draggable="true" data-word="{esc}" '
f'style="color:{color};cursor:grab;padding:0 2px;border-radius:4px" '
f'title="拖拽到另一项=批量替换;双击=修改">{_html.escape(w)}</span>'
f'<span style="color:#6b7280;font-size:11px">×{n}</span>')
def _label_backup_path(base_dir):
"""③打标 caption 快照路径:随素材目录走({素材目录}/.caption_backup/captions.json)。
放在素材目录下而非工具目录——复制/移动素材时备份跟随,多个素材目录互不干扰。"""
return Path(clean_path(base_dir) or ".") / ".caption_backup" / "captions.json"
def _backup_label_captions(base_dir):
"""加载打标目录时,把当前所有 txt 的 caption 快照备份到 LABEL_BACKUP(每次覆盖)。
返回备份的 txt 数量。供「重置所有修改」按钮恢复。"""
base_dir = clean_path(base_dir)
if not base_dir or not Path(base_dir).exists():
return 0
txts = sorted(p for p in Path(base_dir).rglob("*.txt")
if ".bak" not in p.name and "淘汰" not in str(p) and "未选中" not in str(p))
snap = {}
for t in txts:
try:
snap[str(t)] = t.read_text(encoding="utf-8")
except Exception:
continue
bp = _label_backup_path(base_dir)
bp.parent.mkdir(parents=True, exist_ok=True)
bp.write_text(json.dumps(snap, ensure_ascii=False, indent=1), encoding="utf-8")
return len(snap)
def _keyword_stats_from_texts(texts):
"""统计给定 caption 文本列表的关键字频率,按类别分类展示(树状)。
返回 HTML:类别 → 关键词×次数。词条支持拖拽/双击(由页面 JS 处理)。"""
from collections import Counter, defaultdict
import re
cnt = Counter() # 词 -> 次数(未分类的词也统计,最后归"其他")
for cap in texts:
cap = (cap or "").strip()
if not cap:
continue
# 先保护 "Hello Kitty" 整体(不按空格拆散)
cap = cap.replace("Hello Kitty", "HelloKitty").replace("hello kitty", "HelloKitty")
# 保护完整角度词(不按空格拆成碎片)
for _aw in ("front view", "three-quarter view", "side view", "high angle view", "low angle view"):
cap = cap.replace(_aw, _aw.replace(" ", "_"))
for seg in re.split(r"[,]", cap):
seg = seg.strip()
if not seg:
continue
for w in re.split(r"[\s/、]", seg):
w = w.strip(" .,;:")
if not w or len(w) < 2:
continue
wl = w.lower()
if wl in _FIXED_CAP_WORDS or wl.startswith("lm_face"):
continue
# 还原 HelloKitty 和角度下划线(白名单:只还原受保护的角度词,
# 其他含下划线的词保持原样——否则统计显示与 caption 不一致,拖拽/双击替换会匹配失败)
# HelloKitty 可能内嵌在词条里(如 "粉色沙发及HelloKitty玩偶"),都要还原成 "Hello Kitty"
if "HelloKitty" in w:
w = w.replace("HelloKitty", "Hello Kitty")
elif w in ("front_view", "three-quarter_view", "side_view",
"high_angle_view", "low_angle_view"):
w = w.replace("_", " ")
cnt[w] += 1
if not cnt:
return '<span style="color:#d97706">⚠️ 未提取到有效关键词</span>'
# 按类别归类
by_cat = defaultdict(Counter)
others = Counter()
for w, n in cnt.items():
cat = _categorize_keyword(w)
if cat:
by_cat[cat][w] = n
else:
others[w] = n
# 组装树状 HTML
parts = [f'<div style="font-size:13px;line-height:2">'
f'<b style="color:#e5e7eb">📊 打标关键字统计(按类别)</b>'
f'<span style="color:#9ca3af;font-size:12px">{len(texts)} 张素材)</span>'
f'<div style="color:#93c5fd;font-size:11px;margin:2px 0">'
f'🖱 拖拽某一项到另一项 = 批量替换(目标→被拖项);双击某项 = 修改后批量替换</div>']
for cat, _rules in _KEYWORD_CATS:
if cat not in by_cat:
continue
top = by_cat[cat].most_common(15)
inner = " ".join(_kw_chip(w, n, n >= 5) for w, n in top)
parts.append(f'<div style="margin:4px 0"><span style="color:#60a5fa;font-weight:700">▸ {cat}</span> '
f'<span style="color:#9ca3af;font-size:12px">({sum(by_cat[cat].values())}次)</span><br>'
f'<span style="margin-left:14px">{inner}</span></div>')
if others:
top_oth = others.most_common(15)
inner = " ".join(_kw_chip(w, n, False) for w, n in top_oth)
parts.append(f'<div style="margin:4px 0"><span style="color:#9ca3af;font-weight:700">▸ 其他</span><br>'
f'<span style="margin-left:14px">{inner}</span></div>')
parts.append('</div>')
return "".join(parts)
def _keyword_stats_selected(base_dir):
"""③打标 统计:只统计「选中图」的 caption(与 _collect_label_images 同一口径)。
不含 train_dataset 副本/未选中/淘汰——避免同一 caption 被数两遍。"""
base_dir = clean_path(base_dir)
if not base_dir or not Path(base_dir).exists():
return '<span style="color:#dc2626">❌ 目录不存在</span>'
imgs = _collect_label_images(base_dir)
if not imgs:
return '<span style="color:#d97706">⚠️ 没有可统计的选中图</span>'
texts = []
for img_path, _kind in imgs:
txt = Path(img_path).with_suffix(".txt")
try:
cap = txt.read_text(encoding="utf-8").strip() if txt.exists() else ""
except Exception:
cap = ""
texts.append(cap)
return _keyword_stats_from_texts(texts)
def _keyword_stats(train_ds):
"""读目录下所有 txt → 转文本列表 → _keyword_stats_from_texts。
支持:训练集目录(img_*.txt)或打标目录(特写/半身/全身 子目录)。"""
train_ds = clean_path(train_ds)
if not train_ds or not Path(train_ds).exists():
return '<span style="color:#dc2626">❌ 目录不存在</span>'
# 收集目录下所有 *.txt(递归子目录),跳过 .bak 和 auto 流水线内部文件
txts = sorted(p for p in Path(train_ds).rglob("*.txt")
if ".bak" not in p.name and "淘汰" not in str(p) and "未选中" not in str(p))
if not txts:
return '<span style="color:#d97706">⚠️ 没有可统计的 txt(打标目录或训练集目录)</span>'
texts = []
for t in txts:
try:
texts.append(t.read_text(encoding="utf-8").strip())
except Exception:
continue
return _keyword_stats_from_texts(texts)
def _prepare_trainset(base_dir):
"""一键整理训练集:把 auto 输出目录的 已选素材(特写/半身/全身 + caption)复制成训练集 img_001..。
⚠️ 训练集目录必须从 base_dir 派生(base_dir/train_dataset),不能读 train_config.json 的旧值——
否则换目录后会把新素材写进上一次的旧路径(bug:2026-08-09 修复)。"""
from PIL import Image as _Img
base_dir = clean_path(base_dir)
if not base_dir or not Path(base_dir).exists():
return '<span style="color:#dc2626">❌ 目录不存在</span>'
selected, _ = _swap_collect(base_dir)
if not selected:
return '<span style="color:#dc2626">❌ 没有已选素材</span>'
train_ds = Path(base_dir) / "train_dataset" # 永远跟随当前素材目录
cfg = _train_cfg()
train_ds.mkdir(parents=True, exist_ok=True)
for f in train_ds.glob("*"):
if f.is_file():
f.unlink()
copied, skipped = 0, []
for i, (img_path, _label) in enumerate(selected, 1):
txt = Path(img_path).with_suffix(".txt")
if not txt.exists():
skipped.append(Path(img_path).name)
continue
stem = f"img_{i:03d}"
_Img.open(img_path).convert("RGB").save(train_ds / f"{stem}.jpg", quality=95)
(train_ds / f"{stem}.txt").write_text(txt.read_text(encoding="utf-8").strip(), encoding="utf-8")
copied += 1
cfg["train_dataset"] = str(train_ds)
TRAIN_CONFIG.write_text(json.dumps(cfg, ensure_ascii=False, indent=1), encoding="utf-8")
msg = (f'<span style="color:#16a34a">✅ 已整理 {copied} 张到 {train_ds}</span>'
f'<br>➡️ 下一步:到 <b>④ 训练</b> Tab(目录已自动同步),先点「🔄 重新预缓存」,再启动训练')
if skipped:
msg += f'<br><span style="color:#d97706">⚠️ {len(skipped)} 张缺 caption{", ".join(skipped[:5])}),先保存再整理</span>'
return msg
def _find_source_by_orb(crop_path, base_dir):
"""用 ORB 特征匹配找到裁剪图的源图(在 未选中/换出区/final 里),返回源文件路径或 None"""
import cv2 as _cv2
import numpy as np # 模块顶部未导入 numpy,此处必须局部导入(否则 NameError 被 try 吞掉 → 永远返回 None)
try:
orb = _cv2.ORB_create(1500)
bf = _cv2.BFMatcher(_cv2.NORM_HAMMING)
from PIL import Image as _I
crop_rgb = _I.open(crop_path).convert("RGB")
crop_gray = _cv2.cvtColor(np.array(crop_rgb), _cv2.COLOR_RGB2GRAY)
k1, d1 = orb.detectAndCompute(crop_gray, None)
if d1 is None or len(k1) < 10:
return None
base = Path(base_dir)
ud = base / "未选中"
candidates = []
if ud.exists():
for p in ud.iterdir():
if p.suffix.lower() in {".jpg", ".jpeg", ".png", ".webp"} and p.is_file():
candidates.append(p)
fin = Path(r"K:\AI\training\ldf\singled\final")
if fin.exists():
for p in fin.iterdir():
if p.suffix.lower() in {".jpg", ".jpeg", ".png"} and p.is_file():
candidates.append(p)
best, bn = None, 0
for p in candidates:
try:
rgb = _I.open(p).convert("RGB")
gray = _cv2.cvtColor(np.array(rgb), _cv2.COLOR_RGB2GRAY)
k2, d2 = orb.detectAndCompute(gray, None)
if d2 is None:
continue
ms = bf.knnMatch(d1, d2, k=2)
g = [m for m, n in ms if m.distance < 0.75 * n.distance] if ms else []
if len(g) > bn:
bn, best = len(g), p
except Exception:
continue
return best if bn >= 30 else None
except Exception:
return None
def _prep_candidate(cand_path, kind, backend):
"""候选图预处理:检测人脸 + 按类型裁剪(特写裁脸/全身构图/半身原图)+ 统一 caption 生成。
失败抛异常(调用方保证不动任何文件)。"""
import numpy as np
img = fc.load_image(cand_path)
rgb = np.array(img.convert("RGB"))
faces = fc._detect_faces_fast(rgb, _swap_detector())
if not faces:
# 检测器漏检但 VLM 可能确认有脸(如侧脸/遮挡):先让 VLM 确认。
# vlm_caption 返回完整 caption 说明确认有人脸 → 整图作为构图素材加入(不裁脸);否则报错。
caption, angle_en = fc.vlm_caption(str(cand_path), kind, backend=backend, trigger=TRIGGER_DEFAULT)
if caption:
# VLM 确认有脸:构图降级为半身(原图),不裁脸
return img, caption, angle_en or "front view"
raise ValueError(f"{Path(cand_path).name} 未检测到人脸(且 VLM 无法确认)")
f = max(faces, key=lambda x: x[2] * x[3]) if len(faces) > 1 else faces[0]
if kind == "特写":
crop, _ = fc.crop_face_portrait(img, f, margin=2.5)
elif kind == "全身":
crop, _ = fc.crop_fullbody(img, f)
else:
crop = img
caption, angle_en = fc.vlm_caption(str(cand_path), kind, backend=backend, trigger=TRIGGER_DEFAULT)
if not caption:
angle_en = fc.face_angle(f)[1]
caption = f"{TRIGGER_DEFAULT}, {angle_en} {fc.AUTO_TEMPLATES[kind]}"
return crop, caption, angle_en
def build_swap_tab(demo, vlm_backend):
with gr.Tab("② 候选换图"):
gr.Markdown("**② 候选换图**:左边点选要换出的已选素材,右边点选要换入的未选中候选 → 点交换;"
"点「↩️ 取消已选」把选中的已选素材移回未选中候选。"
"换入时自动按类型裁剪(特写裁脸/全身构图)+ VLM 重新打标;换出的图移回未选中。"
"➡️ 换完到 **③ 打标** 审改 caption")
with gr.Row():
base_tb = gr.Textbox(label="auto 输出目录(① 跑完自动填好,含 特写/半身/全身/未选中)", scale=3,
value=state_val("auto_out", ""))
load_btn = gr.Button("🔄 加载", scale=1)
with gr.Row():
prepare_btn = gr.Button("📦 整理训练集(💡 打标全部确认后再点:把 已选素材+caption 复制成训练集 img_001.. → 训练目录)",
variant="primary")
status_html = gr.HTML('<span style="color:#666">先加载目录</span>')
# ── 已选素材筛选(按 水平/垂直/表情/构图 4 维过滤,只影响 已选画廊+操作列表)──
with gr.Row(elem_classes="filter-row"):
gr.Markdown("**筛选已选**", elem_classes="x-label")
f_h = gr.Dropdown(["全部"] + _ANGLE_H_OPTS, value="全部", container=False, scale=1, elem_classes="fld-dd", info="水平")
f_v = gr.Dropdown(["全部"] + _ANGLE_V_OPTS, value="全部", container=False, scale=1, elem_classes="fld-dd", info="垂直")
f_expr = gr.Dropdown(["全部"] + _EXPR_OPTS, value="全部", container=False, scale=1, elem_classes="fld-dd", info="表情")
f_comp = gr.Dropdown(["全部"] + _COMP_OPTS, value="全部", container=False, scale=1, elem_classes="fld-dd", info="构图")
filter_info = gr.HTML('<span class="filter-cnt" style="font-size:12px;color:#9ca3af">未筛选</span>')
clear_filter = gr.Button("✕ 清除筛选", size="sm", scale=0)
with gr.Row():
with gr.Column(scale=1):
sel_gal = gr.Gallery(label="已选素材(点击选择要换出的)", columns=4, height=420, elem_id="sel_gallery",
object_fit="contain", allow_preview=True, show_fullscreen_button=True)
with gr.Column(scale=1):
cand_gal = gr.Gallery(label="未选中候选(点击选择要换入的)", columns=4, height=420, elem_id="cand_gallery",
object_fit="contain", allow_preview=True, show_fullscreen_button=True)
with gr.Row():
with gr.Column(scale=1):
with gr.Accordion("🎛️ 角度/表情/构图 + Caption(可折叠,选完图后点右侧按钮或下方手动裁剪)", open=True):
with gr.Row(elem_classes="field-row"):
gr.Markdown("**水平角度**", elem_classes="x-label")
h_dd = gr.Dropdown(_ANGLE_H_OPTS, value="front view", container=False, scale=1, elem_classes="fld-dd", allow_custom_value=True)
h_dist = gr.HTML('<span class="field-stats"><span style="font-size:11px;color:#9ca3af">正面:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">前侧:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">侧面:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span></span>')
with gr.Row(elem_classes="field-row"):
gr.Markdown("**垂直角度**", elem_classes="x-label")
v_dd = gr.Dropdown(_ANGLE_V_OPTS, value="平视", container=False, scale=1, elem_classes="fld-dd", allow_custom_value=True)
v_dist = gr.HTML('<span class="field-stats"><span style="font-size:11px;color:#9ca3af">平视:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">俯拍:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">仰拍:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span></span>')
with gr.Row(elem_classes="field-row"):
gr.Markdown("**表情**", elem_classes="x-label")
expr_dd = gr.Dropdown(_EXPR_OPTS, value="微笑", container=False, scale=1, elem_classes="fld-dd",
allow_custom_value=True) # 页面加载状态恢复时下拉可能传空值,允许临时值防报错
expr_dist = gr.HTML('<span class="field-stats"><span style="font-size:11px;color:#9ca3af">微笑:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">露齿笑:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">大笑:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">中性:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">严肃:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">惊讶:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span></span>')
with gr.Row(elem_classes="field-row"):
gr.Markdown("**构图**", elem_classes="x-label")
comp_dd = gr.Dropdown(_COMP_OPTS, value="特写", container=False, scale=1, elem_classes="fld-dd", allow_custom_value=True)
comp_dist = gr.HTML('<span class="field-stats"><span style="font-size:11px;color:#9ca3af">特写:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">半身:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span> <span style="font-size:11px;color:#9ca3af">全身:</span><span style="font-size:14px;color:#f3f4f6;font-weight:700">0</span></span>')
advice_html = gr.HTML('<span style="font-size:12px;color:#9ca3af">加载后显示素材建议(总数/角度/表情/构图是否合理)</span>')
cap_tb = gr.Textbox(label="完整 Caption", lines=3, interactive=True,
placeholder="点选图后显示其 caption(可直接编辑)")
with gr.Row():
cap_regen = gr.Button("🔄 重新打标(用顶部 VLM", scale=1)
cap_save = gr.Button("💾 保存 Caption", variant="primary", scale=1)
cap_status = gr.HTML('<span style="color:#666">点选图后编辑 caption</span>')
with gr.Column(scale=1):
cancel_btn = gr.Button("↩️ 取消已选(移入候选)", variant="secondary", elem_id="btn_cancel_sel")
swap_btn = gr.Button("⇄ 交换选中项", variant="primary")
add_btn = gr.Button("➕ 添加选中候选(不替换,追加为新素材)", elem_id="btn_add_cand")
kind_dd = gr.Dropdown(["特写", "半身", "全身"], value="特写", elem_id="cand_kind_dd",
label="添加为类型(系统按类型自动裁剪+打标)")
sel_idx = gr.State(-1)
cand_idx = gr.State(-1)
sel_paths = gr.State([])
cand_paths = gr.State([])
# ---- 手动裁剪面板:点选任一图后在此精修 ----
with gr.Accordion("✂️ 手动裁剪(点选 已选/候选 图后,在这里调整并保存)", open=True):
with gr.Row():
with gr.Column(scale=1):
crop_src = gr.Image(label="待裁剪图(点选左/右栏图片自动载入)", type="filepath", height=340)
with gr.Row():
cr_x0 = gr.Slider(0, 100, value=0, step=1, label="左 %")
cr_y0 = gr.Slider(0, 100, value=0, step=1, label="上 %")
with gr.Row():
cr_x1 = gr.Slider(0, 100, value=100, step=1, label="右 %")
cr_y1 = gr.Slider(0, 100, value=100, step=1, label="下 %")
with gr.Column(scale=1):
crop_prev = gr.Image(label="裁剪预览", type="filepath", height=340)
crop_file = gr.State("")
crop_btn = gr.Button("✂️ 应用裁剪(自动备份原图)", variant="primary")
crop_status = gr.HTML('<span style="color:#666">点选图片后设置裁剪框</span>')
def load(base_dir):
base_dir = clean_path(base_dir)
if not base_dir or not Path(base_dir).exists():
return [], [], [], [], '<span style="color:#dc2626">❌ 目录不存在</span>', ''
selected, cands = _swap_collect(base_dir)
sp = [s[0] for s in selected]
cp = [c[0] for c in cands]
# 画廊显示用时间戳镜像(路径每次不同,强制浏览器重拉)
display_sel = _bust(selected)
display_cand = _bust(cands)
msg = f'<span style="color:#16a34a">✅ 已选 {len(sp)} 张 / 候选 {len(cp)} 张</span>'
return display_sel, display_cand, sp, cp, msg, *_stats_fields(base_dir)
def _filter_selected(base_dir, fh, fv, fexpr, fcomp):
"""按 水平/垂直/表情/构图 4 维筛选已选素材。
返回 (已选画廊显示, 过滤后的 sel_paths, 筛选状态文案)。
只读 caption 解析,不改任何文件;筛选结果作为后续点选/交换/取消的操作列表。"""
base_dir = clean_path(base_dir)
if not base_dir or not Path(base_dir).exists():
return gr.update(), [], '<span style="color:#dc2626">❌ 目录不存在</span>'
selected, _ = _swap_collect(base_dir)
sp = [s[0] for s in selected]
if fh == "全部" and fv == "全部" and fexpr == "全部" and fcomp == "全部":
return _bust(selected), sp, f'<span style="font-size:12px;color:#9ca3af">未筛选 · 已选 {len(sp)} 张</span>'
out = []
for real in sp:
cap = ""
txt = Path(real).with_suffix(".txt")
if txt.exists():
try:
cap = txt.read_text(encoding="utf-8").strip()
except Exception:
cap = ""
d = _parse_caption(cap)
if fh != "全部" and d["h"] != fh:
continue
if fv != "全部" and d["v"] != fv:
continue
if fexpr != "全部" and fexpr not in d["expr"]:
continue
if fcomp != "全部" and d["comp"] != fcomp:
continue
out.append((real, Path(real).name))
conds = [c for c, v in (("水平", fh), ("垂直", fv), ("表情", fexpr), ("构图", fcomp)) if v != "全部"]
info = (f'<span style="font-size:12px;color:#16a34a">筛选 {len(out)}/{len(sp)} 张'
f'{"、".join(conds)}</span>')
return _bust(out), [p for p, _ in out], info
def _clear_filter(base_dir):
"""清除筛选:恢复全部已选"""
return _filter_selected(base_dir, "全部", "全部", "全部", "全部")
def _reapply_filter(base_dir, fh, fv, fexpr, fcomp, cands):
"""写操作后统一重刷画廊:已选按当前筛选条件过滤,候选保持全量。
返回 (_bust(已选筛选结果), _bust(cands全量), 过滤后paths, 全量cand_paths)"""
sel_disp, sp2, _ = _filter_selected(base_dir, fh, fv, fexpr, fcomp)
cp2 = [c[0] for c in cands]
return sel_disp, _bust(cands), sp2, cp2
def _load_crop(base_dir, paths, idx):
"""合并处理器:evt.index 直接用(避免 State 竞态导致 caption 错位)"""
if idx < 0 or idx >= len(paths):
return None, "", "", "", "front view", "平视", "微笑", "特写"
real = paths[idx]
cap = ""
txt = Path(real).with_suffix(".txt")
if txt.exists():
cap = txt.read_text(encoding="utf-8").strip()
d = _parse_caption(cap)
return _bust([(real, Path(real).name)])[0][0], real, cap, \
f'<span style="color:#333">{Path(real).name}</span>', d["h"], d["v"], d["expr"], d["comp"]
def on_sel(evt: gr.SelectData, base_dir, sp):
return evt.index, *_load_crop(base_dir, sp, evt.index)
def on_cand(evt: gr.SelectData, base_dir, cp):
return evt.index, *_load_crop(base_dir, cp, evt.index)
def cap_rebuild(h, v, expr, comp, cap):
"""任一下拉框变化 → 重建完整 caption(页面加载状态恢复时下拉可能传空值,需兜底)"""
d = _parse_caption(cap)
if expr not in _EXPR_OPTS:
expr = d["expr"] if d["expr"] in _EXPR_OPTS else "微笑"
if h not in _ANGLE_H_OPTS:
h = d["h"] if d["h"] in _ANGLE_H_OPTS else "front view"
if v not in _ANGLE_V_OPTS:
v = d["v"] if d["v"] in _ANGLE_V_OPTS else "平视"
if comp not in _COMP_OPTS:
comp = d["comp"] if d["comp"] in _COMP_OPTS else "特写"
return _build_caption(h, v, comp, expr, d["tail"])
def cap_regen_fn(real, backend):
"""用 VLM 重新生成 caption(按图类型自动决定特写/半身描述)"""
if not real or not Path(real).exists():
return gr.update(), '<span style="color:#dc2626">❌ 请先点选一张图</span>'
kind = Path(real).parent.name # 特写/半身/全身
if kind not in ("特写", "半身", "全身"):
return gr.update(), f'<span style="color:#dc2626">❌ 无法识别类型: {kind}</span>'
caption, angle_en = fc.vlm_caption(str(real), kind, backend=backend, trigger=TRIGGER_DEFAULT)
if not caption:
return gr.update(), f'<span style="color:#d97706">⚠️ VLM 判定失败({backend}),未生成</span>'
return caption, f'<span style="color:#16a34a">✅ 已重新打标({backend},角度 {angle_en}</span>'
def cap_save_fn(base_dir, real, cap_text, comp, sp, cp, fh, fv, fexpr, fcomp):
"""保存 caption;若构图与当前目录不符,则移动文件到对应目录并重命名(face/half/full)。
保存后同步界面当前图:若该图仍在筛选结果则保持,否则自动切到筛选后第一张。"""
if not real or not Path(real).exists():
return (gr.update(), gr.update(), sp, cp, gr.update(), gr.update(), gr.update(),
'<span style="color:#dc2626">❌ 请先点选一张图</span>',
gr.update(), gr.update(), gr.update(), gr.update(),
gr.update(), gr.update(), gr.update(), gr.update(), gr.update())
base_dir = clean_path(base_dir)
p = Path(real)
cur_dir = p.parent.name # 特写/半身/全身
prefix_map = {"特写": "face", "半身": "half", "全身": "full"}
msg_extra = ""
if cur_dir != comp:
# 需要移动:目标目录 + 新编号
tgt_dir = Path(base_dir) / comp
tgt_dir.mkdir(parents=True, exist_ok=True)
pre = prefix_map.get(comp, "face")
nums = []
for f in tgt_dir.glob(f"{pre}_*.jpg"):
parts = f.stem.split("_")
if len(parts) == 2 and parts[1].isdigit():
nums.append(int(parts[1]))
n = max(nums) + 1 if nums else 1
new_name = f"{pre}_{n:03d}.jpg"
new_path = tgt_dir / new_name
p.rename(new_path)
new_path.with_suffix(".txt").write_text(cap_text.strip(), encoding="utf-8")
msg_extra = f';已移到 {comp}/{new_name}(原 {cur_dir}/{p.name}'
real = str(new_path)
else:
Path(real).with_suffix(".txt").write_text(cap_text.strip(), encoding="utf-8")
# 刷新画廊(已选按当前筛选条件重过滤)
selected, cands = _swap_collect(base_dir)
sel_disp, cand_disp, sp2, cp2 = _reapply_filter(base_dir, fh, fv, fexpr, fcomp, cands)
saved_real = str(real)
# 同步界面当前图:若保存的图仍在筛选结果中 → 保持它;若已移出筛选 → 自动切到筛选后第一张
if saved_real in sp2:
cur = saved_real
status_note = f'✅ Caption 已保存{msg_extra}'
elif sp2:
cur = sp2[0]
status_note = (f'⚠️ 原图已移出当前筛选(属性不符),已自动切到筛选后的第一张:'
f'{Path(cur).name}。Caption 已保存{msg_extra}')
else:
cur = ""
status_note = f'⚠️ Caption 已保存{msg_extra},但当前筛选下已无已选素材'
if cur:
cap2 = Path(cur).with_suffix(".txt").read_text(encoding="utf-8").strip() if Path(cur).with_suffix(".txt").exists() else ""
d2 = _parse_caption(cap2)
crop_src_v = _bust([(cur, Path(cur).name)])[0][0]
else:
cap2, d2, crop_src_v = "", _parse_caption(""), None
return sel_disp, cand_disp, sp2, cp2, crop_src_v, cur, cap2, \
f'<span style="color:#16a34a">{status_note}</span>', \
d2["h"], d2["v"], d2["expr"], d2["comp"], \
*_stats_fields(base_dir)
def crop_preview(img_path, x0, y0, x1, y1):
if not img_path or not Path(img_path).exists():
return None
from PIL import Image as _I
im = _I.open(img_path)
w, h = im.size
bx0 = int(w * x0 / 100); by0 = int(h * y0 / 100)
bx1 = int(w * x1 / 100); by1 = int(h * y1 / 100)
if bx1 - bx0 < 20 or by1 - by0 < 20:
return str(img_path)
crop = im.crop((bx0, by0, bx1, by1))
tmp = Path(tempfile.gettempdir()) / f"swap_crop_preview_{int(time.time())}.jpg"
crop.save(tmp, quality=95)
return str(tmp)
def crop_apply(base_dir, img_path, x0, y0, x1, y1, sp, cp, fh, fv, fexpr, fcomp):
# img_path 来自 crop_file State(真实路径,非显示副本)
if not img_path or not Path(img_path).exists():
return gr.update(), gr.update(), gr.update(), gr.update(), \
'<span style="color:#dc2626">❌ 请先点选一张图片</span>', gr.update(), gr.update(), gr.update(), gr.update(), gr.update()
p = Path(img_path)
from PIL import Image as _I
im = _I.open(p)
w, h = im.size
bx0 = int(w * x0 / 100); by0 = int(h * y0 / 100)
bx1 = int(w * x1 / 100); by1 = int(h * y1 / 100)
if bx1 - bx0 < 20 or by1 - by0 < 20:
return gr.update(), gr.update(), gr.update(), gr.update(), \
'<span style="color:#dc2626">❌ 裁剪区域太小(至少 20px</span>', gr.update(), gr.update(), gr.update(), gr.update(), gr.update()
crop = im.crop((bx0, by0, bx1, by1))
bak = p.with_name(f"{p.stem}.bak_{int(time.time())}.jpg")
if not bak.exists():
p.rename(bak)
crop.save(p, quality=95)
# 刷新画廊 + 预览(带时间戳防 Gradio 缓存);已选按当前筛选条件重过滤
base_dir = clean_path(base_dir)
selected, cands = _swap_collect(base_dir)
sel_disp, cand_disp, sp2, cp2 = _reapply_filter(base_dir, fh, fv, fexpr, fcomp, cands)
fresh = f"{p}?t={int(time.time())}" if False else str(p) # 路径本身变化不大,直接返回
msg = (f'<span style="color:#16a34a">✅ 已裁剪 {p.name}{crop.size[0]}×{crop.size[1]} '
f'(原图备份 {bak.name}</span>')
return sel_disp, cand_disp, sp2, cp2, msg, *_stats_fields(base_dir)
def do_remove(base_dir, si, sp, fh, fv, fexpr, fcomp):
"""移除已选素材:删除裁剪版,恢复原图到未选中(换出_ 文件恢复原名,final 源复制进来)"""
if si < 0 or si >= len(sp):
return gr.update(), gr.update(), sp, sp, \
'<span style="color:#dc2626">❌ 请先点选一张已选素材</span>', gr.update(), gr.update(), gr.update(), gr.update(), gr.update()
base_dir = clean_path(base_dir)
p = Path(sp[si])
if not p.exists():
return gr.update(), gr.update(), sp, sp, '<span style="color:#dc2626">❌ 文件不存在</span>', gr.update(), gr.update(), gr.update(), gr.update(), gr.update()
removed_name = p.name
# 1. 找源图(在删除前,用 ORB)
src = _find_source_by_orb(str(p), base_dir)
# 2. 删除裁剪版 + caption
p.unlink()
txt = p.with_suffix(".txt")
if txt.exists():
txt.unlink()
# 3. 恢复原图到未选中
unused = Path(base_dir) / "未选中"
unused.mkdir(exist_ok=True)
restored = None
if src is not None:
if src.parent.name == "final":
# 源在 final:复制到未选中(原名)
dst = unused / src.name
if not dst.exists():
import shutil as _sh
_sh.copy2(src, dst)
restored = src.name
else:
# 源在未选中:若是 换出_ 前缀,恢复原名
if src.name.startswith("换出_"):
orig_name = src.name.split("_", 2)[-1] if src.name.count("_") >= 2 else src.name
dst = unused / orig_name
if not dst.exists():
src.rename(dst)
restored = orig_name
else:
src.unlink()
restored = orig_name
else:
restored = src.name # 已在未选中,无需动
else:
restored = "(源未找到,已移除裁剪版)"
# 4. 刷新画廊(已选按当前筛选条件重过滤)
selected, cands = _swap_collect(base_dir)
sel_disp, cand_disp, sp2, cp2 = _reapply_filter(base_dir, fh, fv, fexpr, fcomp, cands)
msg = (f'<span style="color:#16a34a">✅ 已移除 {removed_name};原图恢复: {restored}</span>')
return sel_disp, cand_disp, sp2, cp2, msg, *_stats_fields(base_dir)
def do_swap(base_dir, si, ci, backend, sp, cp, fh, fv, fexpr, fcomp):
if si < 0 or ci < 0 or si >= len(sp) or ci >= len(cp):
return gr.update(), gr.update(), sp, cp, gr.update(), gr.update(), gr.update(), gr.update(), gr.update(), \
'<span style="color:#dc2626">❌ 请先在两边各点选一张图</span>'
base_dir = clean_path(base_dir)
sel_path = Path(sp[si])
cand_path = Path(cp[ci])
kind = sel_path.parent.name # 特写/半身/全身
try:
# 1-3. 检测 + 按类型裁剪 + VLM caption(失败则不动任何文件)
crop, caption, angle_en = _prep_candidate(cand_path, kind, backend)
# 4. 换出:旧文件移回 未选中(带时间戳防重名)
import time
unused = Path(base_dir) / "未选中"
unused.mkdir(exist_ok=True)
out_name = f"换出_{int(time.time())}_{sel_path.name}"
sel_path.rename(unused / out_name)
old_txt = sel_path.with_suffix(".txt")
if old_txt.exists():
old_txt.unlink()
# 5. 换入:写新图 + caption(沿用原文件名,保持编号连续)
crop.save(sel_path, quality=95)
sel_path.with_suffix(".txt").write_text(caption, encoding="utf-8")
# 6. 候选原图从未选中移除(已换入)
cand_path.unlink()
except Exception as e:
return gr.update(), gr.update(), sp, cp, gr.update(), gr.update(), gr.update(), gr.update(), gr.update(), \
f'<span style="color:#dc2626">❌ 交换失败: {e}</span>'
selected, cands = _swap_collect(base_dir)
sel_disp, cand_disp, sp2, cp2 = _reapply_filter(base_dir, fh, fv, fexpr, fcomp, cands)
msg = (f'<span style="color:#16a34a">✅ 已交换:{cand_path.name}{kind}/{sel_path.name} '
f'(角度 {angle_en},caption 已重新生成,可到 ③ 打标微调)</span>')
return sel_disp, cand_disp, sp2, cp2, msg, *_stats_fields(base_dir)
def do_add(base_dir, ci, kind, backend, cp, fh, fv, fexpr, fcomp):
"""添加(非交换):候选图按所选类型裁剪+打标,追加为新编号素材;候选保留可继续添加到其他类型"""
if ci < 0 or ci >= len(cp):
return gr.update(), gr.update(), gr.update(), gr.update(), \
'<span style="color:#dc2626">❌ 请先在右边点选一张候选图</span>', gr.update(), gr.update(), gr.update(), gr.update(), gr.update()
base_dir = clean_path(base_dir)
cand_path = Path(cp[ci])
try:
crop, caption, angle_en = _prep_candidate(cand_path, kind, backend)
prefix = {"特写": "face", "半身": "half", "全身": "full"}[kind]
kind_dir = Path(base_dir) / kind
kind_dir.mkdir(exist_ok=True)
nums = []
for p in kind_dir.glob(f"{prefix}_*.jpg"):
parts = p.stem.split("_")
if len(parts) == 2 and parts[1].isdigit():
nums.append(int(parts[1]))
n = max(nums) + 1 if nums else 1
dst = kind_dir / f"{prefix}_{n:03d}.jpg"
crop.save(dst, quality=95)
dst.with_suffix(".txt").write_text(caption, encoding="utf-8")
except Exception as e:
return gr.update(), gr.update(), gr.update(), gr.update(), \
f'<span style="color:#dc2626">❌ 添加失败: {e}</span>', gr.update(), gr.update(), gr.update(), gr.update(), gr.update()
selected, cands = _swap_collect(base_dir)
sel_disp, cand_disp, sp2, cp2 = _reapply_filter(base_dir, fh, fv, fexpr, fcomp, cands)
msg = (f'<span style="color:#16a34a">✅ 已添加:{cand_path.name}{kind}/{dst.name} '
f'(角度 {angle_en},caption 已生成)。候选仍保留在未选中,可继续添加到其他类型</span>')
return sel_disp, cand_disp, sp2, cp2, msg, *_stats_fields(base_dir)
load_btn.click(load, inputs=base_tb,
outputs=[sel_gal, cand_gal, sel_paths, cand_paths, status_html, h_dist, v_dist, expr_dist, comp_dist, advice_html])
# 页面加载/刷新时自动从 gui_state 同步 auto 输出目录(①跑完自动带过来)
demo.load(lambda: state_val("auto_out", ""), outputs=base_tb)
# 页面加载时也刷新统计(无需先点选图)
# 筛选:4 个下拉任一变化 → 过滤已选(输出 画廊+操作列表+状态);清除按钮恢复全部
for _fd in (f_h, f_v, f_expr, f_comp):
_fd.change(_filter_selected, inputs=[base_tb, f_h, f_v, f_expr, f_comp],
outputs=[sel_gal, sel_paths, filter_info])
clear_filter.click(_clear_filter, inputs=[base_tb], outputs=[sel_gal, sel_paths, filter_info])
sel_gal.select(on_sel, inputs=[base_tb, sel_paths], outputs=[sel_idx, crop_src, crop_file, cap_tb, crop_status, h_dd, v_dd, expr_dd, comp_dd])
cand_gal.select(on_cand, inputs=[base_tb, cand_paths], outputs=[cand_idx, crop_src, crop_file, cap_tb, crop_status, h_dd, v_dd, expr_dd, comp_dd])
for _dd in (h_dd, v_dd, expr_dd, comp_dd):
_dd.change(cap_rebuild, inputs=[h_dd, v_dd, expr_dd, comp_dd, cap_tb], outputs=cap_tb)
cap_regen.click(cap_regen_fn, inputs=[crop_file, vlm_backend], outputs=[cap_tb, cap_status])
cap_save.click(cap_save_fn, inputs=[base_tb, crop_file, cap_tb, comp_dd, sel_paths, cand_paths, f_h, f_v, f_expr, f_comp],
outputs=[sel_gal, cand_gal, sel_paths, cand_paths, crop_src, crop_file, cap_tb, cap_status,
h_dd, v_dd, expr_dd, comp_dd, h_dist, v_dist, expr_dist, comp_dist, advice_html])
for s in (cr_x0, cr_y0, cr_x1, cr_y1):
s.change(crop_preview, inputs=[crop_file, cr_x0, cr_y0, cr_x1, cr_y1], outputs=crop_prev)
crop_btn.click(crop_apply,
inputs=[base_tb, crop_file, cr_x0, cr_y0, cr_x1, cr_y1, sel_paths, cand_paths, f_h, f_v, f_expr, f_comp],
outputs=[sel_gal, cand_gal, sel_paths, cand_paths, crop_status, h_dist, v_dist, expr_dist, comp_dist, advice_html])
swap_btn.click(do_swap, inputs=[base_tb, sel_idx, cand_idx, vlm_backend, sel_paths, cand_paths, f_h, f_v, f_expr, f_comp],
outputs=[sel_gal, cand_gal, sel_paths, cand_paths, status_html, h_dist, v_dist, expr_dist, comp_dist, advice_html])
add_btn.click(do_add, inputs=[base_tb, cand_idx, kind_dd, vlm_backend, cand_paths, f_h, f_v, f_expr, f_comp],
outputs=[sel_gal, cand_gal, sel_paths, cand_paths, status_html, h_dist, v_dist, expr_dist, comp_dist, advice_html])
cancel_btn.click(do_remove, inputs=[base_tb, sel_idx, sel_paths, f_h, f_v, f_expr, f_comp],
outputs=[sel_gal, cand_gal, sel_paths, cand_paths, status_html, h_dist, v_dist, expr_dist, comp_dist, advice_html])
prepare_btn.click(_prepare_trainset, inputs=[base_tb], outputs=status_html)
# ================= Tab 5: 融合工具 =================
def build_merge_tab():
with gr.Tab("🔗 融合(可选)"):
gr.Markdown("**merge_fixed**:原图未修改区 + ComfyUI 修改图修改区 无缝合并(解决 inpainting 全图清晰度损失)")
with gr.Row():
m_orig = gr.Textbox(label="原图目录", value=state_val("merge_orig", ""), scale=2)
m_mod = gr.Textbox(label="修改图目录(IMG_0463*.png", value=state_val("merge_mod", ""), scale=2)
m_out = gr.Textbox(label="输出目录", value=state_val("merge_out", ""), scale=2)
with gr.Row():
m_threshold = gr.Slider(4, 20, value=float(state_clamp("merge_threshold", 8, 4, 20)), step=0.5, label="修改幅度阈值(人物移除diff高、降质diff低;误判时调)")
m_softness = gr.Slider(0.5, 5, value=float(state_clamp("merge_softness", 1.5, 0.5, 5)), step=0.5, label="过渡带宽度(越小修改区越彻底,防残影)")
with gr.Row():
m_protect = gr.Checkbox(value=bool(state_val("merge_protect", True)), label="保护目标人物面部(强制保留原图像素,零改变)")
m_btn = gr.Button("🔗 开始融合", variant="primary")
m_stats = gr.HTML("<span style='color:#666'>就绪</span>")
m_log = gr.Textbox(label="日志", lines=15, interactive=False)
def do_merge(o_dir, m_dir, out_dir, threshold, softness, protect):
o_dir, m_dir, out_dir = clean_path(o_dir), clean_path(m_dir), clean_path(out_dir)
save_state("merge_orig", o_dir)
save_state("merge_mod", m_dir)
save_state("merge_out", out_dir)
save_state("merge_threshold", threshold)
save_state("merge_softness", softness)
save_state("merge_protect", protect)
if not Path(o_dir).exists() or not Path(m_dir).exists():
return '<span style="color:#dc2626">❌ 原图或修改图目录不存在</span>', ""
log_txt, _ = run_with_log(mf.process, o_dir, m_dir, out_dir, threshold, softness, protect)
return f'<span style="color:#16a34a">✅ 融合完成</span>', log_txt
m_btn.click(do_merge, [m_orig, m_mod, m_out, m_threshold, m_softness, m_protect], [m_stats, m_log])
# ================= Tab 6: 训练 =================
TRAIN_SCRIPT = PROJ_DIR / "config" / "训练脚本.py"
STATUS_FILE = PROJ_DIR / "output" / "train_status.json"
TRAIN_CONFIG = PROJ_DIR / "config" / "train_config.json"
MODELS_DIR = Path(r"D:\AI\sd\models\qwen-edit-2511")
_train_proc = {"proc": None}
def _train_cfg():
"""读 train_config.json(训练目录的唯一事实来源),失败返回 {}"""
if TRAIN_CONFIG.exists():
try:
return json.loads(TRAIN_CONFIG.read_text(encoding="utf-8-sig"))
except Exception:
return {}
return {}
def build_train_tab():
with gr.Tab("④ 训练"):
gr.Markdown(
"""**④ 训练**:启动脸部 LoRA 训练(后台隔夜跑)。**三个目录说明:**
| 目录 | 用途 | 里面放什么 |
|---|---|---|
| **训练集目录** | 训练的素材(必须已打标) | `img_001.jpg` + `img_001.txt`(③ 打标一键整理自动写入,此处自动同步) |
| **输出目录** | 训练产物 | `myface_lora-*.safetensors`LoRA+ `sample/`(每 10 epoch 自动样本图) |
| **模型目录** | 训练底模(固定,勿改) | DiT 分片 + VAE + 文本编码器(`D:\\AI\\sd\\models\\qwen-edit-2511` |
**点击顺序**:③ 打标 💾保存 → 📦一键整理 → 回到这里确认目录已同步 → 🔄 重新预缓存(整理后必点,否则训旧图)→ ▶️ 启动。"""
)
_cfg = _train_cfg()
with gr.Row():
td = gr.Textbox(label="训练集目录", value=_cfg.get("train_dataset", state_val("train_ds", str(TRAIN_DATASET))), scale=2)
od = gr.Textbox(label="输出目录(LoRA + 样本图)", value=_cfg.get("output_dir", state_val("train_out", str(CHECKPOINT_DIR))), scale=2)
gr.Markdown(f"**模型目录(固定)**: `{MODELS_DIR}`(不要改)")
# ── 打标关键字统计(整理训练集之后,用于挑选测试 prompt 关键词)──
with gr.Accordion("📊 打标关键字统计(选词做测试 prompt)", open=True):
kw_stats = gr.HTML('<span style="color:#9ca3af">输入训练集目录后自动统计(也可点按钮刷新)</span>')
kw_refresh = gr.Button("🔄 刷新关键字统计", size="sm")
with gr.Row():
cache_btn = gr.Button("🔄 重新预缓存(整理后必点)")
start_btn = gr.Button("▶️ 启动训练(后台隔夜跑)", variant="primary")
stop_btn = gr.Button("⏹️ 停止训练")
refresh_btn = gr.Button("🔄 刷新状态")
t_status = gr.HTML("<span style='color:#666'>未启动</span>")
t_log = gr.Textbox(label="日志(训练/缓存输出尾部)", lines=8, interactive=False)
t_gallery = gr.Gallery(label="训练样本图(每 10 epoch 自动生成)", columns=4, height=350)
def start_train(train_ds, out_dir):
if _train_proc["proc"] and _train_proc["proc"].poll() is None:
return '<span style="color:#d97706">⚠️ 训练已在运行</span>', None
train_ds, out_dir = clean_path(train_ds), clean_path(out_dir)
save_state("train_ds", train_ds)
save_state("train_out", out_dir)
train_ds, out_dir = Path(train_ds), Path(out_dir)
if not list(train_ds.glob("img_*.jpg")):
return '<span style="color:#dc2626">❌ 训练集为空(没有 img_*.jpg),先去 ③ 打标整理</span>', None
if not MODELS_DIR.exists():
return '<span style="color:#dc2626">❌ 模型目录不存在(底模未下载?)</span>', None
# 写训练配置(动态目录)
TRAIN_CONFIG.write_text(
json.dumps({"train_dataset": str(train_ds), "output_dir": str(out_dir)},
ensure_ascii=False, indent=1), encoding="utf-8")
out_dir.mkdir(parents=True, exist_ok=True)
logf = open(PROJ_DIR / "output" / "train.log", "w", encoding="utf-8")
p = subprocess.Popen(
[sys.executable, str(TRAIN_SCRIPT)],
stdout=logf, stderr=subprocess.STDOUT, cwd=str(PROJ_DIR), creationflags=subprocess.CREATE_NO_WINDOW,
)
_train_proc["proc"] = p
return ('<span style="color:#16a34a">✅ 训练已启动(后台运行)。训练集: '
f'{train_ds.name},输出: {out_dir.name}</span>'), None
def stop_train():
p = _train_proc.get("proc")
if p and p.poll() is None:
p.terminate()
return '<span style="color:#d97706">⏹️ 已发送停止请求</span>'
return '<span style="color:#666">没有运行中的训练</span>'
def refresh():
samples = sorted(glob.glob(str(SAMPLE_DIR / "*.png")))[-16:] if SAMPLE_DIR.exists() else []
log_txt = ""
logf = PROJ_DIR / "output" / "train.log"
if logf.exists():
lines = logf.read_text(encoding="utf-8", errors="replace").strip().splitlines()
log_txt = "\n".join(lines[-8:])
p = _train_proc.get("proc")
if p and p.poll() is None:
status = '<span style="color:#16a34a">🟢 训练运行中(样本图 ' + str(len(samples)) + ' 张)</span>'
else:
status = '<span style="color:#666">训练未在运行</span>'
return status, log_txt, samples
def recache(train_ds):
"""重新预缓存(VAE latent + TextEncoder)。素材整理后必须重缓存,否则训练用旧缓存。"""
train_ds = clean_path(train_ds)
save_state("train_ds", train_ds)
# 先把目录写进 train_configrun_cache 读 dataset_active.toml,由训练脚本从 config 生成)
cfg = {}
if TRAIN_CONFIG.exists():
try:
cfg = json.loads(TRAIN_CONFIG.read_text(encoding="utf-8-sig"))
except Exception:
pass
cfg["train_dataset"] = train_ds
TRAIN_CONFIG.write_text(json.dumps(cfg, ensure_ascii=False, indent=1), encoding="utf-8")
# dataset_active.toml 由训练脚本生成;这里手动同步一份保证缓存脚本读的是新目录
toml_path = PROJ_DIR / "config" / "dataset_active.toml"
toml_path.write_text(
'# 由 ⑤ 训练 Tab 重新预缓存生成\n[general]\nresolution = 1024\ncaption_extension = ".txt"\n'
'batch_size = 1\nenable_bucket = true\nbucket_no_upscale = false\n\n'
f'[[datasets]]\nimage_directory = "{train_ds.replace(chr(92), "/")}"\nnum_repeats = 1\n',
encoding="utf-8")
cache_script = PROJ_DIR / "output" / "run_cache.py"
r = subprocess.run([sys.executable, str(cache_script)],
capture_output=True, text=True, encoding="utf-8", errors="replace")
tail = "\n".join(((r.stdout or "") + (r.stderr or "")).splitlines()[-6:])
if r.returncode == 0:
return f'<span style="color:#16a34a">✅ 预缓存完成({train_ds}),可以启动训练</span>', tail
return f'<span style="color:#dc2626">❌ 预缓存失败,见日志</span>', tail
cache_btn.click(recache, [td], [t_status, t_log])
start_btn.click(start_train, [td, od], [t_status, t_gallery])
stop_btn.click(stop_train, outputs=[t_status, t_gallery])
refresh_btn.click(refresh, outputs=[t_status, t_log, t_gallery])
# 关键字统计:目录变化或点按钮时刷新
td.change(lambda d: _keyword_stats(d), inputs=td, outputs=kw_stats)
kw_refresh.click(lambda d: _keyword_stats(d), inputs=td, outputs=kw_stats)
# 自动刷新:每 15 秒更新状态/日志/样本图,不用手动点
timer = gr.Timer(15)
timer.tick(refresh, outputs=[t_status, t_log, t_gallery])
def build_demo():
with gr.Blocks(title="脸部 LoRA 控制台", theme=gr.themes.Soft(),
css="""/* 字段行:label 按内容宽,下拉框固定窄,分布小字占剩余空间 */
.field-row { flex-wrap: wrap !important; align-items: center; }
.field-row .x-label { flex: 0 0 auto !important; width: auto !important; min-width: 0 !important; margin-right: 8px; }
.field-row .fld-dd { flex: 0 0 auto !important; width: 110px !important; min-width: 0 !important; max-width: 110px !important; }
.field-row .gr-html { flex: 1 1 auto !important; min-width: 0 !important; overflow: visible !important; }
.field-stats { white-space: nowrap; }
.field-warn { font-size: 11px; color: #ff6b6b; font-weight: 600; margin-left: 8px; white-space: nowrap; }
/* 筛选栏:一行紧凑排列(label 窄 + 下拉窄 + 计数 + 清除按钮) */
.filter-row { flex-wrap: wrap !important; align-items: center; gap: 4px !important; }
.filter-row .x-label { flex: 0 0 auto !important; margin-right: 6px !important; font-size: 13px !important; }
.filter-row .fld-dd { flex: 0 0 auto !important; width: 96px !important; min-width: 0 !important; max-width: 96px !important; }
.filter-row .filter-cnt { flex: 1 1 auto !important; min-width: 0 !important; white-space: nowrap; }
/* 全屏预览(slide show)里的操作按钮:fixed 定位在视口右上角空白处(不遮挡画廊/缩略图/全屏小方框) */
.slide-bar {
position: fixed !important;
top: 70px !important;
right: 24px !important;
z-index: 9999 !important;
display: none;
align-items: center !important;
gap: 8px !important;
background: rgba(20,20,20,0.85) !important;
padding: 8px 12px !important;
border-radius: 10px !important;
box-shadow: 0 2px 10px rgba(0,0,0,0.4) !important;
}
.slide-bar.show { display: inline-flex !important; }
/* ③打标 拖拽替换的内部通道:保留 DOM 供 JS 使用,但视觉隐藏(编辑面板可见) */
#kw_old_tb, #kw_new_tb, #kw_apply_btn { display: none !important; }
.slide-bar select {
padding: 4px 6px !important;
border-radius: 6px !important;
border: none !important;
font-size: 12px !important;
background: #374151 !important;
color: #f3f4f6 !important;
}
.slide-bar button {
padding: 5px 12px !important;
border-radius: 6px !important;
border: none !important;
font-size: 12px !important;
font-weight: 600 !important;
cursor: pointer !important;
}
.slide-bar .sel-cancel { background: #dc2626 !important; color: #fff !important; }
.slide-bar .cand-add { background: #16a34a !important; color: #fff !important; }
""",
js="""// 刷新/重载时恢复滚动位置(Gradio SPA 默认拉到顶)。
// ⚠️ Gradio 会把这个 js 当成「函数」执行(new AsyncFunction + (${js})()),所以必须定义函数而非 IIFE!
() => {
const KEY = 'caption_gui_scrollY';
const save = () => { try { sessionStorage.setItem(KEY, String(window.scrollY)); } catch(e){} };
const restore = () => {
try {
const y = parseInt(sessionStorage.getItem(KEY) || '0', 10);
if (y > 0) {
window.scrollTo(0, y);
setTimeout(() => window.scrollTo(0, y), 300); // Gradio 渲染完成后补一次
setTimeout(() => window.scrollTo(0, y), 1000);
}
} catch(e){}
};
window.addEventListener('beforeunload', save);
window.addEventListener('load', restore);
window.addEventListener('scroll', () => { try { sessionStorage.setItem(KEY, String(window.scrollY)); } catch(e){} }, {passive:true});
// ── 日志框 tail-follow:auto 日志更新时自动滚到底(用户上翻历史则暂停跟随)──
const LOG_ID = 'auto_log';
let logPinned = true;
const findLog = () => document.getElementById(LOG_ID)?.querySelector('textarea');
const isNearBottom = (el) => el && (el.scrollHeight - el.scrollTop - el.clientHeight) < 60;
// 用户滚动时更新 pinned 状态:滚到底部 → 恢复跟随;上翻 → 暂停
const attachScroll = () => {
const el = findLog();
if (!el || el.__logTailAttached) return;
el.__logTailAttached = true;
el.addEventListener('scroll', () => { logPinned = isNearBottom(el); }, { passive: true });
};
// 内容变化(Gradio 更新 value)后:若在底部则强制滚到底
const tailFollow = () => {
const el = findLog();
if (!el) return;
attachScroll();
if (logPinned && !isNearBottom(el)) {
el.scrollTop = el.scrollHeight;
}
};
// Gradio 更新 value 是设置 textarea.value + dispatch,用 interval 兜底(组件懒渲染时也能抓到)
setInterval(() => { attachScroll(); tailFollow(); }, 500);
// 同时监听 textarea 区域变化,立即响应
const mo = new MutationObserver(tailFollow);
const startObserve = () => {
const el = findLog();
if (el && !el.__logTailObserved) {
el.__logTailObserved = true;
mo.observe(el, { childList: true, subtree: true, characterData: true });
}
};
setInterval(startObserve, 500);
// ── 全屏预览(slide show)内嵌操作栏:已选预览→取消已选;候选预览→添加选中候选(带类型下拉)──
// 按钮 fixed 定位在视口右上角(空白处,不挡图),append 到 body;显隐跟随对应画廊的预览状态
const injectSlideButtons = () => {
const mkBar = (cls) => {
let bar = document.querySelector(cls);
if (!bar) {
bar = document.createElement('div');
bar.className = 'slide-bar ' + cls.replace('.', '');
document.body.appendChild(bar);
}
return bar;
};
// 已选画廊:「取消已选」按钮
const selBar = mkBar('.sel-bar');
if (!selBar.querySelector('button')) {
const b = document.createElement('button');
b.className = 'sel-cancel';
b.textContent = '↩️ 取消已选(移入候选)';
b.addEventListener('click', () => {
const real = document.getElementById('btn_cancel_sel');
if (real) real.click();
});
selBar.appendChild(b);
}
// 候选画廊:「添加选中候选」+ 类型下拉
const candBar = mkBar('.cand-bar');
if (!candBar.querySelector('select')) {
const sel = document.createElement('select');
['特写', '半身', '全身'].forEach(k => {
const o = document.createElement('option');
o.value = k; o.textContent = k;
sel.appendChild(o);
});
const b = document.createElement('button');
b.className = 'cand-add';
b.textContent = ' 添加选中候选';
b.addEventListener('click', () => {
// 同步类型下拉到页面真实 kind_ddelem_id=cand_kind_dd),再触发添加按钮
const dd = document.querySelector('#cand_kind_dd select, #cand_kind_dd input');
if (dd) {
const setter = Object.getOwnPropertyDescriptor(
dd.tagName === 'SELECT' ? window.HTMLSelectElement.prototype : window.HTMLInputElement.prototype,
'value').set;
setter.call(dd, sel.value);
dd.dispatchEvent(new Event('change', { bubbles: true }));
}
const real = document.getElementById('btn_add_cand');
if (real) real.click();
});
candBar.appendChild(sel);
candBar.appendChild(b);
}
// 显隐:对应画廊预览打开(.preview 存在)→ 显示该操作栏。
// 两个画廊的 preview 可同时开(Gradio 两个独立组件),操作栏固定在同一位置会重叠——
// 用「最后点击的画廊」互斥(document click 记录,见下),比时间戳可靠。
const selG = document.getElementById('sel_gallery');
const candG = document.getElementById('cand_gallery');
const selOpen = !!(selG && selG.querySelector('.preview'));
const candOpen = !!(candG && candG.querySelector('.preview'));
if (selOpen && candOpen) {
// 两个都开:显示最后被点击的那个画廊的操作栏
const showSel = __lastActive === 'sel';
selBar.classList.toggle('show', showSel);
candBar.classList.toggle('show', !showSel);
} else {
selBar.classList.toggle('show', selOpen);
candBar.classList.toggle('show', candOpen);
}
};
// 记录用户最后点击的画廊(document 捕获阶段监听,最可靠反映用户意图)
let __lastActive = null;
document.addEventListener('click', (e) => {
const t = e.target;
if (!t || !t.closest) return;
if (t.closest('#sel_gallery')) __lastActive = 'sel';
else if (t.closest('#cand_gallery')) __lastActive = 'cand';
}, true);
setInterval(injectSlideButtons, 400);
const mob = new MutationObserver(injectSlideButtons);
const startSlideObserve = () => {
for (const id of ['sel_gallery', 'cand_gallery']) {
const el = document.getElementById(id);
if (el && !el.__slideObserved) {
el.__slideObserved = true;
mob.observe(el, { childList: true, subtree: true });
}
}
};
setInterval(startSlideObserve, 400);
// ===== 打标统计:拖拽/双击 批量替换(只作用于 #label_kw_stats 容器) =====
function kwInit() {
const statsEl = document.getElementById('label_kw_stats');
if (!statsEl || statsEl.dataset.kwBound) return;
statsEl.dataset.kwBound = '1';
// 拖拽开始:记录源词
statsEl.addEventListener('dragstart', (e) => {
const el = e.target.closest('.kw-item');
if (!el) return;
e.dataTransfer.setData('text/plain', el.dataset.word);
e.dataTransfer.effectAllowed = 'move';
});
// 拖拽悬停:允许放下
statsEl.addEventListener('dragover', (e) => {
if (!e.target.closest('.kw-item')) return;
e.preventDefault();
e.dataTransfer.dropEffect = 'move';
});
// 放下:把「目标词」替换成「被拖的词」(例:拖"黑蕾丝抹胸裙"到"黑白抹胸"上 = 所有 caption 中"黑白抹胸"→"黑蕾丝抹胸裙"
statsEl.addEventListener('drop', (e) => {
const tgt = e.target.closest('.kw-item');
if (!tgt) return;
e.preventDefault();
const src = e.dataTransfer.getData('text/plain');
if (!src) return;
kwSetReplace(tgt.dataset.word, src);
});
// 双击:填入编辑面板(inline 编辑,比 prompt 舒服),用户改完点「应用修改」
statsEl.addEventListener('dblclick', (e) => {
const el = e.target.closest('.kw-item');
if (!el) return;
e.preventDefault();
const old = el.dataset.word;
const oldEl = document.getElementById('kw_old_tb');
const editEl = document.getElementById('kw_edit_tb');
if (!oldEl || !editEl) return;
const setVal = (root, v) => {
const ta = root.querySelector('textarea') || root.querySelector('input');
if (ta) { ta.value = v; ta.dispatchEvent(new Event('input', {bubbles: true})); }
};
setVal(oldEl, old);
setVal(editEl, old);
const ta = editEl.querySelector('textarea');
if (ta) ta.focus();
});
}
function kwSetReplace(oldWord, newWord) {
if (!oldWord || !newWord || oldWord === newWord) return;
const oldEl = document.getElementById('kw_old_tb');
const newEl = document.getElementById('kw_new_tb');
const btn = document.getElementById('kw_apply_btn');
if (!oldEl || !newEl || !btn) return;
const setVal = (root, v) => {
const ta = root.querySelector('textarea') || root.querySelector('input');
if (ta) { ta.value = v; ta.dispatchEvent(new Event('input', {bubbles: true})); }
};
setVal(oldEl, oldWord);
setVal(newEl, newWord);
const bt = btn.querySelector('button') || btn;
bt.click();
}
setInterval(kwInit, 800);
// 缩略图点击 → 滚动到下方对应图片卡片 + 金色高亮闪烁 + 聚焦 Caption 框
// 必须挂到 windowGradio Blocks js 作用域不暴露内部函数,onclick 需全局函数
window.scrollToKwCard = function(idx) {
const card = document.getElementById('img_card_' + idx);
if (!card) return;
card.scrollIntoView({ behavior: 'smooth', block: 'center' });
// 高亮闪烁(2 次)
let n = 0;
const flash = setInterval(() => {
card.style.boxShadow = n % 2 === 0 ? '0 0 0 4px #fbbf24' : '';
n++;
if (n >= 4) { clearInterval(flash); card.style.boxShadow = ''; }
}, 300);
// 聚焦 Caption 文本框
const tb = document.getElementById('cap_tb_' + idx);
const ta = tb ? tb.querySelector('textarea') : null;
if (ta) {
setTimeout(() => { ta.focus(); }, 500);
}
};
return true;
}""") as demo:
gr.Markdown(f"# 🎛️ 脸部 LoRA 一站式控制台 <span style='font-size:14px;color:#888'>v{VERSION}</span>")
gr.Markdown(
"**主线工作流(按序号一步步来,参数自动传递到下一步)**:\n"
"① 素材自动处理(final/ → 分类裁剪打标筛选) → "
"② 候选换图(调整素材 + 审改 caption + 整理训练集) → "
"④ 训练(重新预缓存 → 启动)\n\n"
"🔗 融合 / 🔍 审核 / 🎯 选图 是可选工具,一般不需要用。"
)
with gr.Row():
vlm_backend = gr.CheckboxGroup(
["ollama", "omlx", "omlx-32b", "sensenova"],
value=["omlx-32b"],
label="VLM 模型组合(勾选参与判定的模型:多选=投票更准,单选=只用一个。"
"ollama=本地8B / omlx=小果30B / omlx-32b=小果qwen2.5-VL-32B-Q8 / sensenova=云端)",
info="打标/描述自动用优先级最高的勾选模型(omlx-32b > omlx > ollama > sensenova);"
"omlx 与 omlx-32b 同属小果一台服务器,同时勾选时自动只保留 omlx-32b 一路")
with gr.Tabs():
build_auto_tab(vlm_backend)
build_swap_tab(demo, vlm_backend)
# ③ 打标:恢复独立 tab,专用于 LLM 批量整理 tag(拆短句+合并同义词)
build_label_tab(demo, vlm_backend)
build_train_tab()
build_merge_tab()
build_check_tab()
build_pick_tab()
return demo
if __name__ == "__main__":
demo = build_demo()
# generatoryield 流式输出)必须启用 queue 才能实时逐帧刷新,否则点击按钮无响应/等全部完成才显示
demo.queue()
print(f"\n[OK] 控制台启动 v{VERSION}: http://127.0.0.1:7860")
demo.launch(server_name="127.0.0.1", server_port=7860, inbrowser=False, show_error=True,
# Gradio 安全限制:返回给 Gallery/Image 的文件路径必须在 allowed_paths 内,
# 否则报 InvalidPathErrorK:\ 素材盘 + 项目 photos/output 都放行)
allowed_paths=[r"K:\AI\training", str(PHOTOS), str(PROJ_DIR / "output")])