bjooo's picture
Upload folder using huggingface_hub
9c98083 verified
Raw History Blame Contribute Delete
18.5 kB
"""Local web UI for the Wan 2.2 LoRA dataset builder + trainer.
Thin Gradio layer over the already-verified CLI scripts (build.py,
train/train.py, train/download_checkpoints.py, train/setup_env.ps1) -- this
file does not reimplement any of their logic, it only shells out to them and
streams stdout live. Heavy ML models (captioners, torch) never load inside
this long-lived server process; each run is its own subprocess so a crash or
this environment's known torch-teardown quirk can't take the UI down with it.
Run:
python_embeded\\python.exe dataset_builder\\webui.py [--port 7860]
"""
import argparse
import os
import subprocess
import sys
import gradio as gr
HERE = os.path.dirname(os.path.abspath(__file__))
TRAIN_DIR = os.path.join(HERE, "train")
VENV_PYTHON = os.path.join(TRAIN_DIR, "venv", "Scripts", "python.exe")
PYTHON = sys.executable # python_embeded's own interpreter (has build.py's deps)
sys.path.insert(0, TRAIN_DIR)
from download_checkpoints import FILES, COMMON_FILES # noqa: E402
# ComfyUI/custom_nodes/dolphin_nodes/dataset_builder/webui.py -> ComfyUI/models
MODELS_ROOT = os.path.normpath(os.path.join(HERE, "..", "..", "..", "models"))
# Any python_embeded process that imports torch on this machine can exit with
# this NTSTATUS during CUDA teardown *after* finishing its real work (confirmed
# for both dataset_builder/train/train.py and build.py itself, e.g. via
# caption.py's torch import) -- see dataset_builder/train/README.md's bug list.
# Reported as either the signed or unsigned 32-bit form depending on context.
BENIGN_TORCH_TEARDOWN_NTSTATUS = 0xC0000409 # STATUS_STACK_BUFFER_OVERRUN
EMBEDED_DIR = os.path.dirname(PYTHON)
def child_env():
"""Environment for subprocesses, scrubbed of this server's own Python.
python_embeded is Python 3.13 and ships python313.dll in its own directory.
A child running the training venv's Python 3.10 will happily load that
3.13 DLL if python_embeded is reachable via PATH or the inherited working
directory, and then every stdlib C extension fails with "Module use of
python313.dll conflicts with this version of Python" (confirmed twice
through this UI). Dropping python_embeded from PATH -- combined with the
fixed cwd below -- keeps the child's DLL search away from it.
"""
env = {**os.environ, "PYTHONIOENCODING": "utf-8", "PYTHONUNBUFFERED": "1"}
parts = env.get("PATH", "").split(os.pathsep)
embeded = os.path.normcase(EMBEDED_DIR)
env["PATH"] = os.pathsep.join(
p for p in parts if not os.path.normcase(p).startswith(embeded)
)
return env
def stream_subprocess(cmd):
"""Runs cmd, yielding the accumulated stdout text as new lines arrive."""
proc = subprocess.Popen(
cmd, env=child_env(), cwd=HERE, stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
text=True, encoding="utf-8", errors="replace", bufsize=1,
)
log = ""
for line in proc.stdout:
log += line
yield log
proc.wait()
if proc.returncode not in (0, None):
if (proc.returncode & 0xFFFFFFFF) == BENIGN_TORCH_TEARDOWN_NTSTATUS:
log += "\n(non-fatal: torch/CUDA teardown crash after finishing its work -- known environment quirk)"
else:
log += f"\n[exit code {proc.returncode}]"
yield log
# ---------------------------------------------------------------------------
# Checkpoint scanning (shared by Setup + Training tabs)
# ---------------------------------------------------------------------------
def checkpoint_status_rows():
rows = []
for _repo, path, subdir, size in FILES["fp16"] + COMMON_FILES:
dest = os.path.join(MODELS_ROOT, subdir, os.path.basename(path))
ok = os.path.exists(dest)
actual = f"{os.path.getsize(dest) / 1e9:.1f}GB" if ok else "-"
rows.append([os.path.basename(path), "์žˆ์Œ" if ok else "์—†์Œ", f"{size / 1e9:.1f}GB", actual])
return rows
def venv_status_text():
return "์„ค์น˜๋จ" if os.path.exists(VENV_PYTHON) else "์„ค์น˜ ์•ˆ ๋จ"
def list_checkpoints(subdir, name_filter=None, exclude=None):
d = os.path.join(MODELS_ROOT, subdir)
if not os.path.isdir(d):
return []
files = os.listdir(d)
if name_filter:
files = [f for f in files if name_filter in f.lower()]
if exclude:
files = [f for f in files if exclude not in f.lower()]
return [os.path.join(d, f) for f in sorted(files)]
def list_dit_choices():
# fp8_scaled DiT files can't be used for training (confirmed this session:
# musubi-tuner's loader rejects ComfyUI's pre-quantized scaled-fp8 format
# both with and without --fp8-scaled) -- excluded here so the dropdown
# can't steer anyone into that dead end.
return list_checkpoints("diffusion_models", name_filter="wan2.2_i2v", exclude="fp8_scaled")
def list_t5_choices():
return list_checkpoints("text_encoders", name_filter="t5")
def list_vae_choices():
return list_checkpoints("vae", name_filter="wan2.1_vae")
# ---------------------------------------------------------------------------
# Setup tab
# ---------------------------------------------------------------------------
def run_setup_env():
yield from stream_subprocess(["powershell.exe", "-File", os.path.join(TRAIN_DIR, "setup_env.ps1"), "-Yes"])
def run_download_checkpoints():
yield from stream_subprocess([
PYTHON, os.path.join(TRAIN_DIR, "download_checkpoints.py"),
"--models-root", MODELS_ROOT, "--precision", "fp16", "--yes",
])
def build_setup_tab():
with gr.Tab("Setup"):
gr.Markdown(
"ํ•™์Šต์—๋Š” fp16 ์ฒดํฌํฌ์ธํŠธ๊ฐ€ ํ•„์š”ํ•ฉ๋‹ˆ๋‹ค (fp8_scaled๋Š” ComfyUI ์ถ”๋ก  ์ „์šฉ โ€” "
"musubi-tuner ํ•™์Šต ๋กœ๋”๊ฐ€ ๋ชป ์ฝ๋Š” ๊ฒƒ์œผ๋กœ ํ™•์ธ๋จ). ์•„๋ž˜์—์„œ ํ•™์Šต ํ™˜๊ฒฝ/์ฒดํฌํฌ์ธํŠธ "
"์ƒํƒœ๋ฅผ ํ™•์ธํ•˜๊ณ  ์—†๋Š” ๊ฒƒ๋งŒ ๋ฐ›์œผ์„ธ์š”."
)
with gr.Row():
venv_box = gr.Textbox(label="ํ•™์Šต ํ™˜๊ฒฝ(venv)", value=venv_status_text(), interactive=False, scale=1)
setup_btn = gr.Button("ํ•™์Šต ํ™˜๊ฒฝ ์„ค์น˜ (venv + musubi-tuner)", scale=2)
setup_log = gr.Textbox(label="์„ค์น˜ ๋กœ๊ทธ", lines=10, max_lines=20, interactive=False)
setup_btn.click(run_setup_env, outputs=setup_log).then(venv_status_text, outputs=venv_box)
gr.Markdown("### ์ฒดํฌํฌ์ธํŠธ")
ckpt_table = gr.Dataframe(
headers=["ํŒŒ์ผ", "์ƒํƒœ", "์˜ˆ์ƒ ํฌ๊ธฐ", "์‹ค์ œ ํฌ๊ธฐ"],
value=checkpoint_status_rows(), interactive=False,
)
with gr.Row():
refresh_btn = gr.Button("์ƒํƒœ ์ƒˆ๋กœ๊ณ ์นจ")
confirm_dl = gr.Checkbox(label="์œ„ ๋ชฉ๋ก ์ค‘ ์—†๋Š” fp16 ์ฒดํฌํฌ์ธํŠธ๋ฅผ ๋‹ค์šด๋กœ๋“œํ•˜๊ฒ ์Šต๋‹ˆ๋‹ค (์ด ์ˆ˜์‹ญGB)")
dl_btn = gr.Button("๋‹ค์šด๋กœ๋“œ", interactive=False)
dl_log = gr.Textbox(label="๋‹ค์šด๋กœ๋“œ ๋กœ๊ทธ", lines=10, max_lines=20, interactive=False)
refresh_btn.click(checkpoint_status_rows, outputs=ckpt_table)
confirm_dl.change(lambda v: gr.update(interactive=v), inputs=confirm_dl, outputs=dl_btn)
dl_btn.click(run_download_checkpoints, outputs=dl_log).then(checkpoint_status_rows, outputs=ckpt_table)
# ---------------------------------------------------------------------------
# Dataset Builder tab
# ---------------------------------------------------------------------------
def run_build(input_dir, output_dir, trigger, captioner, max_tier, face_pad, person_pad,
video_fps, clip_min_sec, clip_max_sec, image_format, clips_per_video,
body_region, body_region_frac):
if not input_dir or not output_dir:
yield "์ž…๋ ฅ/์ถœ๋ ฅ ํด๋”๋ฅผ ๋ชจ๋‘ ์ž…๋ ฅํ•˜์„ธ์š”."
return
cmd = [
PYTHON, os.path.join(HERE, "build.py"),
"--input", input_dir, "--output", output_dir, "--trigger", trigger,
"--captioner", captioner, "--max-tier", max_tier,
"--face-pad", str(face_pad), "--person-pad", str(person_pad),
"--video-fps", str(video_fps), "--clip-min-sec", str(clip_min_sec),
"--clip-max-sec", str(clip_max_sec), "--image-format", image_format,
"--body-region", body_region, "--body-region-frac", str(body_region_frac),
"--clips-per-video", str(int(clips_per_video or 0)),
]
yield from stream_subprocess(cmd)
def preview_gallery(output_dir):
img_dir = os.path.join(output_dir or "", "images")
if not os.path.isdir(img_dir):
return []
files = sorted(f for f in os.listdir(img_dir) if f.lower().endswith((".jpg", ".png", ".webp")))
return [os.path.join(img_dir, f) for f in files[:12]]
def build_dataset_tab(dataset_dir_handoff):
with gr.Tab("Dataset Builder"):
gr.Markdown("์ž„์˜์˜ ์ด๋ฏธ์ง€/์˜์ƒ ํด๋”๋ฅผ ํฌ๋กญ+์บก์…˜ ๋ฐ์ดํ„ฐ์…‹์œผ๋กœ ๋ณ€ํ™˜ํ•ฉ๋‹ˆ๋‹ค.")
with gr.Row():
input_dir = gr.Textbox(label="์ž…๋ ฅ ํด๋”", placeholder=r"D:\raw_media")
output_dir = gr.Textbox(label="์ถœ๋ ฅ ํด๋”", placeholder=r"D:\wan_lora_dataset")
with gr.Row():
trigger = gr.Textbox(label="ํŠธ๋ฆฌ๊ฑฐ ์›Œ๋“œ", placeholder="ohwx man")
captioner = gr.Radio(["joycaption", "florence2", "wd14"], value="joycaption", label="์บก์…”๋„ˆ")
with gr.Accordion("๊ณ ๊ธ‰ ์„ค์ •", open=False):
with gr.Row():
max_tier = gr.Radio(["480p", "720p", "1080p"], value="1080p", label="์ตœ๋Œ€ ํ•ด์ƒ๋„ ํ‹ฐ์–ด")
image_format = gr.Radio(["jpg", "png"], value="jpg", label="์ด๋ฏธ์ง€ ์ €์žฅ ํฌ๋งท")
with gr.Row():
face_pad = gr.Number(value=2.2, label="์–ผ๊ตด ํฌ๋กญ ํŒจ๋”ฉ")
person_pad = gr.Number(value=1.3, label="์‚ฌ๋žŒ ํฌ๋กญ ํŒจ๋”ฉ")
with gr.Row():
body_region = gr.Radio(
["auto", "upper", "lower"], value="auto", label="์‹ ์ฒด ๋ถ€์œ„ ์šฐ์„  ํฌ๋กญ",
info="๋Œ€๋ถ€๋ถ„ ์†Œ์Šค๊ฐ€ ์ด๋ฏธ ํŠน์ • ๋ถ€์œ„ ํด๋กœ์ฆˆ์—…์ด๋ฉด upper/lower๋กœ ๊ทธ์ชฝ๋งŒ ์ขํ˜€์„œ ํฌ๋กญ",
)
body_region_frac = gr.Number(value=0.45, label="๋ถ€์œ„ ํฌ๋กญ ๋น„์œจ (bbox ๋†’์ด ์ค‘ ๋‚จ๊ธธ ๋น„์œจ)")
with gr.Row():
video_fps = gr.Number(value=16, label="์˜์ƒ ์ถœ๋ ฅ FPS")
clip_min_sec = gr.Number(value=2.0, label="ํด๋ฆฝ ์ตœ์†Œ ๊ธธ์ด(์ดˆ)")
clip_max_sec = gr.Number(value=6.0, label="ํด๋ฆฝ ์ตœ๋Œ€ ๊ธธ์ด(์ดˆ)")
clips_per_video = gr.Number(
value=5, label="์˜์ƒ 1๊ฐœ๋‹น ํด๋ฆฝ ์ˆ˜ ์ œํ•œ (0์ด๋ฉด ์ œํ•œ ์—†์Œ)",
info="๊ธด ์˜์ƒ(์˜ˆ: 5๋ถ„)์—์„œ ์”ฌ ์ปท์ด ๋งŽ์ด ์žกํ˜€ ํด๋ฆฝ์ด ๊ณผํ•˜๊ฒŒ ๋Š˜์–ด๋‚  ๋•Œ, "
"์˜์ƒ ์ „์ฒด ๊ตฌ๊ฐ„์— ๊ฑธ์ณ ๊ท ๋“ฑํ•˜๊ฒŒ N๊ฐœ๋งŒ ๋ฝ‘๋„๋ก ์ œํ•œํ•ฉ๋‹ˆ๋‹ค.",
)
build_btn = gr.Button("๋ฐ์ดํ„ฐ์…‹ ๋งŒ๋“ค๊ธฐ", variant="primary")
build_log = gr.Textbox(label="๋กœ๊ทธ", lines=12, max_lines=25, interactive=False)
gallery = gr.Gallery(label="๊ฒฐ๊ณผ ๋ฏธ๋ฆฌ๋ณด๊ธฐ (์ด๋ฏธ์ง€)", columns=4, height=300)
send_btn = gr.Button("โ†’ ํ•™์Šต ํƒญ์— ์ด ์ถœ๋ ฅ ํด๋” ์ฑ„์šฐ๊ธฐ")
build_btn.click(
run_build,
inputs=[input_dir, output_dir, trigger, captioner, max_tier, face_pad, person_pad,
video_fps, clip_min_sec, clip_max_sec, image_format, clips_per_video,
body_region, body_region_frac],
outputs=build_log,
).then(preview_gallery, inputs=output_dir, outputs=gallery)
send_btn.click(lambda d: d, inputs=output_dir, outputs=dataset_dir_handoff)
# ---------------------------------------------------------------------------
# Training tab
# ---------------------------------------------------------------------------
def pick_default(choices, name_hint):
"""Prefer a choice whose filename contains name_hint (e.g. 'low_noise');
plain alphabetical order would otherwise default the low-noise dropdown
to the high_noise file since 'h' < 'l'."""
for c in choices:
if name_hint in os.path.basename(c).lower():
return c
return choices[0] if choices else None
def infer_mixed_precision(dit_low_path):
name = (dit_low_path or "").lower()
if "bf16" in name:
return "bf16"
return "fp16"
def run_train(dataset_dir, dit_low, dit_high, t5, vae, network_dim, learning_rate,
use_steps, max_train_steps, max_train_epochs, mixed_precision,
separate_passes, skip_cache, output_dir):
if not all([dataset_dir, dit_low, t5, vae, output_dir]):
yield "dataset-dir / DiT-low / T5 / VAE / output-dir๋Š” ํ•„์ˆ˜์ž…๋‹ˆ๋‹ค."
return
if not os.path.exists(VENV_PYTHON):
yield f"ํ•™์Šต ํ™˜๊ฒฝ(venv)์ด ์—†์Šต๋‹ˆ๋‹ค: {VENV_PYTHON}\nSetup ํƒญ์—์„œ ๋จผ์ € ์„ค์น˜ํ•˜์„ธ์š”."
return
cmd = [
# train.py itself has no heavy deps, but running it with python_embeded
# (which on this machine is Python 3.13, sitting on a PATH with several
# other Python installs) leaked into the venv subprocess's inherited
# environment and broke stdlib C-extension loading in the venv (Python
# 3.10) with "Module use of python313.dll conflicts with this version
# of Python" -- confirmed via a real run through this UI. Running
# train.py with the venv's own interpreter keeps the whole chain
# self-consistent, matching every successful CLI run this session.
VENV_PYTHON, os.path.join(TRAIN_DIR, "train.py"),
"--dataset-dir", dataset_dir, "--dit-low", dit_low, "--t5", t5, "--vae", vae,
"--network-dim", str(int(network_dim)), "--learning-rate", str(learning_rate),
"--mixed-precision", mixed_precision, "--output-dir", output_dir,
]
if dit_high:
cmd += ["--dit-high", dit_high]
if use_steps:
cmd += ["--max-train-steps", str(int(max_train_steps))]
else:
cmd += ["--max-train-epochs", str(int(max_train_epochs))]
if separate_passes:
cmd.append("--separate-passes")
if skip_cache:
cmd.append("--skip-cache")
yield from stream_subprocess(cmd)
def build_training_tab(dataset_dir_handoff):
with gr.Tab("Training"):
gr.Markdown(
"โš ๏ธ ComfyUI๋ฅผ ์ผœ ๋‘” ์ฑ„๋กœ ํ•™์Šตํ•˜๋ฉด VRAM์ด ๋ถ€์กฑํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค โ€” ๊ฐ€๋Šฅํ•˜๋ฉด "
"ํ•™์Šต ์ค‘์—” ComfyUI ์„œ๋ฒ„๋ฅผ ๋‚ด๋ ค๋‘์„ธ์š”."
)
dataset_dir = gr.Textbox(label="๋ฐ์ดํ„ฐ์…‹ ํด๋” (dataset.toml์ด ์žˆ๋Š” ๊ณณ)")
dataset_dir_handoff.change(lambda d: d, inputs=dataset_dir_handoff, outputs=dataset_dir)
dit_choices = list_dit_choices()
t5_choices = list_t5_choices()
vae_choices = list_vae_choices()
with gr.Row():
dit_low = gr.Dropdown(label="DiT (low-noise)", choices=dit_choices,
value=pick_default(dit_choices, "low_noise"))
dit_high = gr.Dropdown(label="DiT (high-noise, ์„ ํƒ)", choices=dit_choices,
value=pick_default(dit_choices, "high_noise"))
with gr.Row():
t5 = gr.Dropdown(label="T5 ํ…์ŠคํŠธ ์ธ์ฝ”๋”", choices=t5_choices,
value=t5_choices[0] if t5_choices else None)
vae = gr.Dropdown(label="VAE", choices=vae_choices,
value=pick_default(vae_choices, "wan2.1_vae"))
refresh_ckpt_btn = gr.Button("์ฒดํฌํฌ์ธํŠธ ๋ชฉ๋ก ์ƒˆ๋กœ๊ณ ์นจ (Setup ํƒญ์—์„œ ๋‹ค์šด๋กœ๋“œํ•œ ๋’ค ๋ˆŒ๋Ÿฌ์ฃผ์„ธ์š”)")
with gr.Row():
network_dim = gr.Number(value=32, label="LoRA rank (network_dim)")
learning_rate = gr.Number(value=2e-4, label="Learning rate")
mixed_precision = gr.Radio(["fp16", "bf16", "no"], value="fp16", label="Mixed precision")
with gr.Row():
use_steps = gr.Checkbox(label="์Šคํ… ์ˆ˜๋กœ ์ œํ•œ (์Šค๋ชจํฌ ํ…Œ์ŠคํŠธ์šฉ)", value=False)
max_train_steps = gr.Number(value=5, label="max_train_steps")
max_train_epochs = gr.Number(value=16, label="max_train_epochs")
with gr.Row():
separate_passes = gr.Checkbox(label="high/low ๋”ฐ๋กœ ํ•™์Šต (ํŒŒ์ผ 2๊ฐœ, ์‹œ๊ฐ„ 2๋ฐฐ)", value=False)
skip_cache = gr.Checkbox(label="์บ์‹ฑ ๋‹จ๊ณ„ ๊ฑด๋„ˆ๋›ฐ๊ธฐ (์ด๋ฏธ ์บ์‹ฑ๋œ ๋ฐ์ดํ„ฐ์…‹)", value=False)
output_dir = gr.Textbox(label="์ถœ๋ ฅ ํด๋”", placeholder=r"D:\wan_lora_output")
train_btn = gr.Button("ํ•™์Šต ์‹œ์ž‘", variant="primary")
train_log = gr.Textbox(label="๋กœ๊ทธ", lines=16, max_lines=30, interactive=False)
def refresh_dit_low():
c = list_dit_choices()
return gr.update(choices=c, value=pick_default(c, "low_noise"))
def refresh_dit_high():
c = list_dit_choices()
return gr.update(choices=c, value=pick_default(c, "high_noise"))
def refresh_t5():
c = list_t5_choices()
return gr.update(choices=c, value=c[0] if c else None)
def refresh_vae():
c = list_vae_choices()
return gr.update(choices=c, value=pick_default(c, "wan2.1_vae"))
refresh_ckpt_btn.click(refresh_dit_low, outputs=dit_low)
refresh_ckpt_btn.click(refresh_dit_high, outputs=dit_high)
refresh_ckpt_btn.click(refresh_t5, outputs=t5)
refresh_ckpt_btn.click(refresh_vae, outputs=vae)
dit_low.change(infer_mixed_precision, inputs=dit_low, outputs=mixed_precision)
train_btn.click(
run_train,
inputs=[dataset_dir, dit_low, dit_high, t5, vae, network_dim, learning_rate,
use_steps, max_train_steps, max_train_epochs, mixed_precision,
separate_passes, skip_cache, output_dir],
outputs=train_log,
)
# ---------------------------------------------------------------------------
def build_app():
with gr.Blocks(title="Wan 2.2 LoRA Builder") as app:
gr.Markdown("# ๐Ÿฌ Wan 2.2 LoRA Dataset + Training")
build_setup_tab()
dataset_dir_handoff = gr.Textbox(visible=False)
build_dataset_tab(dataset_dir_handoff)
build_training_tab(dataset_dir_handoff)
return app
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--port", type=int, default=7860)
return p.parse_args()
if __name__ == "__main__":
args = parse_args()
build_app().queue().launch(server_name="127.0.0.1", server_port=args.port, share=False)