yeq6x commited on
Commit
d36832a
·
1 Parent(s): 73d06d9

Add layered training functionality and enhance error handling in app.py

Browse files

Introduce 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.

Files changed (1) hide show
  1. app.py +435 -78
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
- if image_type == "layered":
887
- if not any(control_upload_sets):
888
- log_buf += "[ERROR] Layered では control_0〜7 のいずれかが必要です。\n"
889
- yield (log_buf, ckpts, artifacts, run_out_dir)
890
- return
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=ctrl_w,
947
- control_res_h=ctrl_h,
948
- multiple_target=(image_type == "layered"),
949
  )
950
- if image_type == "layered":
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("Training"):
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()