"""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)