From 53bb05abf672c4fcf4b19cf90495ddca2c74a14b Mon Sep 17 00:00:00 2001 From: Haoming Date: Sat, 27 Dec 2025 13:22:53 +0800 Subject: [PATCH] timer --- modules/errors.py | 35 +++++++---------- modules/initialize.py | 65 ++++++++++++++++++++++++------- modules/initialize_util.py | 35 ++++++++--------- modules_forge/initialization.py | 69 ++++++++++----------------------- webui.py | 41 +++++++++++--------- 5 files changed, 126 insertions(+), 119 deletions(-) diff --git a/modules/errors.py b/modules/errors.py index 151c4655..2f83b0aa 100644 --- a/modules/errors.py +++ b/modules/errors.py @@ -2,7 +2,6 @@ import sys import textwrap import traceback - exception_records = [] @@ -55,10 +54,10 @@ def print_error_explanation(message): lines = message.strip().split("\n") max_len = max([len(x) for x in lines]) - print('=' * max_len, file=sys.stderr) + print("=" * max_len, file=sys.stderr) for line in lines: print(line, file=sys.stderr) - print('=' * max_len, file=sys.stderr) + print("=" * max_len, file=sys.stderr) def display(e: Exception, task, *, full_traceback=False): @@ -71,13 +70,6 @@ def display(e: Exception, task, *, full_traceback=False): te.stack = traceback.StackSummary(traceback.extract_stack()[:-2] + te.stack) print(*te.format(), sep="", file=sys.stderr) - message = str(e) - if "copying a param with shape torch.Size([640, 1024]) from checkpoint, the shape in current model is torch.Size([640, 768])" in message: - print_error_explanation(""" -The most likely cause of this is you are trying to load Stable Diffusion 2.0 model without specifying its config file. -See https://github.com/AUTOMATIC1111/stable-diffusion-webui/wiki/Features#stable-diffusion-20 for how to solve this. - """) - already_displayed = {} @@ -103,20 +95,22 @@ def run(code, task): def check_versions(): import gradio import torch - from modules import shared from packaging import version - expected_torch = "2.7.0" - expected_xformers = "0.0.30" + from modules import shared + + expected_torch = "2.9.0" + expected_xformers = "0.0.33" expected_gradio = "4.40.0" + _outdated = False + if version.parse(torch.__version__) < version.parse(expected_torch): + _outdated = True print_error_explanation( f""" - You are running torch {torch.__version__}, which is really outdated + You are running torch {torch.__version__}, which is really outdated. To install the latest version, run with commandline flag --reinstall-torch. - - Use --skip-version-check commandline argument to disable this check. """.strip() ) @@ -124,22 +118,23 @@ def check_versions(): import xformers if version.parse(xformers.__version__) < version.parse(expected_xformers): + _outdated = True print_error_explanation( f""" You are running xformers {xformers.__version__}, which is really outdated. To install the latest version, run with commandline flag --reinstall-xformers. - - Use --skip-version-check commandline argument to disable this check. """.strip() ) if version.parse(gradio.__version__) < version.parse(expected_gradio): + _outdated = True print_error_explanation( f""" You are running gradio {gradio.__version__}. This program was built on gradio {expected_gradio}. Using a different version of gradio is likely to break the program. - - Use --skip-version-check commandline argument to disable this check. """.strip() ) + + if _outdated: + print("\nUse --skip-version-check commandline argument to disable the version check(s).\n") diff --git a/modules/initialize.py b/modules/initialize.py index 1e358e9e..02432efa 100644 --- a/modules/initialize.py +++ b/modules/initialize.py @@ -1,8 +1,8 @@ import importlib import logging +import os import sys import warnings -import os from modules.timer import startup_timer @@ -11,36 +11,61 @@ def shush(): logging.getLogger("torch.distributed.nn").setLevel(logging.ERROR) logging.getLogger("xformers").addFilter(lambda record: "triton" not in record.getMessage().lower()) warnings.filterwarnings(action="ignore", category=DeprecationWarning, module="pytorch_lightning") + warnings.filterwarnings(action="ignore", category=UserWarning, module="torchvision.transforms.functional_tensor") warnings.filterwarnings(action="ignore", category=UserWarning, module="torchvision") + startup_timer.record("filter logging") + + +def shush_nunchaku(): + _original = logging.basicConfig + logging.basicConfig = lambda *args, **kwargs: None + + try: + import nunchaku + except ImportError: + pass + + logging.basicConfig = _original + startup_timer.record("bypass basicConfig") def imports(): os.environ.setdefault("GRADIO_ANALYTICS_ENABLED", "False") import gradio # noqa: F401 + startup_timer.record("import gradio") - from modules import paths, timer, errors # noqa: F401 + from modules import errors, paths, timer # noqa: F401 + startup_timer.record("setup paths") from modules import shared_init - shared_init.initialize() - startup_timer.record("initialize shared") - from modules import processing, gradio_extensions, ui # noqa: F401 - startup_timer.record("other imports") + shared_init.initialize() + startup_timer.record("shared init") + + from modules import gradio_extensions, processing, ui # noqa: F401 + + startup_timer.record("misc. imports") def check_versions(): from modules.shared_cmd_options import cmd_opts - if not cmd_opts.skip_version_check: - from modules import errors - errors.check_versions() + if cmd_opts.skip_version_check: + return + + from modules import errors + + errors.check_versions() + + startup_timer.record("version check") def initialize(): from modules import initialize_util + initialize_util.fix_torch_version() initialize_util.fix_asyncio_event_loop_policy() initialize_util.validate_tls_options() @@ -48,17 +73,17 @@ def initialize(): initialize_util.configure_opts_onchange() from modules import sd_models - sd_models.setup_model() - startup_timer.record("setup SD model") - from modules.shared_cmd_options import cmd_opts + sd_models.setup_model() from modules import codeformer_model - warnings.filterwarnings(action="ignore", category=UserWarning, module="torchvision.transforms.functional_tensor") + from modules.shared_cmd_options import cmd_opts + codeformer_model.setup_model(cmd_opts.codeformer_models_path) startup_timer.record("setup codeformer") from modules import gfpgan_model + gfpgan_model.setup_model(cmd_opts.gfpgan_models_path) startup_timer.record("setup gfpgan") @@ -69,31 +94,36 @@ def initialize_rest(*, reload_script_modules=False): """ Called both from initialize() and when reloading the webui. """ + from modules import sd_samplers from modules.shared_cmd_options import cmd_opts - from modules import sd_samplers sd_samplers.set_samplers() startup_timer.record("set samplers") from modules import extensions + extensions.list_extensions() startup_timer.record("list extensions") from modules import initialize_util + initialize_util.restore_config_state_file() startup_timer.record("restore config state file") - from modules import shared, upscaler, scripts + from modules import scripts, shared, upscaler + if cmd_opts.ui_debug_mode: shared.sd_upscalers = upscaler.UpscalerLanczos().scalers scripts.load_scripts() return from modules import sd_models + sd_models.list_models() startup_timer.record("list SD models") from modules import localization + localization.list_localizations(cmd_opts.localizations_dir) startup_timer.record("list localizations") @@ -106,21 +136,26 @@ def initialize_rest(*, reload_script_modules=False): startup_timer.record("reload script modules") from modules import modelloader + modelloader.load_upscalers() startup_timer.record("load upscalers") from modules import sd_vae + sd_vae.refresh_vae_list() startup_timer.record("refresh VAE") from modules import sd_unet + sd_unet.list_unets() startup_timer.record("scripts list_unets") from modules import ui_extra_networks + ui_extra_networks.initialize() ui_extra_networks.register_default_pages() from modules import extra_networks + extra_networks.initialize() startup_timer.record("initialize extra networks") diff --git a/modules/initialize_util.py b/modules/initialize_util.py index d1503698..5855ce9e 100644 --- a/modules/initialize_util.py +++ b/modules/initialize_util.py @@ -1,8 +1,8 @@ import json import os +import re import signal import sys -import re import starlette @@ -21,20 +21,19 @@ def gradio_server_name(): def fix_torch_version(): import torch - # Truncate version number of nightly/local build of PyTorch to not cause exceptions with CodeFormer or Safetensors + # truncate version number of nightly/local build of PyTorch if ".dev" in torch.__version__ or "+git" in torch.__version__: torch.__long_version__ = torch.__version__ - torch.__version__ = re.search(r'[\d.]+[\d]', torch.__version__).group(0) + torch.__version__ = re.search(r"[\d.]+[\d]", torch.__version__).group(0) def fix_asyncio_event_loop_policy(): """ - The default `asyncio` event loop policy only automatically creates - event loops in the main threads. Other threads must create event - loops explicitly or `asyncio.get_event_loop` (and therefore - `.IOLoop.current`) will fail. Installing this policy allows event - loops to be created automatically on any thread, matching the - behavior of Tornado versions prior to 5.0 (or 5.0 on Python 2). + The default `asyncio` event loop policy only automatically creates + event loops in the main threads. Other threads must create event + loops explicitly or `asyncio.get_event_loop` and `.IOLoop.current` + will fail. Installing this policy allows event loops to be created + automatically on any thread, matching the behavior of Tornado prior to 5.0 """ import asyncio @@ -68,7 +67,7 @@ def fix_asyncio_event_loop_policy(): def restore_config_state_file(): - from modules import shared, config_states + from modules import config_states, shared config_state_file = shared.opts.restore_config_state_file if config_state_file == "": @@ -117,18 +116,18 @@ def get_gradio_auth_creds(): s = s.strip() if not s: return None - return tuple(s.split(':', 1)) + return tuple(s.split(":", 1)) if cmd_opts.gradio_auth: - for cred in cmd_opts.gradio_auth.split(','): + for cred in cmd_opts.gradio_auth.split(","): cred = process_credential_line(cred) if cred: yield cred if cmd_opts.gradio_auth_path: - with open(cmd_opts.gradio_auth_path, 'r', encoding="utf8") as file: + with open(cmd_opts.gradio_auth_path, "r", encoding="utf8") as file: for line in file.readlines(): - for cred in line.strip().split(','): + for cred in line.strip().split(","): cred = process_credential_line(cred) if cred: yield cred @@ -151,12 +150,12 @@ def dumpstacks(): def configure_sigint_handler(): - # make the program just exit at ctrl+c without waiting for anything + # make the program just exit at Ctrl + C without waiting for anything from modules import shared def sigint_handler(sig, frame): - print(f'Interrupted with signal {sig} in {frame}') + print(f"Interrupted with signal {sig} in {frame}") if shared.opts.dump_stacks_on_signal: dumpstacks() @@ -187,6 +186,7 @@ def setup_middleware(app): def configure_cors_middleware(app): from starlette.middleware.cors import CORSMiddleware + from modules.shared_cmd_options import cmd_opts cors_options = { @@ -195,9 +195,8 @@ def configure_cors_middleware(app): "allow_credentials": True, } if cmd_opts.cors_allow_origins: - cors_options["allow_origins"] = cmd_opts.cors_allow_origins.split(',') + cors_options["allow_origins"] = cmd_opts.cors_allow_origins.split(",") if cmd_opts.cors_allow_origins_regex: cors_options["allow_origin_regex"] = cmd_opts.cors_allow_origins_regex app.user_middleware.insert(0, starlette.middleware.Middleware(CORSMiddleware, **cors_options)) - diff --git a/modules_forge/initialization.py b/modules_forge/initialization.py index 2023a376..480635e8 100644 --- a/modules_forge/initialization.py +++ b/modules_forge/initialization.py @@ -1,51 +1,13 @@ import os import sys +from modules.timer import startup_timer + INITIALIZED = False -MONITOR_MODEL_MOVING = False -def monitor_module_moving(): - if not MONITOR_MODEL_MOVING: - return - - import torch - import traceback - - old_to = torch.nn.Module.to - - def new_to(*args, **kwargs): - traceback.print_stack() - print("Model Movement") - - return old_to(*args, **kwargs) - - torch.nn.Module.to = new_to - return - - -def fix_logging(): - import logging - - logging.getLogger("nunchaku.caching.teacache").addFilter(lambda record: "deprecated" not in record.getMessage().lower()) - logging.getLogger("nunchaku.models.pulid.pulid_forward").addFilter(lambda record: "deprecated" not in record.getMessage().lower()) - logging.getLogger("nunchaku.models.transformers.transformer_flux").addFilter(lambda record: "deprecated" not in record.getMessage().lower()) - - _original = logging.basicConfig - - logging.basicConfig = lambda *args, **kwargs: None - - try: - import nunchaku - except ImportError: - pass - - logging.basicConfig = _original - - -def initialize_forge(startup_timer): +def initialize_forge(): global INITIALIZED - if INITIALIZED: return @@ -63,23 +25,30 @@ def initialize_forge(startup_timer): from modules_forge.cuda_malloc import try_cuda_malloc try_cuda_malloc() + startup_timer.record("cuda_malloc") from backend import memory_management + + startup_timer.record("memory_management") + + import pytorch_lightning # noqa: F401 import torch import torchvision # noqa: F401 - import pytorch_lightning # noqa: F401 startup_timer.record("import torch") - monitor_module_moving() - device = memory_management.get_torch_device() torch.zeros((1, 1)).to(device, torch.float32) memory_management.soft_empty_cache() + startup_timer.record("tensor warmup") + from backend import stream + print("CUDA Using Stream:", stream.should_use_stream()) + startup_timer.record("stream") + from modules_forge.shared import diffusers_dir if "HF_HOME" not in os.environ: @@ -97,12 +66,16 @@ def initialize_forge(startup_timer): if "HF_HUB_CACHE" not in os.environ: os.environ["HF_HUB_CACHE"] = diffusers_dir - import modules_forge.patch_basic - modules_forge.patch_basic.patch_all_basics() + startup_timer.record("diffusers_dir") - fix_logging() + from modules_forge import patch_basic + + patch_basic.patch_all_basics() + + startup_timer.record("patch basics") from backend.huggingface import process + process() - startup_timer.record("forge init") + startup_timer.record("decompress tokenizers") diff --git a/webui.py b/webui.py index c3c79216..8b4d4da4 100644 --- a/webui.py +++ b/webui.py @@ -2,25 +2,25 @@ from __future__ import annotations import os import time +from threading import Thread from fastapi import Request from fastapi.encoders import jsonable_encoder from fastapi.responses import JSONResponse -from modules import timer -from modules import initialize_util -from modules import initialize -from threading import Thread -from modules_forge.initialization import initialize_forge +from modules import initialize, initialize_util, timer from modules_forge import main_thread - +from modules_forge.initialization import initialize_forge startup_timer = timer.startup_timer startup_timer.record("launcher") initialize.shush() -initialize_forge(startup_timer) +with startup_timer.subcategory("forge init"): + initialize_forge() + +initialize.shush_nunchaku() initialize.imports() @@ -50,6 +50,7 @@ def create_api(app): def api_only_worker(): from fastapi import FastAPI + from modules.shared_cmd_options import cmd_opts app = FastAPI(exception_handlers={Exception: _handle_exception}) @@ -57,15 +58,12 @@ def api_only_worker(): api = create_api(app) from modules import script_callbacks + script_callbacks.before_ui_callback() script_callbacks.app_started_callback(None, app) print(f"Startup time: {startup_timer.summary()}.") - api.launch( - server_name=initialize_util.gradio_server_name(), - port=cmd_opts.port if cmd_opts.port else 7861, - root_path=f"/{cmd_opts.subpath}" if cmd_opts.subpath else "" - ) + api.launch(server_name=initialize_util.gradio_server_name(), port=cmd_opts.port if cmd_opts.port else 7861, root_path=f"/{cmd_opts.subpath}" if cmd_opts.subpath else "") def webui_worker(): @@ -73,7 +71,14 @@ def webui_worker(): launch_api = cmd_opts.api - from modules import shared, ui_tempdir, script_callbacks, ui, progress, ui_extra_networks + from modules import ( + progress, + script_callbacks, + shared, + ui, + ui_extra_networks, + ui_tempdir, + ) while 1: if shared.opts.clean_temp_dir_at_start: @@ -92,7 +97,7 @@ def webui_worker(): gradio_auth_creds = list(initialize_util.get_gradio_auth_creds()) or None auto_launch_browser = False - if os.getenv('SD_WEBUI_RESTARTING') != '1': + if os.getenv("SD_WEBUI_RESTARTING") != "1": if shared.opts.auto_launch_browser == "Remote" or cmd_opts.autolaunch: auto_launch_browser = True elif shared.opts.auto_launch_browser == "Local": @@ -126,7 +131,7 @@ def webui_worker(): # an attacker to trick the user into opening a malicious HTML page, which makes a request to the # running web ui and do whatever the attacker wants, including installing an extension and # running its code. We disable this here. Suggested by RyotaK. - app.user_middleware = [x for x in app.user_middleware if x.cls.__name__ != 'CORSMiddleware'] + app.user_middleware = [x for x in app.user_middleware if x.cls.__name__ != "CORSMiddleware"] initialize_util.setup_middleware(app) @@ -155,7 +160,7 @@ def webui_worker(): else: print(f"Unknown server command: {server_command}") except KeyboardInterrupt: - print('Caught KeyboardInterrupt, stopping...') + print("Caught KeyboardInterrupt, stopping...") server_command = "stop" if server_command == "stop": @@ -165,9 +170,9 @@ def webui_worker(): break # disable auto launch webui in browser for subsequent UI Reload - os.environ.setdefault('SD_WEBUI_RESTARTING', '1') + os.environ.setdefault("SD_WEBUI_RESTARTING", "1") - print('Restarting UI...') + print("Restarting UI...") shared.demo.close() time.sleep(0.5) startup_timer.reset()