forked from Akegarasu/lora-scripts
-
Notifications
You must be signed in to change notification settings - Fork 15
Expand file tree
/
Copy pathgui.py
More file actions
223 lines (192 loc) · 8.76 KB
/
Copy pathgui.py
File metadata and controls
223 lines (192 loc) · 8.76 KB
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
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
import argparse
import os
import platform
import subprocess
import sys
import time
from mikazuki.launch_utils import (base_dir_path, catch_exception, git_tag,
prepare_environment, check_port_avaliable,
ensure_requirements_installed)
from mikazuki.log import log
from mikazuki.portable_utils import sanitize_embedded_deps, train_env_overrides
parser = argparse.ArgumentParser(description="GUI for stable diffusion training")
parser.add_argument("--host", type=str, default="127.0.0.1")
parser.add_argument("--port", type=int, default=28000, help="Port to run the server on")
parser.add_argument("--listen", action="store_true")
parser.add_argument("--skip-prepare-environment", action="store_true")
parser.add_argument("--skip-prepare-onnxruntime", action="store_true")
parser.add_argument("--disable-tensorboard", action="store_true", default=False)
parser.add_argument("--disable-train-monitor", action="store_true")
parser.add_argument("--disable-auto-mirror", action="store_true")
parser.add_argument("--tensorboard-host", type=str, default="127.0.0.1", help="Port to run the tensorboard")
parser.add_argument("--tensorboard-port", type=int, default=6006, help="Port to run the tensorboard")
parser.add_argument("--train-monitor-port", type=int, default=6008, help="Port to run the train status monitor")
parser.add_argument("--localization", type=str)
parser.add_argument("--browser", type=str, default=None,
choices=["chrome", "edge", "default"],
help="Browser to open GUI: chrome, edge, or default (system default)")
parser.add_argument("--dev", action="store_true")
def _popen(command: list[str], **kwargs) -> subprocess.Popen:
if sys.platform.startswith("linux"):
command = [
sys.executable,
str(base_dir_path() / "mikazuki" / "child_process.py"),
str(os.getpid()),
*command,
]
return subprocess.Popen(command, **kwargs)
def ensure_port_available(
port: int,
fallback_start: int,
fallback_end: int,
label: str,
reserved_ports: set[int],
preferred_reserved_port: int | None = None,
) -> int:
if (port == preferred_reserved_port or port not in reserved_ports) and check_port_avaliable(port):
reserved_ports.add(port)
return port
for candidate in range(fallback_start, fallback_end):
if candidate in reserved_ports and candidate != preferred_reserved_port:
continue
if check_port_avaliable(candidate):
reserved_ports.add(candidate)
log.warning(f"{label} port {port} is already in use, using {candidate} instead.")
return candidate
log.error(f"{label}: no available port in range {fallback_start}-{fallback_end}.")
return port
@catch_exception
def run_train_monitor():
env = os.environ.copy()
return _popen([sys.executable, str(base_dir_path() / "train_monitor" / "server.py")], env=env)
@catch_exception
def run_tensorboard():
log.info("Starting tensorboard...")
return _popen([sys.executable, "-m", "tensorboard.main", "--logdir", "logs",
"--host", args.tensorboard_host, "--port", str(args.tensorboard_port)])
def stop_child_processes(
processes: list[tuple[str, subprocess.Popen]], timeout: float = 5.0
) -> None:
running = []
for name, process in processes:
try:
if process.poll() is None:
running.append((name, process))
except Exception as e:
log.warning(f"Could not inspect {name} process: {e}")
if not running:
return
for name, process in running:
log.info(f"Stopping {name} (PID {process.pid})...")
try:
process.terminate()
except Exception as e:
log.warning(f"Could not terminate {name} (PID {process.pid}): {e}")
deadline = time.monotonic() + timeout
remaining = []
for name, process in running:
try:
process.wait(timeout=max(0, deadline - time.monotonic()))
except subprocess.TimeoutExpired:
remaining.append((name, process))
except Exception as e:
log.warning(f"Could not wait for {name} (PID {process.pid}): {e}")
remaining.append((name, process))
for name, process in remaining:
log.warning(f"{name} did not stop in time; killing PID {process.pid}.")
try:
process.kill()
except ProcessLookupError:
continue
except Exception as e:
log.warning(f"Could not kill {name} (PID {process.pid}): {e}")
continue
try:
process.wait()
except Exception as e:
log.warning(f"Could not reap {name} (PID {process.pid}): {e}")
def launch():
sanitize_embedded_deps(log.warning)
from mikazuki.china_hub import enable_china_hub
if enable_china_hub():
log.info("Using ModelScope hub patch for Hugging Face downloads (国内下载走魔搭)")
for key, value in train_env_overrides().items():
os.environ.setdefault(key, value)
log.info("Starting SD-Trainer Mikazuki GUI...")
log.info(f"Base directory: {base_dir_path()}, Working directory: {os.getcwd()}")
log.info(f"{platform.system()} Python {platform.python_version()} {sys.executable}")
if not args.skip_prepare_environment:
prepare_environment(disable_auto_mirror=args.disable_auto_mirror,
prepare_onnxruntime=not args.skip_prepare_onnxruntime)
else:
# Portable launch skips prepare_environment, so requirements.txt is
# otherwise never validated. Run a cheap presence-only guard so newly
# added or missing packages (e.g. onnxruntime-gpu) get repaired instead
# of silently breaking tagging/training.
ensure_requirements_installed("requirements.txt")
# Protect each service's default port before scanning fallbacks. Otherwise
# TensorBoard can claim 6008 as a fallback and make monitor links open it.
protected_default_ports = {args.port}
if not args.disable_tensorboard:
protected_default_ports.add(args.tensorboard_port)
if not args.disable_train_monitor:
protected_default_ports.add(args.train_monitor_port)
reserved_ports: set[int] = set(protected_default_ports)
args.port = ensure_port_available(
args.port, args.port, args.port + 20, "GUI", reserved_ports, preferred_reserved_port=args.port
)
if not args.disable_tensorboard:
args.tensorboard_port = ensure_port_available(
args.tensorboard_port,
args.tensorboard_port,
args.tensorboard_port + 20,
"TensorBoard",
reserved_ports,
preferred_reserved_port=args.tensorboard_port,
)
if not args.disable_train_monitor:
args.train_monitor_port = ensure_port_available(
args.train_monitor_port,
args.train_monitor_port,
args.train_monitor_port + 20,
"Train monitor",
reserved_ports,
preferred_reserved_port=args.train_monitor_port,
)
from mikazuki.update_check import local_version
log.info(f"SD-Trainer Version: {local_version()}")
if args.listen:
args.host = "0.0.0.0"
args.tensorboard_host = "0.0.0.0"
os.environ["MIKAZUKI_HOST"] = args.host
os.environ["MIKAZUKI_PORT"] = str(args.port)
os.environ["MIKAZUKI_TENSORBOARD_HOST"] = args.tensorboard_host
os.environ["MIKAZUKI_TENSORBOARD_PORT"] = str(args.tensorboard_port)
os.environ["TRAIN_MONITOR_HOST"] = args.host
os.environ["TRAIN_MONITOR_PORT"] = str(args.train_monitor_port)
os.environ["TRAIN_MONITOR_ENABLED"] = "0" if args.disable_train_monitor else "1"
os.environ["MIKAZUKI_DEV"] = "1" if args.dev else "0"
if args.browser:
os.environ["MIKAZUKI_BROWSER"] = args.browser
child_processes: list[tuple[str, subprocess.Popen]] = []
try:
if not args.disable_tensorboard:
process = run_tensorboard()
if process is not None:
child_processes.append(("TensorBoard", process))
if not args.disable_train_monitor:
process = run_train_monitor()
if process is not None:
child_processes.append(("train monitor", process))
import uvicorn
log.info(f"Server started at http://{args.host}:{args.port}")
if not args.disable_train_monitor:
log.info(f"Train monitor at http://{args.host}:{args.train_monitor_port}")
else:
log.info("Train monitor disabled (--disable-train-monitor)")
uvicorn.run("mikazuki.app:app", host=args.host, port=args.port, log_level="error", reload=args.dev)
finally:
stop_child_processes(child_processes)
if __name__ == "__main__":
args, _ = parser.parse_known_args()
launch()