mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
reduce imports and do not load ldm in diffusers
This commit is contained in:
+2
-1
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
+36
-30
@@ -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}")
|
||||
|
||||
@@ -2,7 +2,6 @@ import math
|
||||
import functools
|
||||
import torch
|
||||
from modules import shared
|
||||
from modules.sd_hijack_unet import th
|
||||
|
||||
# based on <https://github.com/ljleb/sd-webui-freeu/blob/main/lib_free_u/unet.py>
|
||||
# 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:
|
||||
|
||||
+14
-10
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
+4
-5
@@ -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
|
||||
|
||||
+4
-3
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user