From c55bdffe02c595252ba95956ae03979f6ccd7351 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 8 Jan 2024 12:11:38 -0500 Subject: [PATCH] reduce imports and do not load ldm in diffusers --- CHANGELOG.md | 3 +- modules/hypernetworks/hypernetwork.py | 4 +- modules/processing.py | 8 ++- modules/sd_hijack.py | 66 ++++++++++--------- modules/sd_hijack_freeu.py | 2 +- modules/sd_models.py | 24 ++++--- modules/sd_models_config.py | 3 +- modules/sd_samplers.py | 5 +- modules/sd_samplers_common.py | 2 +- modules/shared.py | 1 + .../textual_inversion/textual_inversion.py | 3 +- modules/textual_inversion/ui.py | 9 ++- modules/ui.py | 9 ++- modules/ui_train.py | 7 +- webui.py | 6 +- 15 files changed, 88 insertions(+), 64 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index c77e2f486..82180c64f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,6 @@ # Change Log for SD.Next -## Update for 2023-01-07 +## Update for 2023-01-08 Following-up on a major release, here is a lot more functionality in new Control module and FaceID & IPAdapter modules Plus welcome additions to UI accessibility and flexibility of deployment @@ -64,6 +64,7 @@ And it also includes fixes for all reported issues so far - faster extension load - faster json parsing - faster lora indexing + - reduced module imports - **offline deployment**: allow deployment without git clone for example, you can now deploy a zip of the sdnext folder - **latent upscale**: updated latent upscalers (some are new) diff --git a/modules/hypernetworks/hypernetwork.py b/modules/hypernetworks/hypernetwork.py index f5d192b05..54f82c36a 100644 --- a/modules/hypernetworks/hypernetwork.py +++ b/modules/hypernetworks/hypernetwork.py @@ -11,7 +11,7 @@ from torch import einsum from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_ from einops import rearrange, repeat from ldm.util import default -from modules import devices, processing, sd_models, shared, hashes, sd_hijack_checkpoint, errors +from modules import devices, processing, sd_models, shared, hashes, errors import modules.textual_inversion.dataset from modules.textual_inversion import textual_inversion, ti_logging from modules.textual_inversion.learn_schedule import LearnRateScheduler @@ -447,7 +447,7 @@ def create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure=None, def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_hypernetwork_every, template_filename, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument # images allows training previews to have infotext. Importing it at the top causes a circular import problem. - from modules import images + from modules import images, sd_hijack_checkpoint save_hypernetwork_every = save_hypernetwork_every or 0 create_image_every = create_image_every or 0 diff --git a/modules/processing.py b/modules/processing.py index ca8ff520a..4804a37b1 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -13,8 +13,6 @@ import numpy as np import cv2 from PIL import Image, ImageOps from skimage import exposure -from ldm.data.util import AddMiDaS -from ldm.models.diffusion.ddpm import LatentDepth2ImageDiffusion from einops import repeat, rearrange from blendmodes.blend import blendLayers, BlendType from installer import git_commit @@ -30,7 +28,6 @@ import modules.extra_networks import modules.face_restoration import modules.images as images import modules.styles -import modules.sd_hijack import modules.sd_hijack_freeu import modules.sd_samplers import modules.sd_samplers_common @@ -42,6 +39,9 @@ import modules.generation_parameters_copypaste from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet, hypertile_set +if shared.backend == shared.Backend.ORIGINAL: + import modules.sd_hijack + opt_C = 4 opt_f = 8 debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -289,6 +289,7 @@ class StableDiffusionProcessing: def depth2img_image_conditioning(self, source_image): # Use the AddMiDaS helper to Format our source image to suit the MiDaS model + from ldm.data.util import AddMiDaS transformer = AddMiDaS(model_type="dpt_hybrid") transformed = transformer({"jpg": rearrange(source_image[0], "c h w -> h w c")}) midas_in = torch.from_numpy(transformed["midas_in"][None, ...]).to(device=shared.device) @@ -352,6 +353,7 @@ class StableDiffusionProcessing: return latent_image.new_zeros(latent_image.shape[0], 5, 1, 1) def img2img_image_conditioning(self, source_image, latent_image, image_mask=None): + from ldm.models.diffusion.ddpm import LatentDepth2ImageDiffusion source_image = devices.cond_cast_float(source_image) # HACK: Using introspection as the Depth2Image model doesn't appear to uniquely # identify itself with a field common to all models. The conditioning_key is also hybrid. diff --git a/modules/sd_hijack.py b/modules/sd_hijack.py index a1975b0c0..816c35d76 100644 --- a/modules/sd_hijack.py +++ b/modules/sd_hijack.py @@ -1,19 +1,25 @@ from types import MethodType, SimpleNamespace +import io +import contextlib import torch from torch.nn.functional import silu -import ldm.modules.attention -import ldm.modules.distributions.distributions -import ldm.modules.diffusionmodules.model -import ldm.modules.diffusionmodules.openaimodel -import ldm.models.diffusion.ddim -import ldm.models.diffusion.plms -import ldm.modules.encoders.modules + +from modules import shared +shared.log.debug('Importing LDM') +stdout = io.StringIO() +with contextlib.redirect_stdout(stdout): + import ldm.modules.attention + import ldm.modules.distributions.distributions + import ldm.modules.diffusionmodules.model + import ldm.modules.diffusionmodules.openaimodel + import ldm.models.diffusion.ddim + import ldm.models.diffusion.plms + import ldm.modules.encoders.modules import modules.textual_inversion.textual_inversion -from modules import devices, sd_hijack_optimizations, shared -from modules.hypernetworks import hypernetwork -from modules.shared import opts +from modules import devices, sd_hijack_optimizations from modules import sd_hijack_clip, sd_hijack_open_clip, sd_hijack_unet, sd_hijack_xlmr, xlmr +from modules.hypernetworks import hypernetwork attention_CrossAttention_forward = ldm.modules.attention.CrossAttention.forward diffusionmodules_model_nonlinearity = ldm.modules.diffusionmodules.model.nonlinearity @@ -37,39 +43,39 @@ def apply_optimizations(): optimization_method = None can_use_sdp = hasattr(torch.nn.functional, "scaled_dot_product_attention") and callable(torch.nn.functional.scaled_dot_product_attention) if devices.device == torch.device("cpu"): - if opts.cross_attention_optimization == "Scaled-Dot-Product": + if shared.opts.cross_attention_optimization == "Scaled-Dot-Product": shared.log.warning("Cross-attention: Scaled dot product is not available on CPU") can_use_sdp = False - if opts.cross_attention_optimization == "xFormers": + if shared.opts.cross_attention_optimization == "xFormers": shared.log.warning("Cross-attention: xFormers is not available on CPU") shared.xformers_available = False - shared.log.info(f"Cross-attention: optimization={opts.cross_attention_optimization} options={opts.cross_attention_options}") - if opts.cross_attention_optimization == "Disabled": + shared.log.info(f"Cross-attention: optimization={shared.opts.cross_attention_optimization} options={shared.opts.cross_attention_options}") + if shared.opts.cross_attention_optimization == "Disabled": optimization_method = 'none' - if can_use_sdp and opts.cross_attention_optimization == "Scaled-Dot-Product" and 'SDP disable memory attention' in opts.cross_attention_options: + if can_use_sdp and shared.opts.cross_attention_optimization == "Scaled-Dot-Product" and 'SDP disable memory attention' in shared.opts.cross_attention_options: ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_no_mem_attention_forward ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_no_mem_attnblock_forward optimization_method = 'sdp-no-mem' - elif can_use_sdp and opts.cross_attention_optimization == "Scaled-Dot-Product": + elif can_use_sdp and shared.opts.cross_attention_optimization == "Scaled-Dot-Product": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_attention_forward ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_attnblock_forward optimization_method = 'sdp' - if shared.xformers_available and opts.cross_attention_optimization == "xFormers": + if shared.xformers_available and shared.opts.cross_attention_optimization == "xFormers": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.xformers_attention_forward ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.xformers_attnblock_forward optimization_method = 'xformers' - if opts.cross_attention_optimization == "Sub-quadratic": + if shared.opts.cross_attention_optimization == "Sub-quadratic": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.sub_quad_attention_forward ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sub_quad_attnblock_forward optimization_method = 'sub-quadratic' - if opts.cross_attention_optimization == "Split attention": + if shared.opts.cross_attention_optimization == "Split attention": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.split_cross_attention_forward_v1 optimization_method = 'v1' - if opts.cross_attention_optimization == "InvokeAI's": + if shared.opts.cross_attention_optimization == "InvokeAI's": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.split_cross_attention_forward_invokeAI optimization_method = 'invokeai' - if opts.cross_attention_optimization == "Doggettx's": + if shared.opts.cross_attention_optimization == "Doggettx's": ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.split_cross_attention_forward ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.cross_attention_attnblock_forward optimization_method = 'doggettx' @@ -148,7 +154,7 @@ class StableDiffusionModelHijack: embedding_db = modules.textual_inversion.textual_inversion.EmbeddingDatabase() def __init__(self): - self.embedding_db.add_embedding_dir(opts.embeddings_dir) + self.embedding_db.add_embedding_dir(shared.opts.embeddings_dir) def hijack(self, m): if type(m.cond_stage_model) == xlmr.BertSeriesModelWithTransformation: @@ -169,7 +175,7 @@ class StableDiffusionModelHijack: if m.cond_stage_key == "edit": sd_hijack_unet.hijack_ddpm_edit() - if opts.ipex_optimize and shared.backend == shared.Backend.ORIGINAL: + if shared.opts.ipex_optimize and shared.backend == shared.Backend.ORIGINAL: try: import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import m.model.training = False @@ -178,22 +184,22 @@ class StableDiffusionModelHijack: except Exception as err: shared.log.warning(f"IPEX Optimize not supported: {err}") - if (opts.cuda_compile or opts.cuda_compile_vae or opts.cuda_compile_upscaler) and shared.opts.cuda_compile_backend != 'none' and shared.backend == shared.Backend.ORIGINAL: + if (shared.opts.cuda_compile or shared.opts.cuda_compile_vae or shared.opts.cuda_compile_upscaler) and shared.opts.cuda_compile_backend != 'none' and shared.backend == shared.Backend.ORIGINAL: try: import logging - shared.log.info(f"Compiling pipeline={m.model.__class__.__name__} mode={opts.cuda_compile_backend}") + shared.log.info(f"Compiling pipeline={m.model.__class__.__name__} mode={shared.opts.cuda_compile_backend}") import torch._dynamo # pylint: disable=unused-import,redefined-outer-name - log_level = logging.WARNING if opts.cuda_compile_verbose else logging.CRITICAL # pylint: disable=protected-access + log_level = logging.WARNING if shared.opts.cuda_compile_verbose else logging.CRITICAL # pylint: disable=protected-access if hasattr(torch, '_logging'): torch._logging.set_logs(dynamo=log_level, aot=log_level, inductor=log_level) # pylint: disable=protected-access - torch._dynamo.config.verbose = opts.cuda_compile_verbose # pylint: disable=protected-access - torch._dynamo.config.suppress_errors = opts.cuda_compile_errors # pylint: disable=protected-access + torch._dynamo.config.verbose = shared.opts.cuda_compile_verbose # pylint: disable=protected-access + torch._dynamo.config.suppress_errors = shared.opts.cuda_compile_errors # pylint: disable=protected-access torch.backends.cudnn.benchmark = True - if opts.cuda_compile_backend == 'hidet': + if shared.opts.cuda_compile_backend == 'hidet': import hidet # pylint: disable=import-error hidet.torch.dynamo_config.use_tensor_core(True) hidet.torch.dynamo_config.search_space(2) - m.model = torch.compile(m.model, mode=opts.cuda_compile_mode, backend=opts.cuda_compile_backend, fullgraph=opts.cuda_compile_fullgraph, dynamic=False) + m.model = torch.compile(m.model, mode=shared.opts.cuda_compile_mode, backend=shared.opts.cuda_compile_backend, fullgraph=shared.opts.cuda_compile_fullgraph, dynamic=False) shared.log.info("Model complilation done.") except Exception as err: shared.log.warning(f"Model compile not supported: {err}") diff --git a/modules/sd_hijack_freeu.py b/modules/sd_hijack_freeu.py index 430c764dc..a1e560485 100644 --- a/modules/sd_hijack_freeu.py +++ b/modules/sd_hijack_freeu.py @@ -2,7 +2,6 @@ import math import functools import torch from modules import shared -from modules.sd_hijack_unet import th # based on # official params are b1,b2,s1,s2 @@ -129,6 +128,7 @@ def ratio_to_region(width: float, offset: float, n: int): def apply_freeu(p, backend_original): + from modules.sd_hijack_unet import th global state_enabled # pylint: disable=global-statement global cat_original # pylint: disable=global-statement if backend_original: diff --git a/modules/sd_models.py b/modules/sd_models.py index 51a3b4adc..b1d29a1ad 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -18,9 +18,8 @@ import diffusers from omegaconf import OmegaConf import tomesd from transformers import logging as transformers_logging -import ldm.modules.midas as midas from ldm.util import instantiate_from_config -from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, errors, hashes, sd_models_config, sd_models_compile, sd_hijack_inpainting +from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, errors, hashes, sd_models_config, sd_models_compile from modules.timer import Timer from modules.memstats import memory_stats from modules.paths import models_path, script_path @@ -123,7 +122,8 @@ def setup_model(): if not os.path.exists(model_path): os.makedirs(model_path, exist_ok=True) list_models() - enable_midas_autodownload() + if shared.backend == shared.Backend.ORIGINAL: + enable_midas_autodownload() def checkpoint_tiles(use_short=False): # pylint: disable=unused-argument @@ -492,29 +492,30 @@ def enable_midas_autodownload(): This function applies a wrapper to download the model to the correct location automatically. """ + import ldm.modules.midas.api midas_path = os.path.join(paths.models_path, 'midas') - for k, v in midas.api.ISL_PATHS.items(): + for k, v in ldm.modules.midas.api.ISL_PATHS.items(): file_name = os.path.basename(v) - midas.api.ISL_PATHS[k] = os.path.join(midas_path, file_name) + ldm.modules.midas.api.ISL_PATHS[k] = os.path.join(midas_path, file_name) midas_urls = { "dpt_large": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_large-midas-2f21e586.pt", "dpt_hybrid": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_hybrid-midas-501f0c75.pt", "midas_v21": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21-f6b98070.pt", "midas_v21_small": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21_small-70d6b9c8.pt", } - midas.api.load_model_inner = midas.api.load_model + ldm.modules.midas.api.load_model_inner = ldm.modules.midas.api.load_model def load_model_wrapper(model_type): - path = midas.api.ISL_PATHS[model_type] + path = ldm.modules.midas.api.ISL_PATHS[model_type] if not os.path.exists(path): if not os.path.exists(midas_path): mkdir(midas_path) shared.log.info(f"Downloading midas model weights for {model_type} to {path}") request.urlretrieve(midas_urls[model_type], path) shared.log.info(f"{model_type} downloaded") - return midas.api.load_model_inner(model_type) + return ldm.modules.midas.api.load_model_inner(model_type) - midas.api.load_model = load_model_wrapper + ldm.modules.midas.api.load_model = load_model_wrapper def repair_config(sd_config): @@ -1108,7 +1109,10 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, current_checkpoint_info = model_data.sd_refiner.sd_checkpoint_info unload_model_weights(op=op) - sd_hijack_inpainting.do_inpainting_hijack() + if shared.backend == shared.Backend.ORIGINAL: + from modules import sd_hijack_inpainting + sd_hijack_inpainting.do_inpainting_hijack() + devices.set_cuda_params() if already_loaded_state_dict is not None: state_dict = already_loaded_state_dict diff --git a/modules/sd_models_config.py b/modules/sd_models_config.py index 3df76caa5..d2bc0ac4c 100644 --- a/modules/sd_models_config.py +++ b/modules/sd_models_config.py @@ -2,7 +2,7 @@ import os import torch -from modules import paths, sd_disable_initialization, devices +from modules import paths, devices sd_repo_configs_path = 'configs' config_default = paths.sd_default_config @@ -21,6 +21,7 @@ def is_using_v_parameterization_for_sd2(state_dict): """ Detects whether unet in state_dict is using v-parameterization. Returns True if it is. You're welcome. """ + from modules import sd_disable_initialization import ldm.modules.diffusionmodules.openaimodel device = devices.cpu diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index 891d83eea..e1cfd0bf6 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -1,8 +1,7 @@ import os -from modules import sd_samplers_compvis, sd_samplers_kdiffusion, sd_samplers_diffusers, shared +from modules import shared from modules.sd_samplers_common import samples_to_image_grid, sample_to_image # pylint: disable=unused-import - debug = shared.log.trace if os.environ.get('SD_SAMPLER_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: SAMPLER') all_samplers = [] @@ -20,8 +19,10 @@ def list_samplers(backend_name = shared.backend): global samplers_for_img2img # pylint: disable=global-statement global samplers_map # pylint: disable=global-statement if backend_name == shared.Backend.ORIGINAL: + from modules import sd_samplers_compvis, sd_samplers_kdiffusion all_samplers = [*sd_samplers_compvis.samplers_data_compvis, *sd_samplers_kdiffusion.samplers_data_k_diffusion] else: + from modules import sd_samplers_diffusers all_samplers = [*sd_samplers_diffusers.samplers_data_diffusers] all_samplers_map = {x.name: x for x in all_samplers} samplers = all_samplers diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index e749aa378..cf4f8b8fb 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -61,7 +61,7 @@ def single_sample_to_image(sample, approximation=None): try: if x_sample.dtype == torch.bfloat16: - x_sample.to(torch.float16) + x_sample.to(torch.float16) transform = T.ToPILImage() image = transform(x_sample) except Exception as e: diff --git a/modules/shared.py b/modules/shared.py index 287b57370..8886d4b82 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -846,6 +846,7 @@ max_workers = 4 if devices.backend == "directml": directml_do_hijack() + class TotalTQDM: # compatibility with previous global-tqdm # import tqdm def __init__(self): diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 0688e37be..5d3715513 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -9,7 +9,7 @@ import safetensors.torch import numpy as np from PIL import Image, PngImagePlugin from torch.utils.tensorboard import SummaryWriter -from modules import shared, devices, sd_hijack, processing, sd_models, images, sd_hijack_checkpoint, errors +from modules import shared, devices, processing, sd_models, images, errors import modules.textual_inversion.dataset from modules.textual_inversion.learn_schedule import LearnRateScheduler from modules.textual_inversion.image_embedding import embedding_to_b64, embedding_from_b64, insert_image_data_embed, extract_image_data_embed, caption_image_overlay @@ -415,6 +415,7 @@ def validate_train_inputs(model_name, learn_rate, batch_size, gradient_step, dat def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_embedding_every, template_filename, save_image_with_stored_embedding, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument + from modules import sd_hijack, sd_hijack_checkpoint shared.log.debug(f'train_embedding: embedding_name={embedding_name}|learn_rate={learn_rate}|batch_size={batch_size}|gradient_step={gradient_step}|data_root={data_root}|log_directory={log_directory}|training_width={training_width}|training_height={training_height}|varsize={varsize}|steps={steps}|clip_grad_mode={clip_grad_mode}|clip_grad_value={clip_grad_value}|shuffle_tags={shuffle_tags}|tag_drop_out={tag_drop_out}|latent_sampling_method={latent_sampling_method}|use_weight={use_weight}|create_image_every={create_image_every}|save_embedding_every={save_embedding_every}|template_filename={template_filename}|save_image_with_stored_embedding={save_image_with_stored_embedding}|preview_from_txt2img={preview_from_txt2img}|preview_prompt={preview_prompt}|preview_negative_prompt={preview_negative_prompt}|preview_steps={preview_steps}|preview_sampler_index={preview_sampler_index}|preview_cfg_scale={preview_cfg_scale}|preview_seed={preview_seed}|preview_width={preview_width}|preview_height={preview_height}') save_embedding_every = save_embedding_every or 0 diff --git a/modules/textual_inversion/ui.py b/modules/textual_inversion/ui.py index 76367b772..1f488848c 100644 --- a/modules/textual_inversion/ui.py +++ b/modules/textual_inversion/ui.py @@ -2,10 +2,11 @@ import html import gradio as gr import modules.textual_inversion.textual_inversion import modules.textual_inversion.preprocess -from modules import sd_hijack, shared +from modules import shared def create_embedding(name, initialization_text, nvpt, overwrite_old): + from modules import sd_hijack filename = modules.textual_inversion.textual_inversion.create_embedding(name, nvpt, overwrite_old, init_text=initialization_text) sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings() return gr.Dropdown.update(choices=sorted(sd_hijack.model_hijack.embedding_db.word_embeddings.keys())), f"Created: {filename}", "" @@ -17,6 +18,7 @@ def preprocess(*args): def train_embedding(*args): + from modules import sd_hijack assert not shared.cmd_opts.lowvram, 'Training models with lowvram not possible' apply_optimizations = False try: @@ -25,8 +27,9 @@ def train_embedding(*args): embedding, filename = modules.textual_inversion.textual_inversion.train_embedding(*args) res = f"Training {'interrupted' if shared.state.interrupted else 'finished'} at {embedding.step} steps. Embedding saved to {html.escape(filename)}" return res, "" - except Exception: - raise + except Exception as e: + shared.log.error(f"Exception in train_embedding: {e}") + raise RuntimeError from e finally: if not apply_optimizations: sd_hijack.apply_optimizations() diff --git a/modules/ui.py b/modules/ui.py index 8780b40bb..9397a2ef6 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -9,8 +9,6 @@ from modules.ui_components import FormRow from modules.paths import script_path, data_path # pylint: disable=unused-import from modules.dml import directml_override_opts import modules.scripts -import modules.textual_inversion.ui -import modules.hypernetworks.ui import modules.errors @@ -143,9 +141,10 @@ def create_ui(startup_timer = None): timer.startup.record("ui-extras") with gr.Blocks(analytics_enabled=False) as train_interface: - from modules import ui_train - ui_train.create_ui() - timer.startup.record("ui-train") + if shared.backend == shared.Backend.ORIGINAL: + from modules import ui_train + ui_train.create_ui() + timer.startup.record("ui-train") with gr.Blocks(analytics_enabled=False) as models_interface: from modules import ui_models diff --git a/modules/ui_train.py b/modules/ui_train.py index 4676ab6fb..ae07b5d16 100644 --- a/modules/ui_train.py +++ b/modules/ui_train.py @@ -5,11 +5,11 @@ from modules.ui_components import FormRow from modules.ui_common import create_refresh_button from modules.ui_sections import create_sampler_inputs from modules.call_queue import wrap_gradio_gpu_call -from modules.textual_inversion import textual_inversion -import modules.errors def create_ui(): + from modules.textual_inversion import textual_inversion + import modules.hypernetworks.ui dummy_component = gr.Label(visible=False) with gr.Row(elem_id="train_tab"): @@ -106,12 +106,13 @@ def create_ui(): process_multicrop_objective = gr.Radio(["Maximize area", "Minimize error"], value="Maximize area", label="Resizing objective") process_multicrop_threshold = gr.Slider(minimum=0, maximum=1, step=0.01, label="Error threshold", value=0.1) + from modules.textual_inversion import ui process_split.change(fn=lambda show: gr_show(show), inputs=[process_split], outputs=[process_split_extra_row]) process_focal_crop.change(fn=lambda show: gr_show(show), inputs=[process_focal_crop], outputs=[process_focal_crop_row]) process_multicrop.change(fn=lambda show: gr_show(show), inputs=[process_multicrop], outputs=[process_multicrop_col]) process_stop.click(fn=lambda: shared.state.interrupt(), inputs=[], outputs=[]) process_run.click( - fn=wrap_gradio_gpu_call(modules.textual_inversion.ui.preprocess, extra_outputs=[gr.update()]), + fn=wrap_gradio_gpu_call(ui.preprocess, extra_outputs=[gr.update()]), _js="startTrainMonitor", inputs=[ dummy_component, diff --git a/webui.py b/webui.py index 8b87989b8..1ab72115d 100644 --- a/webui.py +++ b/webui.py @@ -20,7 +20,6 @@ import modules.devices import modules.sd_samplers import modules.lowvram import modules.scripts -import modules.sd_hijack import modules.sd_models import modules.sd_vae import modules.progress @@ -57,6 +56,11 @@ fastapi_args = { "deepLinking": False, } } + +if shared.backend == shared.Backend.ORIGINAL: + import modules.sd_hijack + timer.startup.record("ldm") + modules.loader.initialized = True