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))