Download custom_nodes/dolphin_nodes/dataset_builder/webui.py from bjooo/tutorials: direct link, hf CLI and curl.
- Browser
- Download file 18.5 kB
-
https://huggingface.co/bjooo/tutorials/resolve/main/custom_nodes/dolphin_nodes/dataset_builder/webui.py
- Command line
-
hf download hf://bjooo/tutorials/custom_nodes/dolphin_nodes/dataset_builder/webui.py
-
curl -L -o webui.py https://huggingface.co/bjooo/tutorials/resolve/main/custom_nodes/dolphin_nodes/dataset_builder/webui.py
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) | |