Spaces:
Runtime error
Runtime error
Add layered training functionality and enhance error handling in app.py
Browse filesIntroduce a new function, run_training_layered, to support layered image training, allowing users to upload multiple layers for processing. Implement error handling for required inputs and improve logging for better user feedback. Update the UI to include a dedicated tab for layered training, enhancing user experience by providing clear options for both standard and layered training processes. Additionally, refactor existing training logic to streamline control flow and maintain consistency across different training modes.
app.py
CHANGED
|
@@ -50,6 +50,8 @@ def _bash_quote(s: str) -> str:
|
|
| 50 |
|
| 51 |
|
| 52 |
_QWEN_IMAGE_TYPES = ("edit-2509", "edit-2511", "layered")
|
|
|
|
|
|
|
| 53 |
|
| 54 |
|
| 55 |
def _get_qwen_image_type() -> str:
|
|
@@ -847,6 +849,10 @@ def run_training(
|
|
| 847 |
image_type = _get_qwen_image_type()
|
| 848 |
log_buf += f"[QIE] Model type: {image_type}\n"
|
| 849 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 850 |
|
| 851 |
# Ensure /auto holds helper files expected by the script
|
| 852 |
_ensure_workspace_auto_files()
|
|
@@ -883,17 +889,11 @@ def run_training(
|
|
| 883 |
_extract_paths(control6_uploads),
|
| 884 |
_extract_paths(control7_uploads),
|
| 885 |
]
|
| 886 |
-
|
| 887 |
-
|
| 888 |
-
|
| 889 |
-
|
| 890 |
-
|
| 891 |
-
else:
|
| 892 |
-
# Require control_0; others optional
|
| 893 |
-
if not control_upload_sets[0]:
|
| 894 |
-
log_buf += "[ERROR] control_0 images are required.\n"
|
| 895 |
-
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 896 |
-
return
|
| 897 |
|
| 898 |
control_dirs: List[Optional[str]] = []
|
| 899 |
for i, uploads in enumerate(control_upload_sets):
|
|
@@ -909,11 +909,6 @@ def run_training(
|
|
| 909 |
log_buf += f"[QIE] Copied {len(uploads)} control_{i} images to {cdir}\n"
|
| 910 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 911 |
|
| 912 |
-
control_dirs_abs = [
|
| 913 |
-
(os.path.join(ds_dir, name) if name else None)
|
| 914 |
-
for name in control_dirs
|
| 915 |
-
]
|
| 916 |
-
|
| 917 |
# Prepare script with user parameters
|
| 918 |
control_folders = [
|
| 919 |
(control_dirs[i] if control_dirs[i] else None)
|
|
@@ -933,24 +928,16 @@ def run_training(
|
|
| 933 |
|
| 934 |
# Update dataset config with requested resolution/batch settings
|
| 935 |
try:
|
| 936 |
-
ctrl_w = int(control_res_w) if control_res_w else None
|
| 937 |
-
ctrl_h = int(control_res_h) if control_res_h else None
|
| 938 |
-
if image_type == "layered":
|
| 939 |
-
ctrl_w = None
|
| 940 |
-
ctrl_h = None
|
| 941 |
_update_dataset_toml(
|
| 942 |
ds_conf,
|
| 943 |
img_res_w=int(train_res_w) if train_res_w else None,
|
| 944 |
img_res_h=int(train_res_h) if train_res_h else None,
|
| 945 |
train_batch_size=int(train_batch_size) if train_batch_size else None,
|
| 946 |
-
control_res_w=
|
| 947 |
-
control_res_h=
|
| 948 |
-
multiple_target=
|
| 949 |
)
|
| 950 |
-
|
| 951 |
-
log_buf += f"[QIE] Updated dataset config: resolution=({train_res_w},{train_res_h}), batch_size={train_batch_size}, multiple_target=true\n"
|
| 952 |
-
else:
|
| 953 |
-
log_buf += f"[QIE] Updated dataset config: resolution=({train_res_w},{train_res_h}), batch_size={train_batch_size}, control_res=({control_res_w},{control_res_h})\n"
|
| 954 |
except Exception as e:
|
| 955 |
log_buf += f"[QIE] WARN: failed to update dataset config: {e}\n"
|
| 956 |
# Expose dataset config for download (if exists)
|
|
@@ -1004,39 +991,8 @@ def run_training(
|
|
| 1004 |
pass
|
| 1005 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1006 |
|
| 1007 |
-
output_json = os.path.join(out_base, "metadata.jsonl")
|
| 1008 |
-
if image_type == "layered":
|
| 1009 |
-
_sync_dataset_config_jsonl(ds_conf, output_json)
|
| 1010 |
-
log_buf += f"[QIE] Generating layered metadata: {output_json}\n"
|
| 1011 |
-
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1012 |
-
try:
|
| 1013 |
-
count = _generate_layered_jsonl(
|
| 1014 |
-
image_dir=img_dir,
|
| 1015 |
-
caption=caption,
|
| 1016 |
-
output_json=output_json,
|
| 1017 |
-
control_dirs=control_dirs_abs,
|
| 1018 |
-
target_prefix=(target_prefix or ""),
|
| 1019 |
-
target_suffix=(target_suffix or ""),
|
| 1020 |
-
control_prefixes=control_prefixes,
|
| 1021 |
-
control_suffixes=control_suffixes,
|
| 1022 |
-
allow_single=True,
|
| 1023 |
-
)
|
| 1024 |
-
log_buf += f"[QIE] Layered metadata written: {count} entries\n"
|
| 1025 |
-
except Exception as e:
|
| 1026 |
-
log_buf += f"[ERROR] Layered metadata generation failed: {e}\n"
|
| 1027 |
-
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1028 |
-
return
|
| 1029 |
-
|
| 1030 |
-
if os.path.isfile(output_json) and output_json not in artifacts:
|
| 1031 |
-
artifacts.append(output_json)
|
| 1032 |
-
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1033 |
-
|
| 1034 |
-
if config_only:
|
| 1035 |
-
log_buf += "[QIE] Config-only mode: skipping cache/training.\n"
|
| 1036 |
-
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1037 |
-
return
|
| 1038 |
-
|
| 1039 |
if config_only:
|
|
|
|
| 1040 |
_sync_dataset_config_jsonl(ds_conf, output_json)
|
| 1041 |
log_buf += f"[QIE] Generating metadata: {output_json}\n"
|
| 1042 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
|
@@ -1142,6 +1098,293 @@ def run_training(
|
|
| 1142 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1143 |
|
| 1144 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1145 |
def build_ui() -> gr.Blocks:
|
| 1146 |
css = """
|
| 1147 |
.pad-section {
|
|
@@ -1167,8 +1410,24 @@ def build_ui() -> gr.Blocks:
|
|
| 1167 |
}
|
| 1168 |
"""
|
| 1169 |
with gr.Blocks(title="Qwen-Image-Edit: Trainer", css=css) as demo:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1170 |
with gr.Tabs() as tabs:
|
| 1171 |
-
with gr.TabItem("
|
| 1172 |
gr.Markdown("""
|
| 1173 |
# Qwen-Image-Edit Trainer
|
| 1174 |
学習に使う画像をアップロードし、必要ならファイル名の前後にある共通の文字(prefix/suffix)を指定して、
|
|
@@ -1355,28 +1614,126 @@ def build_ui() -> gr.Blocks:
|
|
| 1355 |
outputs=[logs, ckpt_files, scripts_files, run_out_dir_box],
|
| 1356 |
)
|
| 1357 |
|
| 1358 |
-
# 回収ボタン: 直近の dataset_ ディレクトリからチェックポイントとスクリプト/設定を再取得
|
| 1359 |
-
def _refresh_all() -> tuple:
|
| 1360 |
-
try:
|
| 1361 |
-
ds_dir = _find_latest_dataset_dir(DATA_ROOT_RUNTIME)
|
| 1362 |
-
except Exception:
|
| 1363 |
-
ds_dir = None
|
| 1364 |
-
try:
|
| 1365 |
-
ck = _list_checkpoints(ds_dir) if ds_dir else []
|
| 1366 |
-
except Exception:
|
| 1367 |
-
ck = []
|
| 1368 |
-
try:
|
| 1369 |
-
sc = _collect_scripts_and_config(ds_dir)
|
| 1370 |
-
except Exception:
|
| 1371 |
-
sc = _collect_scripts_and_config(None)
|
| 1372 |
-
return ck, sc
|
| 1373 |
-
|
| 1374 |
refresh_scripts_btn.click(
|
| 1375 |
fn=_refresh_all,
|
| 1376 |
inputs=[],
|
| 1377 |
outputs=[ckpt_files, scripts_files],
|
| 1378 |
)
|
| 1379 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1380 |
with gr.TabItem("Prompt Generator"):
|
| 1381 |
gr.Markdown("""
|
| 1382 |
# 🎨 A→B 変換プロンプト自動生成
|
|
@@ -1443,7 +1800,7 @@ if __name__ == "__main__":
|
|
| 1443 |
_startup_install_musubi_deps()
|
| 1444 |
|
| 1445 |
# 2) Download models at startup (blocking by design)
|
| 1446 |
-
_startup_download_models()
|
| 1447 |
|
| 1448 |
# 3) Launch Gradio app
|
| 1449 |
ui = build_ui()
|
|
|
|
| 50 |
|
| 51 |
|
| 52 |
_QWEN_IMAGE_TYPES = ("edit-2509", "edit-2511", "layered")
|
| 53 |
+
EDIT_CONTROL_MAX = 8
|
| 54 |
+
LAYER_MAX = 32
|
| 55 |
|
| 56 |
|
| 57 |
def _get_qwen_image_type() -> str:
|
|
|
|
| 849 |
image_type = _get_qwen_image_type()
|
| 850 |
log_buf += f"[QIE] Model type: {image_type}\n"
|
| 851 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 852 |
+
if image_type == "layered":
|
| 853 |
+
log_buf += "[ERROR] Layered は Layered タブで実行してください。QWEN_IMAGE_TYPE=layered を設定して再起動してください。\n"
|
| 854 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 855 |
+
return
|
| 856 |
|
| 857 |
# Ensure /auto holds helper files expected by the script
|
| 858 |
_ensure_workspace_auto_files()
|
|
|
|
| 889 |
_extract_paths(control6_uploads),
|
| 890 |
_extract_paths(control7_uploads),
|
| 891 |
]
|
| 892 |
+
# Require control_0; others optional
|
| 893 |
+
if not control_upload_sets[0]:
|
| 894 |
+
log_buf += "[ERROR] control_0 images are required.\n"
|
| 895 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 896 |
+
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 897 |
|
| 898 |
control_dirs: List[Optional[str]] = []
|
| 899 |
for i, uploads in enumerate(control_upload_sets):
|
|
|
|
| 909 |
log_buf += f"[QIE] Copied {len(uploads)} control_{i} images to {cdir}\n"
|
| 910 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 911 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 912 |
# Prepare script with user parameters
|
| 913 |
control_folders = [
|
| 914 |
(control_dirs[i] if control_dirs[i] else None)
|
|
|
|
| 928 |
|
| 929 |
# Update dataset config with requested resolution/batch settings
|
| 930 |
try:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 931 |
_update_dataset_toml(
|
| 932 |
ds_conf,
|
| 933 |
img_res_w=int(train_res_w) if train_res_w else None,
|
| 934 |
img_res_h=int(train_res_h) if train_res_h else None,
|
| 935 |
train_batch_size=int(train_batch_size) if train_batch_size else None,
|
| 936 |
+
control_res_w=int(control_res_w) if control_res_w else None,
|
| 937 |
+
control_res_h=int(control_res_h) if control_res_h else None,
|
| 938 |
+
multiple_target=False,
|
| 939 |
)
|
| 940 |
+
log_buf += f"[QIE] Updated dataset config: resolution=({train_res_w},{train_res_h}), batch_size={train_batch_size}, control_res=({control_res_w},{control_res_h})\n"
|
|
|
|
|
|
|
|
|
|
| 941 |
except Exception as e:
|
| 942 |
log_buf += f"[QIE] WARN: failed to update dataset config: {e}\n"
|
| 943 |
# Expose dataset config for download (if exists)
|
|
|
|
| 991 |
pass
|
| 992 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 993 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 994 |
if config_only:
|
| 995 |
+
output_json = os.path.join(out_base, "metadata.jsonl")
|
| 996 |
_sync_dataset_config_jsonl(ds_conf, output_json)
|
| 997 |
log_buf += f"[QIE] Generating metadata: {output_json}\n"
|
| 998 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
|
|
|
| 1098 |
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1099 |
|
| 1100 |
|
| 1101 |
+
def run_training_layered(
|
| 1102 |
+
output_name: str,
|
| 1103 |
+
caption: str,
|
| 1104 |
+
image_uploads: Any,
|
| 1105 |
+
target_prefix: str,
|
| 1106 |
+
target_suffix: str,
|
| 1107 |
+
layer_prefix: str,
|
| 1108 |
+
layer_suffix: str,
|
| 1109 |
+
layer1_uploads: Any,
|
| 1110 |
+
layer2_uploads: Any,
|
| 1111 |
+
layer3_uploads: Any,
|
| 1112 |
+
layer4_uploads: Any,
|
| 1113 |
+
layer5_uploads: Any,
|
| 1114 |
+
layer6_uploads: Any,
|
| 1115 |
+
layer7_uploads: Any,
|
| 1116 |
+
layer8_uploads: Any,
|
| 1117 |
+
layer9_uploads: Any,
|
| 1118 |
+
layer10_uploads: Any,
|
| 1119 |
+
layer11_uploads: Any,
|
| 1120 |
+
layer12_uploads: Any,
|
| 1121 |
+
layer13_uploads: Any,
|
| 1122 |
+
layer14_uploads: Any,
|
| 1123 |
+
layer15_uploads: Any,
|
| 1124 |
+
layer16_uploads: Any,
|
| 1125 |
+
layer17_uploads: Any,
|
| 1126 |
+
layer18_uploads: Any,
|
| 1127 |
+
layer19_uploads: Any,
|
| 1128 |
+
layer20_uploads: Any,
|
| 1129 |
+
layer21_uploads: Any,
|
| 1130 |
+
layer22_uploads: Any,
|
| 1131 |
+
layer23_uploads: Any,
|
| 1132 |
+
layer24_uploads: Any,
|
| 1133 |
+
layer25_uploads: Any,
|
| 1134 |
+
layer26_uploads: Any,
|
| 1135 |
+
layer27_uploads: Any,
|
| 1136 |
+
layer28_uploads: Any,
|
| 1137 |
+
layer29_uploads: Any,
|
| 1138 |
+
layer30_uploads: Any,
|
| 1139 |
+
layer31_uploads: Any,
|
| 1140 |
+
layer32_uploads: Any,
|
| 1141 |
+
learning_rate: str,
|
| 1142 |
+
network_dim: int,
|
| 1143 |
+
train_res_w: int,
|
| 1144 |
+
train_res_h: int,
|
| 1145 |
+
train_batch_size: int,
|
| 1146 |
+
te_cache_batch_size: int,
|
| 1147 |
+
seed: int,
|
| 1148 |
+
max_epochs: int,
|
| 1149 |
+
save_every: int,
|
| 1150 |
+
config_only: bool,
|
| 1151 |
+
) -> Iterable[tuple]:
|
| 1152 |
+
log_buf = "[QIE] Start Training invoked.\n"
|
| 1153 |
+
ckpts: List[str] = []
|
| 1154 |
+
artifacts: List[str] = []
|
| 1155 |
+
run_out_dir = ""
|
| 1156 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1157 |
+
if not output_name.strip():
|
| 1158 |
+
log_buf += "[ERROR] OUTPUT NAME is required.\n"
|
| 1159 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1160 |
+
return
|
| 1161 |
+
if not caption.strip():
|
| 1162 |
+
log_buf += "[ERROR] CAPTION is required.\n"
|
| 1163 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1164 |
+
return
|
| 1165 |
+
|
| 1166 |
+
image_type = _get_qwen_image_type()
|
| 1167 |
+
log_buf += f"[QIE] Model type: {image_type}\n"
|
| 1168 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1169 |
+
if image_type != "layered":
|
| 1170 |
+
log_buf += "[ERROR] Layered タブは QWEN_IMAGE_TYPE=layered のみ対応です。環境変数を設定して再起動してください。\n"
|
| 1171 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1172 |
+
return
|
| 1173 |
+
|
| 1174 |
+
_ensure_workspace_auto_files()
|
| 1175 |
+
global DATA_ROOT_RUNTIME
|
| 1176 |
+
DATA_ROOT_RUNTIME = _ensure_data_root(None)
|
| 1177 |
+
|
| 1178 |
+
import time
|
| 1179 |
+
ds_name = f"dataset_{int(time.time())}"
|
| 1180 |
+
ds_dir = os.path.abspath(os.path.join(DATA_ROOT_RUNTIME, ds_name))
|
| 1181 |
+
run_out_dir = os.path.abspath(os.path.join(ds_dir, output_name.strip()))
|
| 1182 |
+
img_folder_name = DEFAULT_IMAGE_FOLDER
|
| 1183 |
+
img_dir = os.path.join(ds_dir, img_folder_name)
|
| 1184 |
+
os.makedirs(img_dir, exist_ok=True)
|
| 1185 |
+
|
| 1186 |
+
base_files = _extract_paths(image_uploads)
|
| 1187 |
+
if not base_files:
|
| 1188 |
+
log_buf += "[ERROR] No images uploaded for IMAGE_FOLDER.\n"
|
| 1189 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1190 |
+
return
|
| 1191 |
+
base_filenames = _copy_uploads(base_files, img_dir)
|
| 1192 |
+
log_buf += f"[QIE] Copied {len(base_filenames)} base images to {img_dir}\n"
|
| 1193 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1194 |
+
|
| 1195 |
+
layer_upload_sets = [
|
| 1196 |
+
_extract_paths(layer1_uploads),
|
| 1197 |
+
_extract_paths(layer2_uploads),
|
| 1198 |
+
_extract_paths(layer3_uploads),
|
| 1199 |
+
_extract_paths(layer4_uploads),
|
| 1200 |
+
_extract_paths(layer5_uploads),
|
| 1201 |
+
_extract_paths(layer6_uploads),
|
| 1202 |
+
_extract_paths(layer7_uploads),
|
| 1203 |
+
_extract_paths(layer8_uploads),
|
| 1204 |
+
_extract_paths(layer9_uploads),
|
| 1205 |
+
_extract_paths(layer10_uploads),
|
| 1206 |
+
_extract_paths(layer11_uploads),
|
| 1207 |
+
_extract_paths(layer12_uploads),
|
| 1208 |
+
_extract_paths(layer13_uploads),
|
| 1209 |
+
_extract_paths(layer14_uploads),
|
| 1210 |
+
_extract_paths(layer15_uploads),
|
| 1211 |
+
_extract_paths(layer16_uploads),
|
| 1212 |
+
_extract_paths(layer17_uploads),
|
| 1213 |
+
_extract_paths(layer18_uploads),
|
| 1214 |
+
_extract_paths(layer19_uploads),
|
| 1215 |
+
_extract_paths(layer20_uploads),
|
| 1216 |
+
_extract_paths(layer21_uploads),
|
| 1217 |
+
_extract_paths(layer22_uploads),
|
| 1218 |
+
_extract_paths(layer23_uploads),
|
| 1219 |
+
_extract_paths(layer24_uploads),
|
| 1220 |
+
_extract_paths(layer25_uploads),
|
| 1221 |
+
_extract_paths(layer26_uploads),
|
| 1222 |
+
_extract_paths(layer27_uploads),
|
| 1223 |
+
_extract_paths(layer28_uploads),
|
| 1224 |
+
_extract_paths(layer29_uploads),
|
| 1225 |
+
_extract_paths(layer30_uploads),
|
| 1226 |
+
_extract_paths(layer31_uploads),
|
| 1227 |
+
_extract_paths(layer32_uploads),
|
| 1228 |
+
]
|
| 1229 |
+
if not layer_upload_sets[0]:
|
| 1230 |
+
log_buf += "[ERROR] Layer 1 (image_path_1) images are required.\n"
|
| 1231 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1232 |
+
return
|
| 1233 |
+
|
| 1234 |
+
layer_dirs: List[Optional[str]] = []
|
| 1235 |
+
for i, uploads in enumerate(layer_upload_sets):
|
| 1236 |
+
if not uploads:
|
| 1237 |
+
layer_dirs.append(None)
|
| 1238 |
+
continue
|
| 1239 |
+
folder_name = f"layer_{i + 1}"
|
| 1240 |
+
cdir = os.path.join(ds_dir, folder_name)
|
| 1241 |
+
os.makedirs(cdir, exist_ok=True)
|
| 1242 |
+
_copy_uploads(uploads, cdir)
|
| 1243 |
+
layer_dirs.append(folder_name)
|
| 1244 |
+
log_buf += f"[QIE] Copied {len(uploads)} layer_{i + 1} images to {cdir}\n"
|
| 1245 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1246 |
+
|
| 1247 |
+
layer_dirs_abs = [
|
| 1248 |
+
(os.path.join(ds_dir, name) if name else None)
|
| 1249 |
+
for name in layer_dirs
|
| 1250 |
+
]
|
| 1251 |
+
|
| 1252 |
+
ds_conf = str(Path(AUTO_DIR_RUNTIME) / "dataset_QIE.toml")
|
| 1253 |
+
try:
|
| 1254 |
+
_update_dataset_toml(
|
| 1255 |
+
ds_conf,
|
| 1256 |
+
img_res_w=int(train_res_w) if train_res_w else None,
|
| 1257 |
+
img_res_h=int(train_res_h) if train_res_h else None,
|
| 1258 |
+
train_batch_size=int(train_batch_size) if train_batch_size else None,
|
| 1259 |
+
multiple_target=True,
|
| 1260 |
+
)
|
| 1261 |
+
log_buf += f"[QIE] Updated dataset config: resolution=({train_res_w},{train_res_h}), batch_size={train_batch_size}, multiple_target=true\n"
|
| 1262 |
+
except Exception as e:
|
| 1263 |
+
log_buf += f"[QIE] WARN: failed to update dataset config: {e}\n"
|
| 1264 |
+
if os.path.isfile(ds_conf):
|
| 1265 |
+
artifacts = [ds_conf]
|
| 1266 |
+
|
| 1267 |
+
models_root = MODELS_ROOT_RUNTIME
|
| 1268 |
+
out_base = ds_dir
|
| 1269 |
+
try:
|
| 1270 |
+
os.makedirs(out_base, exist_ok=True)
|
| 1271 |
+
except Exception:
|
| 1272 |
+
pass
|
| 1273 |
+
|
| 1274 |
+
tmp_script = _prepare_script(
|
| 1275 |
+
dataset_name=ds_name,
|
| 1276 |
+
caption=caption,
|
| 1277 |
+
data_root=DATA_ROOT_RUNTIME,
|
| 1278 |
+
image_folder=img_folder_name,
|
| 1279 |
+
control_folders=[],
|
| 1280 |
+
models_root=models_root,
|
| 1281 |
+
output_dir_base=out_base,
|
| 1282 |
+
dataset_config=ds_conf,
|
| 1283 |
+
override_max_epochs=max_epochs if max_epochs and max_epochs > 0 else None,
|
| 1284 |
+
override_save_every=save_every if save_every and save_every > 0 else None,
|
| 1285 |
+
override_run_name=output_name.strip(),
|
| 1286 |
+
target_prefix=(target_prefix or ""),
|
| 1287 |
+
target_suffix=(target_suffix or ""),
|
| 1288 |
+
control_prefixes=[],
|
| 1289 |
+
control_suffixes=[],
|
| 1290 |
+
override_learning_rate=(learning_rate or None),
|
| 1291 |
+
override_network_dim=int(network_dim) if network_dim is not None else None,
|
| 1292 |
+
override_te_cache_bs=int(te_cache_batch_size) if te_cache_batch_size else None,
|
| 1293 |
+
override_seed=int(seed) if seed is not None else None,
|
| 1294 |
+
)
|
| 1295 |
+
|
| 1296 |
+
out_dir = os.path.join(out_base, output_name.strip())
|
| 1297 |
+
run_out_dir = out_dir
|
| 1298 |
+
ckpts = _list_checkpoints(out_dir)
|
| 1299 |
+
used_script_path = os.path.join(out_base, "train_QIE_used.sh")
|
| 1300 |
+
try:
|
| 1301 |
+
shutil.copy2(str(tmp_script), used_script_path)
|
| 1302 |
+
try:
|
| 1303 |
+
os.chmod(used_script_path, 0o755)
|
| 1304 |
+
except Exception:
|
| 1305 |
+
pass
|
| 1306 |
+
if used_script_path not in artifacts:
|
| 1307 |
+
artifacts.append(used_script_path)
|
| 1308 |
+
except Exception:
|
| 1309 |
+
pass
|
| 1310 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1311 |
+
|
| 1312 |
+
output_json = os.path.join(out_base, "metadata.jsonl")
|
| 1313 |
+
_sync_dataset_config_jsonl(ds_conf, output_json)
|
| 1314 |
+
log_buf += f"[QIE] Generating layered metadata: {output_json}\n"
|
| 1315 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1316 |
+
|
| 1317 |
+
layer_prefixes = [layer_prefix or None for _ in range(LAYER_MAX)]
|
| 1318 |
+
layer_suffixes = [layer_suffix or None for _ in range(LAYER_MAX)]
|
| 1319 |
+
try:
|
| 1320 |
+
count = _generate_layered_jsonl(
|
| 1321 |
+
image_dir=img_dir,
|
| 1322 |
+
caption=caption,
|
| 1323 |
+
output_json=output_json,
|
| 1324 |
+
control_dirs=layer_dirs_abs,
|
| 1325 |
+
target_prefix=(target_prefix or ""),
|
| 1326 |
+
target_suffix=(target_suffix or ""),
|
| 1327 |
+
control_prefixes=layer_prefixes,
|
| 1328 |
+
control_suffixes=layer_suffixes,
|
| 1329 |
+
allow_single=True,
|
| 1330 |
+
)
|
| 1331 |
+
log_buf += f"[QIE] Layered metadata written: {count} entries\n"
|
| 1332 |
+
except Exception as e:
|
| 1333 |
+
log_buf += f"[ERROR] Layered metadata generation failed: {e}\n"
|
| 1334 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1335 |
+
return
|
| 1336 |
+
|
| 1337 |
+
if os.path.isfile(output_json) and output_json not in artifacts:
|
| 1338 |
+
artifacts.append(output_json)
|
| 1339 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1340 |
+
|
| 1341 |
+
if config_only:
|
| 1342 |
+
log_buf += "[QIE] Config-only mode: skipping cache/training.\n"
|
| 1343 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1344 |
+
return
|
| 1345 |
+
|
| 1346 |
+
shell = _pick_shell()
|
| 1347 |
+
log_buf += f"[QIE] Using shell: {shell}\n"
|
| 1348 |
+
log_buf += f"[QIE] Running script: {tmp_script}\n"
|
| 1349 |
+
|
| 1350 |
+
child_env = os.environ.copy()
|
| 1351 |
+
child_env["PYTHONUNBUFFERED"] = "1"
|
| 1352 |
+
child_env["PYTHONIOENCODING"] = "utf-8"
|
| 1353 |
+
|
| 1354 |
+
proc = subprocess.Popen(
|
| 1355 |
+
[shell, str(tmp_script)],
|
| 1356 |
+
stdout=subprocess.PIPE,
|
| 1357 |
+
stderr=subprocess.STDOUT,
|
| 1358 |
+
text=True,
|
| 1359 |
+
bufsize=1,
|
| 1360 |
+
universal_newlines=True,
|
| 1361 |
+
env=child_env,
|
| 1362 |
+
)
|
| 1363 |
+
try:
|
| 1364 |
+
assert proc.stdout is not None
|
| 1365 |
+
i = 0
|
| 1366 |
+
for line in proc.stdout:
|
| 1367 |
+
log_buf += line
|
| 1368 |
+
i += 1
|
| 1369 |
+
if i % 30 == 0:
|
| 1370 |
+
ckpts = _list_checkpoints(out_dir)
|
| 1371 |
+
metadata_json = os.path.join(out_base, "metadata.jsonl")
|
| 1372 |
+
if os.path.isfile(metadata_json) and metadata_json not in artifacts:
|
| 1373 |
+
artifacts.append(metadata_json)
|
| 1374 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1375 |
+
finally:
|
| 1376 |
+
code = proc.wait()
|
| 1377 |
+
try:
|
| 1378 |
+
ckpts = _list_checkpoints(out_dir)
|
| 1379 |
+
except Exception:
|
| 1380 |
+
pass
|
| 1381 |
+
log_buf += f"[QIE] Exit code: {code}\n"
|
| 1382 |
+
metadata_json = os.path.join(out_base, "metadata.jsonl")
|
| 1383 |
+
if os.path.isfile(metadata_json) and metadata_json not in artifacts:
|
| 1384 |
+
artifacts.append(metadata_json)
|
| 1385 |
+
yield (log_buf, ckpts, artifacts, run_out_dir)
|
| 1386 |
+
|
| 1387 |
+
|
| 1388 |
def build_ui() -> gr.Blocks:
|
| 1389 |
css = """
|
| 1390 |
.pad-section {
|
|
|
|
| 1410 |
}
|
| 1411 |
"""
|
| 1412 |
with gr.Blocks(title="Qwen-Image-Edit: Trainer", css=css) as demo:
|
| 1413 |
+
# 回収ボタン: 直近の dataset_ ディレクトリからチェックポイントとスクリプト/設定を再取得
|
| 1414 |
+
def _refresh_all() -> tuple:
|
| 1415 |
+
try:
|
| 1416 |
+
ds_dir = _find_latest_dataset_dir(DATA_ROOT_RUNTIME)
|
| 1417 |
+
except Exception:
|
| 1418 |
+
ds_dir = None
|
| 1419 |
+
try:
|
| 1420 |
+
ck = _list_checkpoints(ds_dir) if ds_dir else []
|
| 1421 |
+
except Exception:
|
| 1422 |
+
ck = []
|
| 1423 |
+
try:
|
| 1424 |
+
sc = _collect_scripts_and_config(ds_dir)
|
| 1425 |
+
except Exception:
|
| 1426 |
+
sc = _collect_scripts_and_config(None)
|
| 1427 |
+
return ck, sc
|
| 1428 |
+
|
| 1429 |
with gr.Tabs() as tabs:
|
| 1430 |
+
with gr.TabItem("Edit"):
|
| 1431 |
gr.Markdown("""
|
| 1432 |
# Qwen-Image-Edit Trainer
|
| 1433 |
学習に使う画像をアップロードし、必要ならファイル名の前後にある共通の文字(prefix/suffix)を指定して、
|
|
|
|
| 1614 |
outputs=[logs, ckpt_files, scripts_files, run_out_dir_box],
|
| 1615 |
)
|
| 1616 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1617 |
refresh_scripts_btn.click(
|
| 1618 |
fn=_refresh_all,
|
| 1619 |
inputs=[],
|
| 1620 |
outputs=[ckpt_files, scripts_files],
|
| 1621 |
)
|
| 1622 |
|
| 1623 |
+
with gr.TabItem("Layered"):
|
| 1624 |
+
gr.Markdown("""
|
| 1625 |
+
# Qwen-Image-Layered Trainer
|
| 1626 |
+
このタブは `QWEN_IMAGE_TYPE=layered` で起動したときのみ有効です。
|
| 1627 |
+
""")
|
| 1628 |
+
|
| 1629 |
+
with gr.Accordion("Settings", elem_classes=["pad-section"]):
|
| 1630 |
+
with gr.Group():
|
| 1631 |
+
with gr.Row():
|
| 1632 |
+
output_name_layered = gr.Textbox(label="OUTPUT NAME", placeholder="my_lora_output", lines=1)
|
| 1633 |
+
caption_layered = gr.Textbox(label="CAPTION", placeholder="A photo of ...", lines=2)
|
| 1634 |
+
|
| 1635 |
+
with gr.Row():
|
| 1636 |
+
lr_input_layered = gr.Textbox(label="Learning rate", value="1e-3")
|
| 1637 |
+
dim_input_layered = gr.Number(label="Network dim", value=4, precision=0)
|
| 1638 |
+
train_bs_layered = gr.Number(label="Batch size (dataset)", value=1, precision=0)
|
| 1639 |
+
seed_input_layered = gr.Number(label="Seed", value=42, precision=0)
|
| 1640 |
+
max_epochs_layered = gr.Number(label="Max epochs", value=100, precision=0)
|
| 1641 |
+
save_every_layered = gr.Number(label="Save every N epochs", value=10, precision=0)
|
| 1642 |
+
|
| 1643 |
+
with gr.Row():
|
| 1644 |
+
tr_w_layered = gr.Number(label="Image resolution W", value=1024, precision=0)
|
| 1645 |
+
tr_h_layered = gr.Number(label="Image resolution H", value=1024, precision=0)
|
| 1646 |
+
te_bs_layered = gr.Number(label="TE cache batch size", value=16, precision=0)
|
| 1647 |
+
|
| 1648 |
+
with gr.Accordion("Base Image (image_path_0)", elem_classes=["pad-section_0"]):
|
| 1649 |
+
with gr.Group():
|
| 1650 |
+
with gr.Row():
|
| 1651 |
+
base_images_input = gr.File(label="Upload base images (image_path_0)", file_count="multiple", type="filepath", height=220, scale=3)
|
| 1652 |
+
base_gallery = gr.Gallery(label="Base preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
|
| 1653 |
+
with gr.Column(scale=1):
|
| 1654 |
+
with gr.Row():
|
| 1655 |
+
base_prefix = gr.Textbox(label="Base prefix", placeholder="e.g., IMG_")
|
| 1656 |
+
base_suffix = gr.Textbox(label="Base suffix", placeholder="e.g., _v2")
|
| 1657 |
+
with gr.Row():
|
| 1658 |
+
layer_prefix = gr.Textbox(label="Layer prefix (all)", placeholder="")
|
| 1659 |
+
layer_suffix = gr.Textbox(label="Layer suffix (all)", placeholder="")
|
| 1660 |
+
with gr.Accordion("prefix/sufixについて", open=False):
|
| 1661 |
+
gr.Markdown("""
|
| 1662 |
+
ファイル名の対応付けのルール:
|
| 1663 |
+
- base画像のファイル名から Base prefix/suffix を取り除いたものを key とします。
|
| 1664 |
+
- レイヤーは `Layer prefix + key + Layer suffix + .png` を探します。
|
| 1665 |
+
- レイヤーが1枚のみのときは全ベース画像に適用します。
|
| 1666 |
+
""")
|
| 1667 |
+
|
| 1668 |
+
layer_files: List[gr.File] = []
|
| 1669 |
+
with gr.Accordion("Layer Images (image_path_1..32)", elem_classes=["pad-section_1"]):
|
| 1670 |
+
with gr.Group():
|
| 1671 |
+
for i in range(1, LAYER_MAX + 1, 2):
|
| 1672 |
+
with gr.Row():
|
| 1673 |
+
lf1 = gr.File(
|
| 1674 |
+
label=f"Layer {i} (image_path_{i})",
|
| 1675 |
+
file_count="multiple",
|
| 1676 |
+
type="filepath",
|
| 1677 |
+
height=120,
|
| 1678 |
+
scale=1,
|
| 1679 |
+
)
|
| 1680 |
+
layer_files.append(lf1)
|
| 1681 |
+
if i + 1 <= LAYER_MAX:
|
| 1682 |
+
lf2 = gr.File(
|
| 1683 |
+
label=f"Layer {i + 1} (image_path_{i + 1})",
|
| 1684 |
+
file_count="multiple",
|
| 1685 |
+
type="filepath",
|
| 1686 |
+
height=120,
|
| 1687 |
+
scale=1,
|
| 1688 |
+
)
|
| 1689 |
+
layer_files.append(lf2)
|
| 1690 |
+
|
| 1691 |
+
with gr.Row():
|
| 1692 |
+
run_btn_layered = gr.Button("Start Training", variant="primary")
|
| 1693 |
+
config_btn_layered = gr.Button("設定のみ生成", variant="secondary")
|
| 1694 |
+
run_out_dir_box_layered = gr.Textbox(label="出力フォルダ", lines=1, interactive=False)
|
| 1695 |
+
scripts_files_layered = gr.Files(label="Scripts & Config (live)", interactive=False)
|
| 1696 |
+
ckpt_files_layered = gr.Files(label="Checkpoints (live)", interactive=False)
|
| 1697 |
+
logs_layered = gr.Textbox(label="Logs", lines=20)
|
| 1698 |
+
with gr.Row():
|
| 1699 |
+
refresh_scripts_btn_layered = gr.Button("ファイルを再取得", variant="secondary")
|
| 1700 |
+
|
| 1701 |
+
base_images_input.change(fn=_files_to_gallery, inputs=base_images_input, outputs=base_gallery)
|
| 1702 |
+
|
| 1703 |
+
config_only_off_layered = gr.State(False)
|
| 1704 |
+
config_only_on_layered = gr.State(True)
|
| 1705 |
+
|
| 1706 |
+
run_btn_layered.click(
|
| 1707 |
+
fn=run_training_layered,
|
| 1708 |
+
inputs=[
|
| 1709 |
+
output_name_layered, caption_layered, base_images_input, base_prefix, base_suffix,
|
| 1710 |
+
layer_prefix, layer_suffix,
|
| 1711 |
+
*layer_files,
|
| 1712 |
+
lr_input_layered, dim_input_layered,
|
| 1713 |
+
tr_w_layered, tr_h_layered, train_bs_layered, te_bs_layered,
|
| 1714 |
+
seed_input_layered, max_epochs_layered, save_every_layered, config_only_off_layered,
|
| 1715 |
+
],
|
| 1716 |
+
outputs=[logs_layered, ckpt_files_layered, scripts_files_layered, run_out_dir_box_layered],
|
| 1717 |
+
)
|
| 1718 |
+
config_btn_layered.click(
|
| 1719 |
+
fn=run_training_layered,
|
| 1720 |
+
inputs=[
|
| 1721 |
+
output_name_layered, caption_layered, base_images_input, base_prefix, base_suffix,
|
| 1722 |
+
layer_prefix, layer_suffix,
|
| 1723 |
+
*layer_files,
|
| 1724 |
+
lr_input_layered, dim_input_layered,
|
| 1725 |
+
tr_w_layered, tr_h_layered, train_bs_layered, te_bs_layered,
|
| 1726 |
+
seed_input_layered, max_epochs_layered, save_every_layered, config_only_on_layered,
|
| 1727 |
+
],
|
| 1728 |
+
outputs=[logs_layered, ckpt_files_layered, scripts_files_layered, run_out_dir_box_layered],
|
| 1729 |
+
)
|
| 1730 |
+
|
| 1731 |
+
refresh_scripts_btn_layered.click(
|
| 1732 |
+
fn=_refresh_all,
|
| 1733 |
+
inputs=[],
|
| 1734 |
+
outputs=[ckpt_files_layered, scripts_files_layered],
|
| 1735 |
+
)
|
| 1736 |
+
|
| 1737 |
with gr.TabItem("Prompt Generator"):
|
| 1738 |
gr.Markdown("""
|
| 1739 |
# 🎨 A→B 変換プロンプト自動生成
|
|
|
|
| 1800 |
_startup_install_musubi_deps()
|
| 1801 |
|
| 1802 |
# 2) Download models at startup (blocking by design)
|
| 1803 |
+
# _startup_download_models()
|
| 1804 |
|
| 1805 |
# 3) Launch Gradio app
|
| 1806 |
ui = build_ui()
|