File size: 18,534 Bytes
9c98083
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
"""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)