yeq6x's picture
Remove default branch assignment for layered image type in _resolve_musubi_target function in app.py. This change simplifies branch handling by eliminating the automatic assignment of a specific branch when no branch is set, enhancing clarity in target resolution logic.
10a4f05
Raw
History Blame Contribute Delete
86 kB
#!/usr/bin/env python3
import os
import re
import shutil
import subprocess
import sys
import tempfile
from pathlib import Path
from typing import Dict, Iterable, List, Optional, Any, Tuple
import time
import json
import gradio as gr
import importlib
import spaces
# Local modules
from download_qwen_image_models import download_all_models, DEFAULT_MODELS_DIR
# Defaults matching train_QIE.sh expectations
DEFAULT_DATA_ROOT = "/data"
DEFAULT_IMAGE_FOLDER = "image"
DEFAULT_OUTPUT_DIR_BASE = "/auto/train_LoRA"
DEFAULT_DATASET_CONFIG = "/auto/dataset_QIE.toml"
DEFAULT_MODELS_ROOT = DEFAULT_MODELS_DIR # "/Qwen-Image_models"
WORKSPACE_AUTO_DIR = "/auto"
# musubi-tuner settings
DEFAULT_MUSUBI_TUNER_DIR = os.environ.get("MUSUBI_TUNER_DIR", "/musubi-tuner")
DEFAULT_MUSUBI_TUNER_REPO = os.environ.get(
"MUSUBI_TUNER_REPO", "https://github.com/kohya-ss/musubi-tuner.git"
)
TRAINING_DIR = Path(__file__).resolve().parent
# Runtime-resolved paths with fallbacks for non-root environments
MUSUBI_TUNER_DIR_RUNTIME = DEFAULT_MUSUBI_TUNER_DIR
MODELS_ROOT_RUNTIME = DEFAULT_MODELS_ROOT
AUTO_DIR_RUNTIME = WORKSPACE_AUTO_DIR
DATA_ROOT_RUNTIME = DEFAULT_DATA_ROOT
def _bash_quote(s: str) -> str:
"""Return a POSIX-safe single-quoted string literal representing s."""
if s is None:
return "''"
return "'" + str(s).replace("'", "'\"'\"'") + "'"
_QWEN_IMAGE_TYPES = ("edit-2509", "edit-2511", "layered")
EDIT_CONTROL_MAX = 8
LAYER_MAX = 32
def _get_qwen_image_type() -> str:
raw = os.environ.get("QWEN_IMAGE_TYPE", "layered")
v = raw.strip().lower()
if v not in _QWEN_IMAGE_TYPES:
print(f"[QIE] Unsupported QWEN_IMAGE_TYPE: {v}. Defaulting to edit-2509.")
return "edit-2509"
return v
def _resolve_musubi_target(image_type: str) -> Tuple[str, str]:
env_dir = os.environ.get("MUSUBI_TUNER_DIR", "").strip()
if env_dir:
target = env_dir
else:
target = "/musubi-tuner-layered" if image_type == "layered" else "/musubi-tuner"
branch = os.environ.get("MUSUBI_TUNER_BRANCH", "").strip()
return target, branch
def _dit_filename_for_type(image_type: str) -> str:
if image_type == "layered":
return "qwen_image_layered_bf16.safetensors"
suffix = image_type.split("-", 1)[1]
return f"qwen_image_edit_{suffix}_bf16.safetensors"
def _vae_filename_for_type(image_type: str) -> str:
if image_type == "layered":
return "qwen_image_layered_vae.safetensors"
return "diffusion_pytorch_model.safetensors"
def _ensure_workspace_auto_files() -> None:
"""Ensure /workspace/auto has required helper files from this repo.
Copies training/create_image_caption_json.py and training/dataset_QIE.toml
into /workspace/auto so that train_QIE.sh can run unmodified.
"""
global AUTO_DIR_RUNTIME
try:
os.makedirs(AUTO_DIR_RUNTIME, exist_ok=True)
except PermissionError:
home_auto = os.path.join(os.path.expanduser("~"), "auto")
os.makedirs(home_auto, exist_ok=True)
AUTO_DIR_RUNTIME = home_auto # type: ignore
src_py = TRAINING_DIR / "create_image_caption_json.py"
src_toml = TRAINING_DIR / "dataset_QIE.toml"
dst_py = Path(AUTO_DIR_RUNTIME) / "create_image_caption_json.py"
dst_toml = Path(AUTO_DIR_RUNTIME) / "dataset_QIE.toml"
try:
shutil.copy2(src_py, dst_py)
except Exception:
pass
try:
if src_toml.exists():
shutil.copy2(src_toml, dst_toml)
except Exception:
pass
def _update_dataset_toml(
path: str,
*,
img_res_w: Optional[int] = None,
img_res_h: Optional[int] = None,
train_batch_size: Optional[int] = None,
control_res_w: Optional[int] = None,
control_res_h: Optional[int] = None,
multiple_target: Optional[bool] = None,
remove_multiple_target: bool = False,
) -> None:
"""Update dataset TOML for resolution/batch/control resolution in-place.
- Updates [general] resolution and batch_size if provided.
- Updates first [[datasets]] qwen_image_edit_control_resolution if provided.
- Creates sections/keys if missing.
"""
try:
txt = Path(path).read_text(encoding="utf-8")
except Exception:
return
def _set_in_general(block: str, key: str, value_line: str) -> str:
import re as _re
if _re.search(rf"(?m)^\s*{_re.escape(key)}\s*=", block):
block = _re.sub(rf"(?m)^\s*{_re.escape(key)}\s*=.*$", value_line, block)
else:
block = block.rstrip() + "\n" + value_line + "\n"
return block
import re
m = re.search(r"(?ms)^\[general\]\s*(.*?)(?=^\[|\Z)", txt)
if not m:
gen = "[general]\n"
if img_res_w and img_res_h:
gen += f"resolution = [{int(img_res_w)}, {int(img_res_h)}]\n"
if train_batch_size is not None:
gen += f"batch_size = {int(train_batch_size)}\n"
txt = gen + "\n" + txt
else:
head, block, tail = txt[:m.start(1)], m.group(1), txt[m.end(1):]
if img_res_w and img_res_h:
block = _set_in_general(block, "resolution", f"resolution = [{int(img_res_w)}, {int(img_res_h)}]")
if train_batch_size is not None:
block = _set_in_general(block, "batch_size", f"batch_size = {int(train_batch_size)}")
txt = head + block + tail
m2 = re.search(r"(?ms)^\[\[datasets\]\]\s*(.*?)(?=^\[\[|\Z)", txt)
if m2:
head, block, tail = txt[:m2.start(1)], m2.group(1), txt[m2.end(1):]
if control_res_w and control_res_h:
line = f"qwen_image_edit_control_resolution = [{int(control_res_w)}, {int(control_res_h)}]"
if re.search(r"(?m)^\s*qwen_image_edit_control_resolution\s*=", block):
block = re.sub(r"(?m)^\s*qwen_image_edit_control_resolution\s*=.*$", line, block)
else:
block = block.rstrip() + "\n" + line + "\n"
if remove_multiple_target:
block = re.sub(r"(?m)^\s*multiple_target\s*=.*$\n?", "", block)
elif multiple_target is not None:
mt_line = f"multiple_target = {'true' if multiple_target else 'false'}"
if re.search(r"(?m)^\s*multiple_target\s*=", block):
block = re.sub(r"(?m)^\s*multiple_target\s*=.*$", mt_line, block)
else:
block = block.rstrip() + "\n" + mt_line + "\n"
txt = head + block + tail
try:
Path(path).write_text(txt, encoding="utf-8")
except Exception:
pass
def _sync_dataset_config_jsonl(path: str, output_json: str) -> None:
"""Ensure dataset TOML points to the generated JSONL and a local cache dir."""
try:
txt = Path(path).read_text(encoding="utf-8")
except Exception:
return
base = os.path.dirname(path)
cache = os.path.join(base, "cache").replace("\\", "/")
image_line = f'image_jsonl_file = "{output_json}"'
new = re.sub(r"(?m)^\s*image_jsonl_file\s*=.*$", lambda _m, r=image_line: r, txt)
if new == txt and "image_jsonl_file" not in txt:
new = txt.rstrip("\n") + f'\nimage_jsonl_file = "{output_json}"\n'
if re.search(r"(?m)^\s*cache_directory\s*=", new):
cache_line = f'cache_directory = "{cache}"'
new = re.sub(r"(?m)^\s*cache_directory\s*=.*$", lambda _m, r=cache_line: r, new)
else:
new = new.rstrip("\n") + f'\ncache_directory = "{cache}"\n'
try:
Path(path).write_text(new, encoding="utf-8")
os.makedirs(os.path.join(base, "cache"), exist_ok=True)
except Exception:
pass
def _ensure_dir_writable(path: str) -> str:
try:
os.makedirs(path, exist_ok=True)
return path
except PermissionError:
home_path = os.path.join(os.path.expanduser("~"), os.path.basename(path.strip("/\\")))
os.makedirs(home_path, exist_ok=True)
return home_path
def _ensure_data_root(candidate: Optional[str]) -> str:
root = (candidate or DEFAULT_DATA_ROOT).strip() or DEFAULT_DATA_ROOT
try:
os.makedirs(root, exist_ok=True)
return root
except PermissionError:
home_root = os.path.join(os.path.expanduser("~"), "data")
os.makedirs(home_root, exist_ok=True)
return home_root
def _extract_paths(files: Any) -> List[Tuple[str, str]]:
"""Extract a list of (abs_path, orig_basename) from Gradio Files input.
Supports various gradio return shapes across versions.
"""
out: List[Tuple[str, str]] = []
if not files:
return out
# Gradio Files often returns a list
if isinstance(files, (list, tuple)):
items = files
else:
items = [files]
for item in items:
p: Optional[str] = None
orig: Optional[str] = None
# dict-like
if isinstance(item, dict):
p = item.get("path") or item.get("name") or item.get("file")
orig = item.get("orig_name") or item.get("name")
else:
# object with attributes
p = getattr(item, "name", None) or getattr(item, "path", None) or str(item)
# best-effort original name attribute
orig = getattr(item, "orig_name", None) or os.path.basename(p) if p else None
if p:
abs_p = os.path.abspath(p)
out.append((abs_p, os.path.basename(orig or abs_p)))
return out
def _norm_key(filename: str, prefix: str, suffix: str) -> str:
stem = os.path.splitext(os.path.basename(filename))[0]
if prefix and stem.startswith(prefix):
stem = stem[len(prefix):]
if suffix and stem.endswith(suffix):
stem = stem[: -len(suffix)]
return stem
IMAGE_EXTENSIONS = (".jpg", ".jpeg", ".png", ".bmp", ".gif")
def _list_image_files(folder: str) -> List[str]:
try:
files = [
fn
for fn in sorted(os.listdir(folder))
if fn.lower().endswith(IMAGE_EXTENSIONS)
and os.path.isfile(os.path.join(folder, fn))
]
return files
except Exception:
return []
def _join_posix(base: str, name: str) -> str:
if not base:
return name.replace("\\", "/")
return base.rstrip("/\\") + "/" + name.replace("\\", "/").lstrip("/")
def _generate_layered_jsonl(
image_dir: str,
caption: str,
output_json: str,
control_dirs: List[Optional[str]],
target_prefix: str = "",
target_suffix: str = "",
control_prefixes: Optional[List[Optional[str]]] = None,
control_suffixes: Optional[List[Optional[str]]] = None,
allow_single: bool = True,
) -> int:
base_files = _list_image_files(image_dir)
if not base_files:
raise ValueError("IMAGE_FOLDER に画像が見つかりません。")
active_controls = [d for d in control_dirs if d]
if not active_controls:
raise ValueError("Layered では少なくとも1つのレイヤー画像が必要です。")
use_name_matching = bool(target_prefix or target_suffix)
if control_prefixes and any((c or "").strip() for c in control_prefixes):
use_name_matching = True
if control_suffixes and any((c or "").strip() for c in control_suffixes):
use_name_matching = True
if not use_name_matching:
missing: List[str] = []
for fname in base_files:
for cdir in active_controls:
expected = os.path.join(cdir, fname)
if not os.path.exists(expected):
missing.append(expected)
if missing:
preview = "\n".join(f" - {p}" for p in missing[:5])
raise ValueError(f"対応するレイヤーが見つかりません(strict):\n{preview}")
out_dir = os.path.dirname(output_json)
if out_dir:
os.makedirs(out_dir, exist_ok=True)
count = 0
with open(output_json, "w", encoding="utf-8") as f:
for base_fn in base_files:
entry: Dict[str, str] = {
"image_path_0": _join_posix(image_dir, base_fn),
"caption": caption,
}
key = _norm_key(base_fn, target_prefix, target_suffix)
layer_index = 1
for i, cdir in enumerate(control_dirs):
if not cdir:
continue
cprefix = ""
csuffix = ""
if control_prefixes and i < len(control_prefixes) and control_prefixes[i]:
cprefix = control_prefixes[i]
if control_suffixes and i < len(control_suffixes) and control_suffixes[i]:
csuffix = control_suffixes[i]
expected_name = base_fn if not use_name_matching else f"{cprefix}{key}{csuffix}.png"
expected_path = os.path.join(cdir, expected_name)
if not os.path.exists(expected_path):
cfiles = _list_image_files(cdir)
if allow_single and len(cfiles) == 1:
expected_name = cfiles[0]
else:
raise ValueError(
f"対応するレイヤーが見つかりません: expected={expected_name}, layer_index={i}"
)
entry[f"image_path_{layer_index}"] = _join_posix(cdir, expected_name)
layer_index += 1
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
count += 1
return count
def _copy_uploads(
uploads: List[Tuple[str, str]],
dest_dir: str,
rename_to: Optional[List[str]] = None,
force_rgba: bool = False,
) -> List[str]:
os.makedirs(dest_dir, exist_ok=True)
used_names: List[str] = []
for idx, (src, orig) in enumerate(uploads):
# Determine target stem
if rename_to and idx < len(rename_to):
stem = os.path.splitext(rename_to[idx])[0]
else:
stem = os.path.splitext(orig)[0]
dst_name = f"{stem}.png"
# ensure unique within this batch
final_name = dst_name
dup_idx = 1
while final_name in used_names:
final_name = f"{stem}_{dup_idx}.png"
dup_idx += 1
dst_path = os.path.join(dest_dir, final_name)
# Convert to PNG during save
try:
try:
from PIL import Image # type: ignore
with Image.open(src) as img:
if force_rgba and img.mode != "RGBA":
img = img.convert("RGBA")
img.save(dst_path, format="PNG")
except Exception:
# Fallback: copy then rename
shutil.copy2(src, dst_path)
except Exception:
# Last resort
shutil.copy(src, dst_path)
used_names.append(final_name)
return used_names
def _list_checkpoints(out_dir: str, limit: int = 20) -> List[str]:
try:
if not out_dir or not os.path.isdir(out_dir):
return []
import time
now = time.time()
min_age_sec = 3.0 # treat files newer than this as possibly in-flight
items: List[Tuple[float, str]] = []
for root, _, files in os.walk(out_dir):
for fn in files:
if fn.lower().endswith('.safetensors'):
full = os.path.join(root, fn)
try:
# Skip zero-length, too-new, or unreadable files (likely in-flight)
size = os.path.getsize(full)
if size <= 0:
continue
mtime = os.path.getmtime(full)
if (now - mtime) < min_age_sec:
continue
# Try opening a small read to ensure readability
with open(full, 'rb') as rf:
rf.read(64)
items.append((mtime, full))
except Exception:
pass
items.sort(reverse=True)
return [p for _, p in items[:limit]]
except Exception:
return []
def _find_latest_dataset_dir(root: str) -> Optional[str]:
try:
if not os.path.isdir(root):
return None
cand: List[Tuple[float, str]] = []
for name in os.listdir(root):
if not name.startswith("dataset_"):
continue
full = os.path.join(root, name)
if os.path.isdir(full):
try:
cand.append((os.path.getmtime(full), full))
except Exception:
pass
if not cand:
return None
cand.sort(reverse=True)
return cand[0][1]
except Exception:
return None
def _collect_scripts_and_config(ds_dir: Optional[str]) -> List[str]:
files: List[str] = []
try:
ds_conf = str(Path(AUTO_DIR_RUNTIME) / "dataset_QIE.toml")
if os.path.isfile(ds_conf):
files.append(ds_conf)
if ds_dir and os.path.isdir(ds_dir):
used_script = os.path.join(ds_dir, "train_QIE_used.sh")
if os.path.isfile(used_script):
files.append(used_script)
meta = os.path.join(ds_dir, "metadata.jsonl")
if os.path.isfile(meta):
files.append(meta)
except Exception:
pass
return files
def _files_to_gallery(files: Any) -> List[str]:
items: List[str] = []
if not files:
return items
seq = files if isinstance(files, (list, tuple)) else [files]
for f in seq:
p = None
if isinstance(f, str):
p = f
elif isinstance(f, dict):
p = f.get("path") or f.get("name")
else:
p = getattr(f, "path", None) or getattr(f, "name", None)
if p:
items.append(p)
return items
def _prepare_script(
dataset_name: str,
caption: str,
data_root: str,
image_folder: str,
control_folders: List[Optional[str]],
models_root: str,
output_dir_base: Optional[str] = None,
dataset_config: Optional[str] = None,
override_max_epochs: Optional[int] = None,
override_save_every: Optional[int] = None,
override_run_name: Optional[str] = None,
target_prefix: Optional[str] = None,
target_suffix: Optional[str] = None,
control_prefixes: Optional[List[Optional[str]]] = None,
control_suffixes: Optional[List[Optional[str]]] = None,
override_learning_rate: Optional[str] = None,
override_network_dim: Optional[int] = None,
override_seed: Optional[int] = None,
override_te_cache_bs: Optional[int] = None,
) -> Path:
"""Create a temporary copy of train_QIE.sh with injected variables.
Only variables that must vary per-run are replaced. The rest of the script
remains as-is to preserve behavior.
"""
src = TRAINING_DIR / "train_QIE.sh"
txt = src.read_text(encoding="utf-8")
# Replace core variables
replacements = {
r"^DATA_ROOT=\".*\"": f"DATA_ROOT={_bash_quote(data_root)}",
r"^DATASET_NAME=\".*\"": f"DATASET_NAME={_bash_quote(dataset_name)}",
r"^CAPTION=\".*\"": f"CAPTION={_bash_quote(caption)}",
r"^IMAGE_FOLDER=\".*\"": f"IMAGE_FOLDER={_bash_quote(image_folder)}",
}
if output_dir_base:
replacements[r"^OUTPUT_DIR_BASE=\".*\""] = (
f"OUTPUT_DIR_BASE={_bash_quote(output_dir_base)}"
)
if dataset_config:
replacements[r"^DATASET_CONFIG=\".*\""] = (
f"DATASET_CONFIG={_bash_quote(dataset_config)}"
)
for pat, val in replacements.items():
txt = re.sub(pat, lambda _m, v=val: v, txt, flags=re.MULTILINE)
# Inject CONTROL_FOLDER_i if provided (uncomment/override or append)
for i in range(8):
val = control_folders[i] if i < len(control_folders) else None
if not val:
continue
# Try to replace commented placeholder first
pattern = rf"^#\s*CONTROL_FOLDER_{i}=\".*\""
if re.search(pattern, txt, flags=re.MULTILINE):
txt = re.sub(
pattern,
f"CONTROL_FOLDER_{i}={_bash_quote(val)}",
txt,
flags=re.MULTILINE,
)
else:
# Append after IMAGE_FOLDER definition
txt = re.sub(
r"^(IMAGE_FOLDER=.*)$",
rf"\1\nCONTROL_FOLDER_{i}={_bash_quote(val)}",
txt,
count=1,
flags=re.MULTILINE,
)
# Point model paths to the selected models_root
def _replace_model_path(txt: str, key: str, rel: str) -> str:
repl = f"--{key} \"{models_root.rstrip('/')}/{rel}\""
return re.sub(
rf"--{key} \"[^\"]+\"",
lambda _m, r=repl: r,
txt,
)
image_type = _get_qwen_image_type()
dit_filename = _dit_filename_for_type(image_type)
vae_filename = _vae_filename_for_type(image_type)
txt = _replace_model_path(txt, "vae", f"vae/{vae_filename}")
txt = _replace_model_path(txt, "text_encoder", "text_encoder/qwen_2.5_vl_7b.safetensors")
txt = _replace_model_path(txt, "dit", f"dit/{dit_filename}")
# Replace working dir for metadata generation to runtime /auto
txt = re.sub(
r"^cd\s+/workspace/auto\s*$",
lambda _m: f"cd {AUTO_DIR_RUNTIME}",
txt,
flags=re.MULTILINE,
)
# Ensure musubi-tuner path matches runtime location
txt = re.sub(
r"^cd\s+/musubi-tuner\s*$",
lambda _m: f"cd {MUSUBI_TUNER_DIR_RUNTIME}",
txt,
flags=re.MULTILINE,
)
# ZeroGPU compatibility: avoid spawning via 'accelerate launch'.
# Run the training module directly in-process so GPU stays attached
# to the same Python request context.
txt = re.sub(
r"\baccelerate\s+launch\s+src/musubi_tuner/qwen_image_train_network.py",
r"python -u src/musubi_tuner/qwen_image_train_network.py",
txt,
flags=re.MULTILINE,
)
# Optionally override epochs and save frequency for ZeroGPU time slicing
if override_max_epochs is not None and override_max_epochs > 0:
txt = re.sub(r"--max_train_epochs\s+\d+",
f"--max_train_epochs {override_max_epochs}", txt)
if override_save_every is not None and override_save_every > 0:
txt = re.sub(r"--save_every_n_epochs\s+\d+",
f"--save_every_n_epochs {override_save_every}", txt)
if override_run_name:
repl = f"RUN_NAME={_bash_quote(override_run_name)}"
txt = re.sub(r"^RUN_NAME=.*$", lambda _m, r=repl: r, txt, flags=re.MULTILINE)
# Inject prefix/suffix flags for metadata creation
extra_lines: List[str] = []
if (target_prefix or ""):
extra_lines.append(f" --target_prefix {_bash_quote(target_prefix)} \\")
if (target_suffix or ""):
extra_lines.append(f" --target_suffix {_bash_quote(target_suffix)} \\")
for i in range(8):
pre = control_prefixes[i] if (control_prefixes and i < len(control_prefixes)) else None
suf = control_suffixes[i] if (control_suffixes and i < len(control_suffixes)) else None
if pre:
extra_lines.append(f" --control_prefix_{i} {_bash_quote(pre)} \\")
if suf:
extra_lines.append(f" --control_suffix_{i} {_bash_quote(suf)} \\")
if extra_lines:
extra_block = "\n".join(extra_lines)
# Insert extra flags just before the CONTROL_ARGS line, preserving indentation.
txt = re.sub(
r'^(\s*)"\$\{CONTROL_ARGS\[@\]\}"',
lambda m: f"{extra_block}\n{m.group(1)}\"${{CONTROL_ARGS[@]}}\"",
txt,
flags=re.MULTILINE,
)
# Override CLI hyperparameters if provided
if override_learning_rate:
txt = re.sub(r"--learning_rate\s+[-+eE0-9\.]+", f"--learning_rate {override_learning_rate}", txt)
if override_network_dim is not None:
txt = re.sub(r"--network_dim\s+\d+", f"--network_dim {override_network_dim}", txt)
if override_seed is not None:
txt = re.sub(r"--seed\s+\d+", f"--seed {override_seed}", txt)
# Optionally override text-encoder cache batch size
if override_te_cache_bs is not None and override_te_cache_bs > 0:
txt = re.sub(
r"(qwen_image_cache_text_encoder_outputs\.py[^\n]*--batch_size\s+)\d+",
rf"\g<1>{int(override_te_cache_bs)}",
txt,
flags=re.MULTILINE,
)
# Prefer overriding variable definitions at top of script (safer than CLI regex)
def _set_var(name: str, value: str) -> None:
nonlocal txt
pattern = rf"(?m)^\s*{name}\s*=.*$"
replacement = f'{name}="{value}"' if not str(value).isdigit() else f'{name}={value}'
if re.search(pattern, txt):
txt = re.sub(pattern, lambda _m, r=replacement: r, txt)
else:
txt = f"{replacement}\n" + txt
if override_learning_rate:
_set_var('LEARNING_RATE', override_learning_rate)
if override_network_dim is not None:
_set_var('NETWORK_DIM', str(override_network_dim))
if override_seed is not None:
_set_var('SEED', str(override_seed))
if override_max_epochs is not None and override_max_epochs > 0:
_set_var('MAX_TRAIN_EPOCHS', str(override_max_epochs))
if override_save_every is not None and override_save_every > 0:
_set_var('SAVE_EVERY_N_EPOCHS', str(override_save_every))
_set_var('MODEL_VERSION', image_type)
# Write to a temp file alongside this repo for easier inspection
run_dir = TRAINING_DIR / ".gradio_runs"
run_dir.mkdir(parents=True, exist_ok=True)
tmp = run_dir / f"train_QIE_run_{os.getpid()}.sh"
tmp.write_text(txt, encoding="utf-8", newline="\n")
try:
os.chmod(tmp, 0o755)
except Exception:
pass
return tmp
def _pick_shell() -> str:
for sh in ("bash", "sh"):
if shutil.which(sh):
return sh
raise RuntimeError("No POSIX shell found. Please install bash or sh.")
def _is_git_repo(path: str) -> bool:
try:
out = subprocess.run(
["git", "-C", path, "rev-parse", "--is-inside-work-tree"],
capture_output=True,
text=True,
check=False,
)
return out.returncode == 0 and out.stdout.strip() == "true"
except Exception:
return False
def _startup_clone_musubi_tuner() -> None:
global MUSUBI_TUNER_DIR_RUNTIME
image_type = _get_qwen_image_type()
target, branch = _resolve_musubi_target(image_type)
MUSUBI_TUNER_DIR_RUNTIME = target
repo = DEFAULT_MUSUBI_TUNER_REPO
parent = os.path.dirname(target.rstrip("/\\")) or "/"
try:
os.makedirs(parent, exist_ok=True)
except PermissionError:
# Fallback to home directory
fallback_name = "musubi-tuner-layered" if image_type == "layered" else "musubi-tuner"
target = os.path.join(os.path.expanduser("~"), fallback_name)
MUSUBI_TUNER_DIR_RUNTIME = target
os.makedirs(os.path.dirname(target), exist_ok=True)
except Exception:
pass
if os.path.isdir(target) and _is_git_repo(target):
print(f"[QIE] musubi-tuner exists at {target}; syncing...")
try:
subprocess.run(["git", "-C", target, "fetch", "--all", "--prune"], check=False)
if branch:
res = subprocess.run(
["git", "-C", target, "rev-parse", "--abbrev-ref", "HEAD"],
capture_output=True,
text=True,
check=False,
)
current = res.stdout.strip()
if current != branch:
checkout = subprocess.run(
["git", "-C", target, "checkout", branch],
capture_output=True,
text=True,
check=False,
)
if checkout.returncode != 0:
subprocess.run(
["git", "-C", target, "checkout", "-B", branch, f"origin/{branch}"],
check=False,
)
subprocess.run(["git", "-C", target, "pull", "--ff-only"], check=False)
else:
subprocess.run(["git", "-C", target, "pull", "--ff-only"], check=False)
except Exception as e:
print(f"[QIE] git pull failed: {e}")
return
if os.path.exists(target) and not _is_git_repo(target):
print(f"[QIE] Warning: {target} exists and is not a git repo. Skipping clone.")
return
if branch:
print(f"[QIE] Cloning musubi-tuner into {target} from {repo} (branch {branch})...")
else:
print(f"[QIE] Cloning musubi-tuner into {target} from {repo} ...")
try:
if branch:
subprocess.run(["git", "clone", "--depth", "1", "--branch", branch, repo, target], check=True)
else:
subprocess.run(["git", "clone", "--depth", "1", repo, target], check=True)
print("[QIE] Clone completed.")
except subprocess.CalledProcessError as e:
print(f"[QIE] Clone failed at {target}: {e}")
# Last-chance fallback into home
if not target.startswith(os.path.expanduser("~")):
fallback_name = "musubi-tuner-layered" if image_type == "layered" else "musubi-tuner"
fallback = os.path.join(os.path.expanduser("~"), fallback_name)
print(f"[QIE] Retrying clone into {fallback}...")
try:
if branch:
subprocess.run(
["git", "clone", "--depth", "1", "--branch", branch, repo, fallback],
check=True,
)
else:
subprocess.run(["git", "clone", "--depth", "1", repo, fallback], check=True)
MUSUBI_TUNER_DIR_RUNTIME = fallback
print("[QIE] Clone completed in fallback.")
except Exception as e2:
print(f"[QIE] Clone failed in fallback as well: {e2}")
def _run_pip(args: List[str], cwd: Optional[str] = None) -> None:
cmd = [sys.executable, "-m", "pip"] + args
try:
print(f"[QIE] pip {' '.join(args)} (cwd={cwd or os.getcwd()})")
subprocess.run(cmd, check=True, cwd=cwd)
except subprocess.CalledProcessError as e:
print(f"[QIE] pip failed: {e}")
def _startup_install_musubi_deps() -> None:
repo_dir = MUSUBI_TUNER_DIR_RUNTIME
if not os.path.isdir(repo_dir):
print(f"[QIE] Skip deps: musubi-tuner not found at {repo_dir}")
return
# Upgrade basic build tooling (best-effort)
try:
_run_pip(["install", "-U", "pip", "setuptools", "wheel"])
except Exception:
pass
# Optional Torch extra via env: MUSUBI_TUNER_TORCH_EXTRA=cu124|cu128
extra = os.environ.get("MUSUBI_TUNER_TORCH_EXTRA", "").strip()
editable_spec = "." if not extra else f".[{extra}]"
# Install musubi-tuner in editable mode to expose entrypoints and deps
try:
_run_pip(["install", "-e", editable_spec], cwd=repo_dir)
except Exception:
# Fallback: plain install without editable
try:
_run_pip(["install", editable_spec], cwd=repo_dir)
except Exception:
print("[QIE] WARN: musubi-tuner installation failed. Continuing.")
@spaces.GPU
def run_training(
output_name: str,
caption: str,
image_uploads: Any,
target_prefix: str,
target_suffix: str,
control0_uploads: Any,
ctrl0_prefix: str,
ctrl0_suffix: str,
control1_uploads: Any,
ctrl1_prefix: str,
ctrl1_suffix: str,
control2_uploads: Any,
ctrl2_prefix: str,
ctrl2_suffix: str,
control3_uploads: Any,
ctrl3_prefix: str,
ctrl3_suffix: str,
control4_uploads: Any,
ctrl4_prefix: str,
ctrl4_suffix: str,
control5_uploads: Any,
ctrl5_prefix: str,
ctrl5_suffix: str,
control6_uploads: Any,
ctrl6_prefix: str,
ctrl6_suffix: str,
control7_uploads: Any,
ctrl7_prefix: str,
ctrl7_suffix: str,
learning_rate: str,
network_dim: int,
train_res_w: int,
train_res_h: int,
train_batch_size: int,
control_res_w: int,
control_res_h: int,
te_cache_batch_size: int,
seed: int,
max_epochs: int,
save_every: int,
config_only: bool,
) -> Iterable[tuple]:
# Basic validation
log_buf = "[QIE] Start Training invoked.\n"
ckpts: List[str] = []
artifacts: List[str] = []
run_out_dir = ""
# Emit an initial line so UI can confirm invocation
yield (log_buf, ckpts, artifacts, run_out_dir)
if not output_name.strip():
log_buf += "[ERROR] OUTPUT NAME is required.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
if not caption.strip():
log_buf += "[ERROR] CAPTION is required.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
image_type = _get_qwen_image_type()
log_buf += f"[QIE] Model type: {image_type}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
if image_type == "layered":
log_buf += "[ERROR] QWEN_IMAGE_TYPE=layered では Edit モードは実行できません。環境変数を edit-2509 / edit-2511 にして再起動してください。\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
# Ensure /auto holds helper files expected by the script
_ensure_workspace_auto_files()
# Resolve data root and create dataset directories (auto-decide)
global DATA_ROOT_RUNTIME
DATA_ROOT_RUNTIME = _ensure_data_root(None)
# Auto-generate dataset directory name
import time
ds_name = f"dataset_{int(time.time())}"
ds_dir = os.path.abspath(os.path.join(DATA_ROOT_RUNTIME, ds_name))
run_out_dir = os.path.abspath(os.path.join(ds_dir, output_name.strip()))
img_folder_name = DEFAULT_IMAGE_FOLDER
img_dir = os.path.join(ds_dir, img_folder_name)
os.makedirs(img_dir, exist_ok=True)
# Ingest uploads into dataset folders
base_files = _extract_paths(image_uploads)
if not base_files:
log_buf += "[ERROR] No images uploaded for IMAGE_FOLDER.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
base_filenames = _copy_uploads(base_files, img_dir)
log_buf += f"[QIE] Copied {len(base_filenames)} base images to {img_dir}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
# Prepare control sets
control_upload_sets = [
_extract_paths(control0_uploads),
_extract_paths(control1_uploads),
_extract_paths(control2_uploads),
_extract_paths(control3_uploads),
_extract_paths(control4_uploads),
_extract_paths(control5_uploads),
_extract_paths(control6_uploads),
_extract_paths(control7_uploads),
]
# Require control_0; others optional
if not control_upload_sets[0]:
log_buf += "[ERROR] control_0 images are required.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
control_dirs: List[Optional[str]] = []
for i, uploads in enumerate(control_upload_sets):
if not uploads:
control_dirs.append(None)
continue
folder_name = f"control_{i}"
cdir = os.path.join(ds_dir, folder_name)
os.makedirs(cdir, exist_ok=True)
# Simply copy; name matching will be handled by create_image_caption_json.py
_copy_uploads(uploads, cdir)
control_dirs.append(folder_name)
log_buf += f"[QIE] Copied {len(uploads)} control_{i} images to {cdir}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
# Prepare script with user parameters
control_folders = [
(control_dirs[i] if control_dirs[i] else None)
for i in range(8)
]
control_prefixes = [
ctrl0_prefix, ctrl1_prefix, ctrl2_prefix, ctrl3_prefix,
ctrl4_prefix, ctrl5_prefix, ctrl6_prefix, ctrl7_prefix,
]
control_suffixes = [
ctrl0_suffix, ctrl1_suffix, ctrl2_suffix, ctrl3_suffix,
ctrl4_suffix, ctrl5_suffix, ctrl6_suffix, ctrl7_suffix,
]
# Decide dataset_config path with fallback to runtime auto dir
ds_conf = str(Path(AUTO_DIR_RUNTIME) / "dataset_QIE.toml")
# Update dataset config with requested resolution/batch settings
try:
_update_dataset_toml(
ds_conf,
img_res_w=int(train_res_w) if train_res_w else None,
img_res_h=int(train_res_h) if train_res_h else None,
train_batch_size=int(train_batch_size) if train_batch_size else None,
control_res_w=int(control_res_w) if control_res_w else None,
control_res_h=int(control_res_h) if control_res_h else None,
remove_multiple_target=True,
)
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"
except Exception as e:
log_buf += f"[QIE] WARN: failed to update dataset config: {e}\n"
# Expose dataset config for download (if exists)
if os.path.isfile(ds_conf):
artifacts = [ds_conf]
# Resolve models_root and set output_dir_base to the unique dataset dir
models_root = MODELS_ROOT_RUNTIME
out_base = ds_dir
try:
os.makedirs(out_base, exist_ok=True)
except Exception:
pass
tmp_script = _prepare_script(
dataset_name=ds_name,
caption=caption,
data_root=DATA_ROOT_RUNTIME,
image_folder=img_folder_name,
control_folders=control_folders,
models_root=models_root,
output_dir_base=out_base,
dataset_config=ds_conf,
override_max_epochs=max_epochs if max_epochs and max_epochs > 0 else None,
override_save_every=save_every if save_every and save_every > 0 else None,
override_run_name=output_name.strip(),
target_prefix=(target_prefix or ""),
target_suffix=(target_suffix or ""),
control_prefixes=control_prefixes,
control_suffixes=control_suffixes,
override_learning_rate=(learning_rate or None),
override_network_dim=int(network_dim) if network_dim is not None else None,
override_te_cache_bs=int(te_cache_batch_size) if te_cache_batch_size else None,
override_seed=int(seed) if seed is not None else None,
)
out_dir = os.path.join(out_base, output_name.strip())
run_out_dir = out_dir
ckpts = _list_checkpoints(out_dir)
# Copy the final script to dataset dir for download
used_script_path = os.path.join(out_base, "train_QIE_used.sh")
try:
shutil.copy2(str(tmp_script), used_script_path)
try:
os.chmod(used_script_path, 0o755)
except Exception:
pass
if used_script_path not in artifacts:
artifacts.append(used_script_path)
except Exception:
pass
yield (log_buf, ckpts, artifacts, run_out_dir)
if config_only:
output_json = os.path.join(out_base, "metadata.jsonl")
_sync_dataset_config_jsonl(ds_conf, output_json)
log_buf += f"[QIE] Generating metadata: {output_json}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
script_path = Path(AUTO_DIR_RUNTIME) / "create_image_caption_json.py"
if not script_path.is_file():
log_buf += f"[ERROR] create_image_caption_json.py not found: {script_path}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
cmd = [
sys.executable,
str(script_path),
"-i",
img_dir,
"-c",
caption,
"-o",
output_json,
"--image-dir",
img_dir,
]
if target_prefix:
cmd += ["--target_prefix", target_prefix]
if target_suffix:
cmd += ["--target_suffix", target_suffix]
for i in range(8):
cdir_name = control_dirs[i] if i < len(control_dirs) else None
if cdir_name:
cmd += [f"--control_dir_{i}", os.path.join(ds_dir, cdir_name)]
if control_prefixes[i]:
cmd += [f"--control_prefix_{i}", control_prefixes[i]]
if control_suffixes[i]:
cmd += [f"--control_suffix_{i}", control_suffixes[i]]
try:
res = subprocess.run(cmd, capture_output=True, text=True, check=True)
if res.stdout:
log_buf += res.stdout
if res.stderr:
log_buf += res.stderr
except subprocess.CalledProcessError as e:
if e.stdout:
log_buf += e.stdout
if e.stderr:
log_buf += e.stderr
log_buf += "[ERROR] Metadata generation failed.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
if os.path.isfile(output_json) and output_json not in artifacts:
artifacts.append(output_json)
log_buf += "[QIE] Config-only mode: skipping cache/training.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
shell = _pick_shell()
log_buf += f"[QIE] Using shell: {shell}\n"
log_buf += f"[QIE] Running script: {tmp_script}\n"
# Run and stream output
# Ensure child Python processes are unbuffered for real-time logs
child_env = os.environ.copy()
child_env["PYTHONUNBUFFERED"] = "1"
child_env["PYTHONIOENCODING"] = "utf-8"
proc = subprocess.Popen(
[shell, str(tmp_script)],
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
bufsize=1,
universal_newlines=True,
env=child_env,
)
try:
assert proc.stdout is not None
i = 0
for line in proc.stdout:
log_buf += line
i += 1
if i % 30 == 0:
ckpts = _list_checkpoints(out_dir)
# Try to add metadata.jsonl once available
metadata_json = os.path.join(out_base, "metadata.jsonl")
if os.path.isfile(metadata_json) and metadata_json not in artifacts:
artifacts.append(metadata_json)
yield (log_buf, ckpts, artifacts, run_out_dir)
finally:
code = proc.wait()
# Try to locate latest LoRA file for download
lora_path = None
try:
ckpts = _list_checkpoints(out_dir)
except Exception:
pass
lora_path = ckpts[0] if ckpts else None
log_buf += f"[QIE] Exit code: {code}\n"
# Final attempt to include metadata.jsonl
metadata_json = os.path.join(out_base, "metadata.jsonl")
if os.path.isfile(metadata_json) and metadata_json not in artifacts:
artifacts.append(metadata_json)
yield (log_buf, ckpts, artifacts, run_out_dir)
def run_training_layered(
output_name: str,
caption: str,
image_uploads: Any,
target_prefix: str,
target_suffix: str,
layer1_uploads: Any,
layer2_uploads: Any,
layer3_uploads: Any,
layer4_uploads: Any,
layer5_uploads: Any,
layer6_uploads: Any,
layer7_uploads: Any,
layer8_uploads: Any,
layer9_uploads: Any,
layer10_uploads: Any,
layer11_uploads: Any,
layer12_uploads: Any,
layer13_uploads: Any,
layer14_uploads: Any,
layer15_uploads: Any,
layer16_uploads: Any,
layer17_uploads: Any,
layer18_uploads: Any,
layer19_uploads: Any,
layer20_uploads: Any,
layer21_uploads: Any,
layer22_uploads: Any,
layer23_uploads: Any,
layer24_uploads: Any,
layer25_uploads: Any,
layer26_uploads: Any,
layer27_uploads: Any,
layer28_uploads: Any,
layer29_uploads: Any,
layer30_uploads: Any,
layer31_uploads: Any,
layer32_uploads: Any,
layer1_prefix: str,
layer2_prefix: str,
layer3_prefix: str,
layer4_prefix: str,
layer5_prefix: str,
layer6_prefix: str,
layer7_prefix: str,
layer8_prefix: str,
layer9_prefix: str,
layer10_prefix: str,
layer11_prefix: str,
layer12_prefix: str,
layer13_prefix: str,
layer14_prefix: str,
layer15_prefix: str,
layer16_prefix: str,
layer17_prefix: str,
layer18_prefix: str,
layer19_prefix: str,
layer20_prefix: str,
layer21_prefix: str,
layer22_prefix: str,
layer23_prefix: str,
layer24_prefix: str,
layer25_prefix: str,
layer26_prefix: str,
layer27_prefix: str,
layer28_prefix: str,
layer29_prefix: str,
layer30_prefix: str,
layer31_prefix: str,
layer32_prefix: str,
layer1_suffix: str,
layer2_suffix: str,
layer3_suffix: str,
layer4_suffix: str,
layer5_suffix: str,
layer6_suffix: str,
layer7_suffix: str,
layer8_suffix: str,
layer9_suffix: str,
layer10_suffix: str,
layer11_suffix: str,
layer12_suffix: str,
layer13_suffix: str,
layer14_suffix: str,
layer15_suffix: str,
layer16_suffix: str,
layer17_suffix: str,
layer18_suffix: str,
layer19_suffix: str,
layer20_suffix: str,
layer21_suffix: str,
layer22_suffix: str,
layer23_suffix: str,
layer24_suffix: str,
layer25_suffix: str,
layer26_suffix: str,
layer27_suffix: str,
layer28_suffix: str,
layer29_suffix: str,
layer30_suffix: str,
layer31_suffix: str,
layer32_suffix: str,
learning_rate: str,
network_dim: int,
train_res_w: int,
train_res_h: int,
train_batch_size: int,
te_cache_batch_size: int,
seed: int,
max_epochs: int,
save_every: int,
config_only: bool,
) -> Iterable[tuple]:
log_buf = "[QIE] Start Training invoked.\n"
ckpts: List[str] = []
artifacts: List[str] = []
run_out_dir = ""
yield (log_buf, ckpts, artifacts, run_out_dir)
if not output_name.strip():
log_buf += "[ERROR] OUTPUT NAME is required.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
if not caption.strip():
log_buf += "[ERROR] CAPTION is required.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
image_type = _get_qwen_image_type()
log_buf += f"[QIE] Model type: {image_type}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
if image_type != "layered":
log_buf += "[ERROR] QWEN_IMAGE_TYPE=layered のときのみ Layered モードを実行できます。環境変数を設定して再起動してください。\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
_ensure_workspace_auto_files()
global DATA_ROOT_RUNTIME
DATA_ROOT_RUNTIME = _ensure_data_root(None)
import time
ds_name = f"dataset_{int(time.time())}"
ds_dir = os.path.abspath(os.path.join(DATA_ROOT_RUNTIME, ds_name))
run_out_dir = os.path.abspath(os.path.join(ds_dir, output_name.strip()))
img_folder_name = DEFAULT_IMAGE_FOLDER
img_dir = os.path.join(ds_dir, img_folder_name)
os.makedirs(img_dir, exist_ok=True)
base_files = _extract_paths(image_uploads)
if not base_files:
log_buf += "[ERROR] No images uploaded for IMAGE_FOLDER.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
base_filenames = _copy_uploads(base_files, img_dir)
log_buf += f"[QIE] Copied {len(base_filenames)} base images to {img_dir}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
layer_uploads = [
layer1_uploads,
layer2_uploads,
layer3_uploads,
layer4_uploads,
layer5_uploads,
layer6_uploads,
layer7_uploads,
layer8_uploads,
layer9_uploads,
layer10_uploads,
layer11_uploads,
layer12_uploads,
layer13_uploads,
layer14_uploads,
layer15_uploads,
layer16_uploads,
layer17_uploads,
layer18_uploads,
layer19_uploads,
layer20_uploads,
layer21_uploads,
layer22_uploads,
layer23_uploads,
layer24_uploads,
layer25_uploads,
layer26_uploads,
layer27_uploads,
layer28_uploads,
layer29_uploads,
layer30_uploads,
layer31_uploads,
layer32_uploads,
]
layer_prefixes = [
layer1_prefix,
layer2_prefix,
layer3_prefix,
layer4_prefix,
layer5_prefix,
layer6_prefix,
layer7_prefix,
layer8_prefix,
layer9_prefix,
layer10_prefix,
layer11_prefix,
layer12_prefix,
layer13_prefix,
layer14_prefix,
layer15_prefix,
layer16_prefix,
layer17_prefix,
layer18_prefix,
layer19_prefix,
layer20_prefix,
layer21_prefix,
layer22_prefix,
layer23_prefix,
layer24_prefix,
layer25_prefix,
layer26_prefix,
layer27_prefix,
layer28_prefix,
layer29_prefix,
layer30_prefix,
layer31_prefix,
layer32_prefix,
]
layer_suffixes = [
layer1_suffix,
layer2_suffix,
layer3_suffix,
layer4_suffix,
layer5_suffix,
layer6_suffix,
layer7_suffix,
layer8_suffix,
layer9_suffix,
layer10_suffix,
layer11_suffix,
layer12_suffix,
layer13_suffix,
layer14_suffix,
layer15_suffix,
layer16_suffix,
layer17_suffix,
layer18_suffix,
layer19_suffix,
layer20_suffix,
layer21_suffix,
layer22_suffix,
layer23_suffix,
layer24_suffix,
layer25_suffix,
layer26_suffix,
layer27_suffix,
layer28_suffix,
layer29_suffix,
layer30_suffix,
layer31_suffix,
layer32_suffix,
]
layer_upload_sets = [_extract_paths(u) for u in layer_uploads]
if not layer_upload_sets[0]:
log_buf += "[ERROR] Layer 1 (image_path_1) images are required.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
layer_dirs: List[Optional[str]] = []
for i, uploads in enumerate(layer_upload_sets):
if not uploads:
layer_dirs.append(None)
continue
folder_name = f"layer_{i + 1}"
cdir = os.path.join(ds_dir, folder_name)
os.makedirs(cdir, exist_ok=True)
_copy_uploads(uploads, cdir, force_rgba=True)
layer_dirs.append(folder_name)
log_buf += f"[QIE] Copied {len(uploads)} layer_{i + 1} images to {cdir}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
layer_dirs_abs = [
(os.path.join(ds_dir, name) if name else None)
for name in layer_dirs
]
ds_conf = str(Path(AUTO_DIR_RUNTIME) / "dataset_QIE.toml")
try:
_update_dataset_toml(
ds_conf,
img_res_w=int(train_res_w) if train_res_w else None,
img_res_h=int(train_res_h) if train_res_h else None,
train_batch_size=int(train_batch_size) if train_batch_size else None,
multiple_target=True,
)
log_buf += f"[QIE] Updated dataset config: resolution=({train_res_w},{train_res_h}), batch_size={train_batch_size}, multiple_target=true\n"
except Exception as e:
log_buf += f"[QIE] WARN: failed to update dataset config: {e}\n"
if os.path.isfile(ds_conf):
artifacts = [ds_conf]
models_root = MODELS_ROOT_RUNTIME
out_base = ds_dir
try:
os.makedirs(out_base, exist_ok=True)
except Exception:
pass
tmp_script = _prepare_script(
dataset_name=ds_name,
caption=caption,
data_root=DATA_ROOT_RUNTIME,
image_folder=img_folder_name,
control_folders=[],
models_root=models_root,
output_dir_base=out_base,
dataset_config=ds_conf,
override_max_epochs=max_epochs if max_epochs and max_epochs > 0 else None,
override_save_every=save_every if save_every and save_every > 0 else None,
override_run_name=output_name.strip(),
target_prefix=(target_prefix or ""),
target_suffix=(target_suffix or ""),
control_prefixes=[],
control_suffixes=[],
override_learning_rate=(learning_rate or None),
override_network_dim=int(network_dim) if network_dim is not None else None,
override_te_cache_bs=int(te_cache_batch_size) if te_cache_batch_size else None,
override_seed=int(seed) if seed is not None else None,
)
out_dir = os.path.join(out_base, output_name.strip())
run_out_dir = out_dir
ckpts = _list_checkpoints(out_dir)
used_script_path = os.path.join(out_base, "train_QIE_used.sh")
try:
shutil.copy2(str(tmp_script), used_script_path)
try:
os.chmod(used_script_path, 0o755)
except Exception:
pass
if used_script_path not in artifacts:
artifacts.append(used_script_path)
except Exception:
pass
yield (log_buf, ckpts, artifacts, run_out_dir)
output_json = os.path.join(out_base, "metadata.jsonl")
_sync_dataset_config_jsonl(ds_conf, output_json)
log_buf += f"[QIE] Generating layered metadata: {output_json}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
try:
count = _generate_layered_jsonl(
image_dir=img_dir,
caption=caption,
output_json=output_json,
control_dirs=layer_dirs_abs,
target_prefix=(target_prefix or ""),
target_suffix=(target_suffix or ""),
control_prefixes=layer_prefixes,
control_suffixes=layer_suffixes,
allow_single=True,
)
log_buf += f"[QIE] Layered metadata written: {count} entries\n"
except Exception as e:
log_buf += f"[ERROR] Layered metadata generation failed: {e}\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
if os.path.isfile(output_json) and output_json not in artifacts:
artifacts.append(output_json)
yield (log_buf, ckpts, artifacts, run_out_dir)
if config_only:
log_buf += "[QIE] Config-only mode: skipping cache/training.\n"
yield (log_buf, ckpts, artifacts, run_out_dir)
return
shell = _pick_shell()
log_buf += f"[QIE] Using shell: {shell}\n"
log_buf += f"[QIE] Running script: {tmp_script}\n"
child_env = os.environ.copy()
child_env["PYTHONUNBUFFERED"] = "1"
child_env["PYTHONIOENCODING"] = "utf-8"
proc = subprocess.Popen(
[shell, str(tmp_script)],
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
bufsize=1,
universal_newlines=True,
env=child_env,
)
try:
assert proc.stdout is not None
i = 0
for line in proc.stdout:
log_buf += line
i += 1
if i % 30 == 0:
ckpts = _list_checkpoints(out_dir)
metadata_json = os.path.join(out_base, "metadata.jsonl")
if os.path.isfile(metadata_json) and metadata_json not in artifacts:
artifacts.append(metadata_json)
yield (log_buf, ckpts, artifacts, run_out_dir)
finally:
code = proc.wait()
try:
ckpts = _list_checkpoints(out_dir)
except Exception:
pass
log_buf += f"[QIE] Exit code: {code}\n"
metadata_json = os.path.join(out_base, "metadata.jsonl")
if os.path.isfile(metadata_json) and metadata_json not in artifacts:
artifacts.append(metadata_json)
yield (log_buf, ckpts, artifacts, run_out_dir)
def build_ui() -> gr.Blocks:
css = """
.pad-section {
padding: 6px;
margin-bottom: 12px;
border: 1px solid var(--color-border, #e5e7eb);
border-radius: 8px;
background: var(--color-background-secondary, #ffffff);
}
.pad-section_0 {
padding: 6px;
margin-bottom: 12px;
border: 1px solid var(--color-border, #e5e7eb);
border-radius: 8px;
background: var(--color-background-secondary, #fafafa);
}
.pad-section_1 {
padding: 6px;
margin-bottom: 12px;
border: 1px solid var(--color-border, #e5e7eb);
border-radius: 8px;
background: var(--color-background-secondary, #eaeaea);
}
.svelte-1nguped {
background: none !important;
}
"""
with gr.Blocks(title="Qwen-Image-Edit: Trainer", css=css) as demo:
# 回収ボタン: 直近の dataset_ ディレクトリからチェックポイントとスクリプト/設定を再取得
def _refresh_all() -> tuple:
try:
ds_dir = _find_latest_dataset_dir(DATA_ROOT_RUNTIME)
except Exception:
ds_dir = None
try:
ck = _list_checkpoints(ds_dir) if ds_dir else []
except Exception:
ck = []
try:
sc = _collect_scripts_and_config(ds_dir)
except Exception:
sc = _collect_scripts_and_config(None)
return ck, sc
image_type = _get_qwen_image_type()
with gr.Tabs() as tabs:
with gr.TabItem("Training"):
if image_type == "layered":
gr.Markdown("""
# Qwen-Image-Layered Trainer
学習に使う画像をアップロードし、必要ならファイル名の前後にある共通の文字(prefix/suffix)を指定して、 自動でデータセットを作成し学習を開始します。難しい操作は不要です。
""")
with gr.Accordion("Settings", elem_classes=["pad-section"]):
with gr.Group():
with gr.Row():
output_name_layered = gr.Textbox(label="OUTPUT NAME", placeholder="my_lora_output", lines=1)
caption_layered = gr.Textbox(label="CAPTION", placeholder="A photo of ...", lines=2)
with gr.Row():
lr_input_layered = gr.Textbox(label="Learning rate", value="1e-3")
dim_input_layered = gr.Number(label="Network dim", value=4, precision=0)
train_bs_layered = gr.Number(label="Batch size (dataset)", value=1, precision=0)
seed_input_layered = gr.Number(label="Seed", value=42, precision=0)
max_epochs_layered = gr.Number(label="Max epochs", value=100, precision=0)
save_every_layered = gr.Number(label="Save every N epochs", value=10, precision=0)
with gr.Row():
tr_w_layered = gr.Number(label="Image resolution W", value=1024, precision=0)
tr_h_layered = gr.Number(label="Image resolution H", value=1024, precision=0)
te_bs_layered = gr.Number(label="TE cache batch size", value=16, precision=0)
with gr.Accordion("Base Image (image_path_0)", elem_classes=["pad-section_0"]):
with gr.Group():
with gr.Row():
base_images_input = gr.File(label="Upload base images (image_path_0)", file_count="multiple", type="filepath", height=220, scale=3)
base_gallery = gr.Gallery(label="Base preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
base_prefix = gr.Textbox(label="Base prefix", placeholder="e.g., IMG_")
base_suffix = gr.Textbox(label="Base suffix", placeholder="e.g., _v2")
with gr.Accordion("prefix/sufixについて", open=False):
gr.Markdown("""
ファイル名の対応付けのルール:
- base画像のファイル名から Base prefix/suffix を取り除いたものを key とします。
- 各レイヤーは `layer_i prefix + key + layer_i suffix + .png` を探します。
- レイヤーが1枚のみのときは全ベース画像に適用します。
""")
layer_files: List[gr.File] = []
layer_galleries: List[gr.Gallery] = []
layer_prefixes: List[gr.Textbox] = []
layer_suffixes: List[gr.Textbox] = []
group_size = 8
for group_start in range(1, LAYER_MAX + 1, group_size):
group_end = min(group_start + group_size - 1, LAYER_MAX)
open_flag = True if group_start == 1 else False
group_label = f"Layers {group_start}-{group_end}"
with gr.Accordion(group_label, open=open_flag, elem_classes=["pad-section_0"]):
with gr.Group():
for i in range(group_start, group_end + 1):
inner_label = f"Layer {i}"
inner_open = True if i == 1 else False
inner_class = "pad-section_1"
with gr.Accordion(inner_label, open=inner_open, elem_classes=[inner_class]):
with gr.Group():
with gr.Row():
file_label = f"Upload layer_{i} images"
if i == 1:
file_label = "Upload layer_1 images (required)"
layer_file = gr.File(
label=file_label,
file_count="multiple",
type="filepath",
height=220,
scale=3,
)
layer_gallery = gr.Gallery(
label=f"layer_{i} preview",
columns=4,
height=220,
object_fit="contain",
preview=True,
scale=3,
)
with gr.Column(scale=1):
with gr.Row():
layer_prefix = gr.Textbox(label=f"layer_{i} prefix", placeholder="")
layer_suffix = gr.Textbox(label=f"layer_{i} suffix", placeholder="")
layer_files.append(layer_file)
layer_galleries.append(layer_gallery)
layer_prefixes.append(layer_prefix)
layer_suffixes.append(layer_suffix)
layer_file.change(fn=_files_to_gallery, inputs=layer_file, outputs=layer_gallery)
with gr.Row():
run_btn_layered = gr.Button("Start Training", variant="primary")
config_btn_layered = gr.Button("設定のみ生成", variant="secondary")
run_out_dir_box_layered = gr.Textbox(label="出力フォルダ", lines=1, interactive=False)
scripts_files_layered = gr.Files(label="Scripts & Config (live)", interactive=False)
ckpt_files_layered = gr.Files(label="Checkpoints (live)", interactive=False)
logs_layered = gr.Textbox(label="Logs", lines=20)
with gr.Row():
refresh_scripts_btn_layered = gr.Button("ファイルを再取得", variant="secondary")
base_images_input.change(fn=_files_to_gallery, inputs=base_images_input, outputs=base_gallery)
config_only_off_layered = gr.State(False)
config_only_on_layered = gr.State(True)
run_btn_layered.click(
fn=run_training_layered,
inputs=[
output_name_layered, caption_layered, base_images_input, base_prefix, base_suffix,
*layer_files,
*layer_prefixes,
*layer_suffixes,
lr_input_layered, dim_input_layered,
tr_w_layered, tr_h_layered, train_bs_layered, te_bs_layered,
seed_input_layered, max_epochs_layered, save_every_layered, config_only_off_layered,
],
outputs=[logs_layered, ckpt_files_layered, scripts_files_layered, run_out_dir_box_layered],
)
config_btn_layered.click(
fn=run_training_layered,
inputs=[
output_name_layered, caption_layered, base_images_input, base_prefix, base_suffix,
*layer_files,
*layer_prefixes,
*layer_suffixes,
lr_input_layered, dim_input_layered,
tr_w_layered, tr_h_layered, train_bs_layered, te_bs_layered,
seed_input_layered, max_epochs_layered, save_every_layered, config_only_on_layered,
],
outputs=[logs_layered, ckpt_files_layered, scripts_files_layered, run_out_dir_box_layered],
)
refresh_scripts_btn_layered.click(
fn=_refresh_all,
inputs=[],
outputs=[ckpt_files_layered, scripts_files_layered],
)
else:
gr.Markdown("""
# Qwen-Image-Edit Trainer
学習に使う画像をアップロードし、必要ならファイル名の前後にある共通の文字(prefix/suffix)を指定して、
自動でデータセットを作成し学習を開始します。難しい操作は不要です。
""")
with gr.Accordion("Settings", elem_classes=["pad-section"]):
with gr.Group():
with gr.Row():
output_name = gr.Textbox(label="OUTPUT NAME", placeholder="my_lora_output", lines=1)
caption = gr.Textbox(label="CAPTION", placeholder="A photo of ...", lines=2)
with gr.Row():
lr_input = gr.Textbox(label="Learning rate", value="1e-3")
dim_input = gr.Number(label="Network dim", value=4, precision=0)
train_bs = gr.Number(label="Batch size (dataset)", value=1, precision=0)
seed_input = gr.Number(label="Seed", value=42, precision=0)
max_epochs = gr.Number(label="Max epochs", value=100, precision=0)
save_every = gr.Number(label="Save every N epochs", value=10, precision=0)
with gr.Row():
tr_w = gr.Number(label="Image resolution W", value=1024, precision=0)
tr_h = gr.Number(label="Image resolution H", value=1024, precision=0)
cr_w = gr.Number(label="Control resolution W", value=1024, precision=0)
cr_h = gr.Number(label="Control resolution H", value=1024, precision=0)
te_bs = gr.Number(label="TE cache batch size", value=16, precision=0)
with gr.Accordion("Target Image", elem_classes=["pad-section_0"]):
with gr.Group():
with gr.Row():
images_input = gr.File(label="Upload target images", file_count="multiple", type="filepath", height=220, scale=3)
main_gallery = gr.Gallery(label="Target preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
main_prefix = gr.Textbox(label="Target prefix", placeholder="e.g., IMG_")
main_suffix = gr.Textbox(label="Target suffix", placeholder="e.g., _v2")
with gr.Accordion("prefix/sufixについて", open=False):
gr.Markdown("""
ファイルの同名判定のため、画像のファイル名から共通の先頭/末尾文字を取り除く指定(例: IMG_ や _v2)
- まずターゲット画像のファイル名(拡張子なし)から、指定した Target prefix/suffix を取り除いたものを key とします。
- 各コントロールは「付加」規則で、期待名 = control_prefix_i + key + control_suffix_i + ".png" を探して対応付けます。
- アップロード時に画像は自動で .png に変換して保存します(元のファイル名のベースは維持)。
- Control 0 は必須、Control 1〜7 は任意。コントロール画像が1枚だけのときは、すべてのターゲット画像に適用します。
""")
# control_0 is required and shown outside the accordion
with gr.Accordion("Control 0", elem_classes=["pad-section_1"]):
with gr.Group():
with gr.Row():
ctrl0_files = gr.File(label="Upload control_0 images (required)", file_count="multiple", type="filepath", height=220, scale=3)
ctrl0_gallery = gr.Gallery(label="control_0 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl0_prefix = gr.Textbox(label="control_0 prefix", placeholder="e.g., C0_")
ctrl0_suffix = gr.Textbox(label="control_0 suffix", placeholder="e.g., _mask")
# Optional controls start from 1, accordion closed by default
with gr.Accordion("Control 1", open=False, elem_classes=["pad-section_0"]):
with gr.Group():
with gr.Row():
ctrl1_files = gr.File(label="Upload control_1 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl1_gallery = gr.Gallery(label="control_1 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl1_prefix = gr.Textbox(label="control_1 prefix", placeholder="")
ctrl1_suffix = gr.Textbox(label="control_1 suffix", placeholder="")
with gr.Accordion("Control 2", open=False, elem_classes=["pad-section_1"]):
with gr.Group():
with gr.Row():
ctrl2_files = gr.File(label="Upload control_2 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl2_gallery = gr.Gallery(label="control_2 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl2_prefix = gr.Textbox(label="control_2 prefix", placeholder="")
ctrl2_suffix = gr.Textbox(label="control_2 suffix", placeholder="")
with gr.Accordion("Control 3", open=False, elem_classes=["pad-section_0"]):
with gr.Group():
with gr.Row():
ctrl3_files = gr.File(label="Upload control_3 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl3_gallery = gr.Gallery(label="control_3 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl3_prefix = gr.Textbox(label="control_3 prefix", placeholder="")
ctrl3_suffix = gr.Textbox(label="control_3 suffix", placeholder="")
with gr.Accordion("Control 4", open=False, elem_classes=["pad-section_1"]):
with gr.Group():
with gr.Row():
ctrl4_files = gr.File(label="Upload control_4 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl4_gallery = gr.Gallery(label="control_4 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl4_prefix = gr.Textbox(label="control_4 prefix", placeholder="")
ctrl4_suffix = gr.Textbox(label="control_4 suffix", placeholder="")
with gr.Accordion("Control 5", open=False, elem_classes=["pad-section_0"]):
with gr.Group():
with gr.Row():
ctrl5_files = gr.File(label="Upload control_5 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl5_gallery = gr.Gallery(label="control_5 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl5_prefix = gr.Textbox(label="control_5 prefix", placeholder="")
ctrl5_suffix = gr.Textbox(label="control_5 suffix", placeholder="")
with gr.Accordion("Control 6", open=False, elem_classes=["pad-section_1"]):
with gr.Group():
with gr.Row():
ctrl6_files = gr.File(label="Upload control_6 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl6_gallery = gr.Gallery(label="control_6 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl6_prefix = gr.Textbox(label="control_6 prefix", placeholder="")
ctrl6_suffix = gr.Textbox(label="control_6 suffix", placeholder="")
with gr.Accordion("Control 7", open=False, elem_classes=["pad-section_0"]):
with gr.Group():
with gr.Row():
ctrl7_files = gr.File(label="Upload control_7 images", file_count="multiple", type="filepath", height=220, scale=3)
ctrl7_gallery = gr.Gallery(label="control_7 preview", columns=4, height=220, object_fit='contain', preview=True, scale=3)
with gr.Column(scale=1):
with gr.Row():
ctrl7_prefix = gr.Textbox(label="control_7 prefix", placeholder="")
ctrl7_suffix = gr.Textbox(label="control_7 suffix", placeholder="")
# Models root / OUTPUT_DIR_BASE / DATASET_CONFIG are auto-resolved at runtime; no user input needed.
with gr.Row():
run_btn = gr.Button("Start Training", variant="primary")
config_btn = gr.Button("設定のみ生成", variant="secondary")
run_out_dir_box = gr.Textbox(label="出力フォルダ", lines=1, interactive=False)
scripts_files = gr.Files(label="Scripts & Config (live)", interactive=False)
ckpt_files = gr.Files(label="Checkpoints (live)", interactive=False)
logs = gr.Textbox(label="Logs", lines=20)
with gr.Row():
refresh_scripts_btn = gr.Button("ファイルを再取得", variant="secondary")
# moved max_epochs/save_every above next to OUTPUT NAME
# Wire previews
images_input.change(fn=_files_to_gallery, inputs=images_input, outputs=main_gallery)
ctrl0_files.change(fn=_files_to_gallery, inputs=ctrl0_files, outputs=ctrl0_gallery)
ctrl1_files.change(fn=_files_to_gallery, inputs=ctrl1_files, outputs=ctrl1_gallery)
ctrl2_files.change(fn=_files_to_gallery, inputs=ctrl2_files, outputs=ctrl2_gallery)
ctrl3_files.change(fn=_files_to_gallery, inputs=ctrl3_files, outputs=ctrl3_gallery)
ctrl4_files.change(fn=_files_to_gallery, inputs=ctrl4_files, outputs=ctrl4_gallery)
ctrl5_files.change(fn=_files_to_gallery, inputs=ctrl5_files, outputs=ctrl5_gallery)
ctrl6_files.change(fn=_files_to_gallery, inputs=ctrl6_files, outputs=ctrl6_gallery)
ctrl7_files.change(fn=_files_to_gallery, inputs=ctrl7_files, outputs=ctrl7_gallery)
config_only_off = gr.State(False)
config_only_on = gr.State(True)
run_btn.click(
fn=run_training,
inputs=[
output_name, caption, images_input, main_prefix, main_suffix,
ctrl0_files, ctrl0_prefix, ctrl0_suffix,
ctrl1_files, ctrl1_prefix, ctrl1_suffix,
ctrl2_files, ctrl2_prefix, ctrl2_suffix,
ctrl3_files, ctrl3_prefix, ctrl3_suffix,
ctrl4_files, ctrl4_prefix, ctrl4_suffix,
ctrl5_files, ctrl5_prefix, ctrl5_suffix,
ctrl6_files, ctrl6_prefix, ctrl6_suffix,
ctrl7_files, ctrl7_prefix, ctrl7_suffix,
lr_input, dim_input,
tr_w, tr_h, train_bs, cr_w, cr_h, te_bs,
seed_input, max_epochs, save_every, config_only_off,
],
outputs=[logs, ckpt_files, scripts_files, run_out_dir_box],
)
config_btn.click(
fn=run_training,
inputs=[
output_name, caption, images_input, main_prefix, main_suffix,
ctrl0_files, ctrl0_prefix, ctrl0_suffix,
ctrl1_files, ctrl1_prefix, ctrl1_suffix,
ctrl2_files, ctrl2_prefix, ctrl2_suffix,
ctrl3_files, ctrl3_prefix, ctrl3_suffix,
ctrl4_files, ctrl4_prefix, ctrl4_suffix,
ctrl5_files, ctrl5_prefix, ctrl5_suffix,
ctrl6_files, ctrl6_prefix, ctrl6_suffix,
ctrl7_files, ctrl7_prefix, ctrl7_suffix,
lr_input, dim_input,
tr_w, tr_h, train_bs, cr_w, cr_h, te_bs,
seed_input, max_epochs, save_every, config_only_on,
],
outputs=[logs, ckpt_files, scripts_files, run_out_dir_box],
)
refresh_scripts_btn.click(
fn=_refresh_all,
inputs=[],
outputs=[ckpt_files, scripts_files],
)
with gr.TabItem("Prompt Generator"):
gr.Markdown("""
# 🎨 A→B 変換プロンプト自動生成
画像A(入力)と画像B(出力)、補足説明を入力すると、
A→B の変換内容を英語プロンプトとして自動生成し、タスク名候補(3件)も提案します。
モデルは `gpt-5` を使用します。
""")
api_key_pg = gr.Textbox(label="OpenAI API Key", type="password", placeholder="sk-...")
with gr.Row():
img_a_pg = gr.Image(type="filepath", label="Image A (Input)", height=300)
img_b_pg = gr.Image(type="filepath", label="Image B (Output)", height=300)
notes_pg = gr.Textbox(label="補足説明(日本語可)", lines=4, value="この画像は例であって、汎用的なプロンプトにする")
want_japanese_pg = gr.Checkbox(label="日本語訳を含める", value=True)
run_btn_pg = gr.Button("生成する", variant="primary")
english_out_pg = gr.Textbox(label="English Prompt", lines=8)
names_out_pg = gr.Textbox(label="Name Suggestions", lines=4)
japanese_out_pg = gr.Textbox(label="日本語訳(任意)", lines=8)
def _on_click_prompt(api_key_in, a_path, b_path, notes_in, ja_flag):
# Lazy import to avoid constructing extra Blocks at startup
qpg = importlib.import_module("QIE_prompt_generator")
a_url = qpg.file_to_data_url(a_path) if a_path else None
b_url = qpg.file_to_data_url(b_path) if b_path else None
return qpg.call_openai_chat(api_key_in, a_url, b_url, notes_in, ja_flag)
run_btn_pg.click(
fn=_on_click_prompt,
inputs=[api_key_pg, img_a_pg, img_b_pg, notes_pg, want_japanese_pg],
outputs=[english_out_pg, names_out_pg, japanese_out_pg],
)
return demo
def _startup_download_models() -> None:
global MODELS_ROOT_RUNTIME
# Pick a writable models directory
candidate = os.environ.get("QWEN_IMAGE_MODELS_DIR", DEFAULT_MODELS_ROOT)
try:
os.makedirs(candidate, exist_ok=True)
MODELS_ROOT_RUNTIME = candidate
except PermissionError:
MODELS_ROOT_RUNTIME = os.path.join(os.path.expanduser("~"), "Qwen-Image_models")
os.makedirs(MODELS_ROOT_RUNTIME, exist_ok=True)
print(f"[QIE] Ensuring models in: {MODELS_ROOT_RUNTIME}")
skip_raw = os.environ.get("QWEN_IMAGE_SKIP_DOWNLOAD", "")
if skip_raw.strip().lower() in ("1", "true", "yes", "on"):
print("[QIE] QWEN_IMAGE_SKIP_DOWNLOAD=1: skipping model download.")
return
try:
download_all_models(MODELS_ROOT_RUNTIME)
except Exception as e:
print(f"[QIE] Model download failed: {e}")
if __name__ == "__main__":
# 1) Ensure musubi-tuner is cloned before anything else
_startup_clone_musubi_tuner()
# 1.1) Install musubi-tuner dependencies (best-effort)
_startup_install_musubi_deps()
# 2) Download models at startup (blocking by design)
_startup_download_models()
# 3) Launch Gradio app
ui = build_ui()
# Limit concurrency (training is heavy). Enable queue for Spaces compatibility.
# Use generic signature to support multiple gradio versions.
try:
ui = ui.queue(max_size=16)
except TypeError:
ui = ui.queue()
# Allow Gradio to serve files saved under our runtime dirs
try:
allowed = [
AUTO_DIR_RUNTIME,
os.path.join(AUTO_DIR_RUNTIME, "train_LoRA"),
DEFAULT_DATA_ROOT,
DATA_ROOT_RUNTIME,
os.path.join(os.path.expanduser("~"), "auto"),
os.path.join(os.path.expanduser("~"), "data"),
]
ui.launch(server_name="0.0.0.0", allowed_paths=allowed, ssr_mode=False)
except TypeError:
# Older gradio without allowed_paths
try:
ui.launch(server_name="0.0.0.0", ssr_mode=False)
except TypeError:
# Very old gradio without ssr_mode
ui.launch(server_name="0.0.0.0")