File size: 8,721 Bytes
5e296ef | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | import os
import sys
import argparse
import subprocess
import contextlib
# Apply PyTorch compatibility before Transformers imports TorchAO.
import musubi_tuner_gui.torch_compat # noqa: F401
import gradio as gr
from musubi_tuner_gui.lora_gui import lora_tab
from musubi_tuner_gui.qwen_image_lora_gui import qwen_image_lora_tab
from musubi_tuner_gui.wan_lora_gui import wan_lora_tab
from musubi_tuner_gui.flux_lora_gui import flux_lora_tab
from musubi_tuner_gui.zimage_lora_gui import zimage_lora_tab
from musubi_tuner_gui.ideogram4_lora_gui import ideogram4_lora_tab
from musubi_tuner_gui.krea2_lora_gui import krea2_lora_tab
from musubi_tuner_gui.image_captioning_gui import image_captioning_tab
from musubi_tuner_gui.model_quantizer_gui import model_quantizer_tab
from musubi_tuner_gui.image_preprocessing_gui import image_preprocessing_tab
from musubi_tuner_gui.changelog_gui import version_history_tab
from musubi_tuner_gui.lora_extractor_gui import lora_extractor_tab
from musubi_tuner_gui.lora_merge_gui import lora_merge_tab
from musubi_tuner_gui.lora_convert_gui import lora_convert_tab
from musubi_tuner_gui.custom_logging import setup_logging
from musubi_tuner_gui.class_gui_config import GUIConfig
from musubi_tuner_gui.class_tab_config_manager import TabConfigManager
import toml
PYTHON = sys.executable
project_dir = os.path.dirname(os.path.abspath(__file__))
# Function to read file content, suppressing any FileNotFoundError
def read_file_content(file_path):
with contextlib.suppress(FileNotFoundError):
with open(file_path, "r", encoding="utf8") as file:
return file.read()
return ""
# Function to initialize the Gradio UI interface
def initialize_ui_interface(config_manager, headless, release_info, readme_content):
# Load custom CSS if available
css = read_file_content("./assets/style.css")
# Create the main Gradio Blocks interface
ui_interface = gr.Blocks(css=css, title="SECourses Musubi Trainer V30.2", theme=gr.themes.Soft())
with ui_interface:
# Add title with Patreon link
gr.Markdown("# SECourses Musubi Trainer V30.2 : [https://www.patreon.com/posts/137551634](https://www.patreon.com/posts/137551634)")
# Create tabs for different functionalities
with gr.Tab("Qwen Image Training"):
qwen_config = config_manager.get_config_for_tab("qwen_image")
qwen_image_lora_tab(headless=headless, config=qwen_config)
with gr.Tab("Wan Models Training"):
wan_config = config_manager.get_config_for_tab("wan")
wan_lora_tab(headless=headless, config=wan_config)
with gr.Tab("FLUX Training"):
flux_config = config_manager.get_config_for_tab("flux")
flux_lora_tab(headless=headless, config=flux_config)
with gr.Tab("Z Image Training"):
zimage_config = config_manager.get_config_for_tab("zimage")
zimage_lora_tab(headless=headless, config=zimage_config)
with gr.Tab("Ideogram 4 Training"):
ideogram4_config = config_manager.get_config_for_tab("ideogram4")
ideogram4_lora_tab(headless=headless, config=ideogram4_config)
with gr.Tab("Krea 2 Training"):
krea2_config = config_manager.get_config_for_tab("krea2")
krea2_lora_tab(headless=headless, config=krea2_config)
with gr.Tab("Image Captioning"):
captioning_config = config_manager.get_config_for_tab("image_captioning")
image_captioning_tab(headless=headless, config=captioning_config)
with gr.Tab("Model Quantizer"):
quant_config = config_manager.get_config_for_tab("model_quantizer")
model_quantizer_tab(headless=headless, config=quant_config)
with gr.Tab("LoRA Extractor"):
lora_extractor_tab(headless=headless, config=None)
with gr.Tab("LoRA Merger"):
lora_merge_tab(headless=headless, config=None)
with gr.Tab("LoRA Converter"):
lora_convert_tab(headless=headless, config=None)
with gr.Tab("Image Preprocessing"):
preprocessing_config = config_manager.get_config_for_tab("image_preprocessing")
image_preprocessing_tab(headless=headless, config=preprocessing_config)
with gr.Tab("Version History"):
version_history_tab(headless=headless, config=None)
with gr.Tab("Musubi Tuner (Deprecated)"):
musubi_config = config_manager.get_config_for_tab("musubi_tuner")
lora_tab(headless=headless, config=musubi_config)
return ui_interface
# Function to configure and launch the UI
def UI(**kwargs):
# Add custom JavaScript if specified
log.info(f"headless: {kwargs.get('headless', False)}")
# Load release and README information
release_info = "v18.0" # Hardcoded version since pyproject.toml is not needed
readme_content = read_file_content("./README.md")
# Initialize tab-aware configuration manager - default to qwen_image_defaults.toml
config_manager = TabConfigManager(config_file_path=kwargs.get("config", "./qwen_image_defaults.toml"))
if config_manager.user_loaded_config:
log.info(f"Loaded user configuration from '{kwargs.get('config', './qwen_image_defaults.toml')}'...")
else:
log.info("No user config loaded - will use tab-specific defaults")
# Initialize the Gradio UI interface
ui_interface = initialize_ui_interface(config_manager, kwargs.get("headless", False), release_info, readme_content)
# Construct launch parameters using dictionary comprehension
launch_params = {
"server_name": kwargs.get("listen"),
"auth": (kwargs["username"], kwargs["password"]) if kwargs.get("username") and kwargs.get("password") else None,
"server_port": kwargs.get("server_port", 0) if kwargs.get("server_port", 0) > 0 else None,
"inbrowser": kwargs.get("inbrowser", True),
"share": kwargs.get("share", False),
"root_path": kwargs.get("root_path", None),
"debug": kwargs.get("debug", False),
}
# This line filters out any key-value pairs from `launch_params` where the value is `None`, ensuring only valid parameters are passed to the `launch` function.
launch_params = {k: v for k, v in launch_params.items() if v is not None}
# Launch the Gradio interface with the specified parameters
ui_interface.launch(**launch_params)
# Function to initialize argument parser for command-line arguments
def initialize_arg_parser():
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=str, default="./qwen_image_defaults.toml", help="Path to the toml config file for interface defaults")
parser.add_argument("--debug", action="store_true", help="Debug on")
parser.add_argument("--listen", type=str, default="127.0.0.1", help="IP to listen on for connections to Gradio")
parser.add_argument("--username", type=str, default="", help="Username for authentication")
parser.add_argument("--password", type=str, default="", help="Password for authentication")
parser.add_argument("--server_port", type=int, default=0, help="Port to run the server listener on")
parser.add_argument("--inbrowser", action="store_true", default=True, help="Open in browser")
parser.add_argument("--share", action="store_true", help="Share the gradio UI")
parser.add_argument("--headless", action="store_true", help="Is the server headless")
parser.add_argument("--language", type=str, default=None, help="Set custom language")
parser.add_argument("--use-ipex", action="store_true", help="Use IPEX environment")
parser.add_argument("--use-rocm", action="store_true", help="Use ROCm environment")
parser.add_argument("--do_not_use_shell", action="store_true", help="Enforce not to use shell=True when running external commands")
parser.add_argument("--do_not_share", action="store_true", help="Do not share the gradio UI")
parser.add_argument("--requirements", type=str, default=None, help="requirements file to use for validation")
parser.add_argument("--root_path", type=str, default=None, help="`root_path` for Gradio to enable reverse proxy support. e.g. /kohya_ss")
parser.add_argument("--noverify", action="store_true", help="Disable requirements verification")
return parser
if __name__ == "__main__":
# Initialize argument parser and parse arguments
parser = initialize_arg_parser()
args = parser.parse_args()
# Set up logging based on the debug flag
log = setup_logging(debug=args.debug)
# Launch the UI with the provided arguments
UI(**vars(args))
|