From 7939a1649d26396e128d210e8f711a2927f293e7 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 20 Apr 2023 23:19:25 -0400 Subject: [PATCH] parse model preload --- TODO.md | 3 ++- modules/devices.py | 12 +++++++++--- modules/processing.py | 2 +- modules/script_loading.py | 2 -- modules/sd_models.py | 4 ++-- modules/sd_vae.py | 31 +++++++++++++------------------ modules/shared.py | 5 +---- modules/xlmr.py | 13 ++++++------- setup.py | 3 +++ 9 files changed, 37 insertions(+), 38 deletions(-) diff --git a/TODO.md b/TODO.md index 03dc58039..661f20753 100644 --- a/TODO.md +++ b/TODO.md @@ -4,7 +4,7 @@ Stuff to be fixed... -- Fix integration with +- Fix extensions imports - ClipSkip not updated on read gen info ## Features @@ -57,3 +57,4 @@ Tech that can be integrated as part of the core workflow... ### Pending Code Updates +- fix parse cmd line args from extensions diff --git a/modules/devices.py b/modules/devices.py index 88c64f690..7675c5bc3 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -72,11 +72,17 @@ def set_cuda_params(): torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = shared.opts.cuda_allow_tf16_reduced global dtype, dtype_vae, dtype_unet, unet_needs_upcast # pylint: disable=global-statement if shared.opts.cuda_dtype == 'FP16': - dtype = dtype_vae = dtype_unet = torch.float16 + dtype = torch.float16 + dtype_vae = torch.float16 + dtype_unet = torch.float16 if shared.opts.cuda_dtype == 'BP16': - dtype = dtype_vae = dtype_unet = torch.bfloat16 + dtype = torch.bfloat16 + dtype_vae = torch.bfloat16 + dtype_unet = torch.bfloat16 if shared.opts.cuda_dtype == 'FP32': - dtype = dtype_vae = dtype_unet = torch.float32 + dtype = torch.float32 + dtype_vae = torch.float32 + dtype_unet = torch.float32 unet_needs_upcast = shared.opts.upcast_sampling diff --git a/modules/processing.py b/modules/processing.py index 43417bcc9..c40681961 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -104,7 +104,7 @@ class StableDiffusionProcessing: """ The first set of paramaters: sd_models -> do_not_reload_embeddings represent the minimum required to create a StableDiffusionProcessing """ - def __init__(self, sd_model=None, outpath_samples=None, outpath_grids=None, prompt: str = "", styles: List[str] = None, seed: int = -1, subseed: int = -1, subseed_strength: float = 0, seed_resize_from_h: int = -1, seed_resize_from_w: int = -1, seed_enable_extras: bool = True, sampler_name: str = None, batch_size: int = 1, n_iter: int = 1, steps: int = 50, cfg_scale: float = 7.0, width: int = 512, height: int = 512, restore_faces: bool = False, tiling: bool = False, do_not_save_samples: bool = False, do_not_save_grid: bool = False, extra_generation_params: Dict[Any, Any] = None, overlay_images: Any = None, negative_prompt: str = None, eta: float = None, do_not_reload_embeddings: bool = False, denoising_strength: float = 0, ddim_discretize: str = None, s_churn: float = 0.0, s_tmax: float = None, s_tmin: float = 0.0, s_noise: float = 1.0, override_settings: Dict[str, Any] = None, override_settings_restore_afterwards: bool = True, sampler_index: int = None, script_args: list = None): + def __init__(self, sd_model=None, outpath_samples=None, outpath_grids=None, prompt: str = "", styles: List[str] = None, seed: int = -1, subseed: int = -1, subseed_strength: float = 0, seed_resize_from_h: int = -1, seed_resize_from_w: int = -1, seed_enable_extras: bool = True, sampler_name: str = None, batch_size: int = 1, n_iter: int = 1, steps: int = 50, cfg_scale: float = 7.0, width: int = 512, height: int = 512, restore_faces: bool = False, tiling: bool = False, do_not_save_samples: bool = False, do_not_save_grid: bool = False, extra_generation_params: Dict[Any, Any] = None, overlay_images: Any = None, negative_prompt: str = None, eta: float = None, do_not_reload_embeddings: bool = False, denoising_strength: float = 0, ddim_discretize: str = None, s_churn: float = 0.0, s_tmax: float = None, s_tmin: float = 0.0, s_noise: float = 1.0, override_settings: Dict[str, Any] = None, override_settings_restore_afterwards: bool = True, sampler_index: int = None, script_args: list = None): # pylint: disable=unused-argument if sampler_index is not None: print("sampler_index argument for StableDiffusionProcessing does not do anything; use sampler_name", file=sys.stderr) diff --git a/modules/script_loading.py b/modules/script_loading.py index d5b809492..abfc05473 100644 --- a/modules/script_loading.py +++ b/modules/script_loading.py @@ -21,11 +21,9 @@ def preload_extensions(extensions_dir, parser): preload_script = os.path.join(extensions_dir, dirname, "preload.py") if not os.path.isfile(preload_script): continue - try: module = load_module(preload_script) if hasattr(module, 'preload'): module.preload(parser) - except Exception as e: errors.display(e, f'Extension preload: {preload_script}') diff --git a/modules/sd_models.py b/modules/sd_models.py index 4f7600da8..f72c7b657 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -102,7 +102,7 @@ def checkpoint_tiles(): def list_models(): - global model_path + global model_path # pylint: disable=global-statement model_path = shared.opts.ckpt_dir checkpoints_list.clear() checkpoint_aliases.clear() @@ -302,6 +302,7 @@ def load_model_weights(model, checkpoint_info: CheckpointInfo, state_dict, timer if depth_model: model.depth_model = depth_model + devices.set_cuda_params() devices.dtype_unet = model.model.diffusion_model.dtype model.first_stage_model.to(devices.dtype_vae) @@ -402,7 +403,6 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None): sd_hijack.model_hijack.undo_hijack(shared.sd_model) shared.sd_model = None - devices.set_cuda_params() gc.collect() devices.torch_gc() diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 6b8652471..516b82b9e 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -4,17 +4,14 @@ import glob from copy import deepcopy from rich import print # pylint: disable=redefined-builtin from modules import paths, shared, devices, script_callbacks, sd_models -from modules import paths_internal + vae_ignore_keys = {"model_ema.decay", "model_ema.num_updates"} vae_dict = {} - - base_vae = None loaded_vae_file = None checkpoint_info = None vae_path = os.path.abspath(os.path.join(paths.models_path, 'VAE')) - checkpoints_loaded = collections.OrderedDict() def get_base_vae(model): @@ -24,7 +21,7 @@ def get_base_vae(model): def store_base_vae(model): - global base_vae, checkpoint_info + global base_vae, checkpoint_info # pylint: disable=global-statement if checkpoint_info != model.sd_checkpoint_info: assert not loaded_vae_file, "Trying to store non-base VAE!" base_vae = deepcopy(model.first_stage_model.state_dict()) @@ -32,13 +29,13 @@ def store_base_vae(model): def delete_base_vae(): - global base_vae, checkpoint_info + global base_vae, checkpoint_info # pylint: disable=global-statement base_vae = None checkpoint_info = None def restore_base_vae(model): - global loaded_vae_file + global loaded_vae_file # pylint: disable=global-statement if base_vae is not None and checkpoint_info == model.sd_checkpoint_info: print("Restoring base VAE") _load_vae_dict(model, base_vae) @@ -51,11 +48,11 @@ def get_filename(filepath): def refresh_vae_list(): - global vae_path + global vae_path # pylint: disable=global-statement vae_path = shared.opts.vae_dir vae_dict.clear() - paths = [ + vae_paths = [ os.path.join(sd_models.model_path, '**/*.vae.ckpt'), os.path.join(sd_models.model_path, '**/*.vae.pt'), os.path.join(sd_models.model_path, '**/*.vae.safetensors'), @@ -63,23 +60,20 @@ def refresh_vae_list(): os.path.join(shared.opts.vae_dir, '**/*.pt'), os.path.join(shared.opts.vae_dir, '**/*.safetensors'), ] - if shared.opts.ckpt_dir is not None and os.path.isdir(shared.opts.ckpt_dir): - paths += [ + vae_paths += [ os.path.join(shared.opts.ckpt_dir, '**/*.vae.ckpt'), os.path.join(shared.opts.ckpt_dir, '**/*.vae.pt'), os.path.join(shared.opts.ckpt_dir, '**/*.vae.safetensors'), ] - if shared.opts.vae_dir is not None and os.path.isdir(shared.opts.vae_dir): - paths += [ + vae_paths += [ os.path.join(shared.opts.vae_dir, '**/*.ckpt'), os.path.join(shared.opts.vae_dir, '**/*.pt'), os.path.join(shared.opts.vae_dir, '**/*.safetensors'), ] - candidates = [] - for path in paths: + for path in vae_paths: candidates += glob.iglob(path, recursive=True) for filepath in candidates: @@ -126,7 +120,7 @@ def load_vae_dict(filename): def load_vae(model, vae_file=None, vae_source="from unknown source"): - global vae_dict, loaded_vae_file + global loaded_vae_file # pylint: disable=global-statement # save_settings = False cache_enabled = shared.opts.sd_vae_checkpoint_cache > 0 @@ -172,7 +166,7 @@ def _load_vae_dict(model, vae_dict_1): def clear_loaded_vae(): - global loaded_vae_file + global loaded_vae_file # pylint: disable=global-statement loaded_vae_file = None @@ -180,11 +174,12 @@ unspecified = object() def reload_vae_weights(sd_model=None, vae_file=unspecified): - from modules import lowvram, devices, sd_hijack + from modules import lowvram, sd_hijack if not sd_model: sd_model = shared.sd_model + global checkpoint_info # pylint: disable=global-statement checkpoint_info = sd_model.sd_checkpoint_info checkpoint_file = checkpoint_info.filename diff --git a/modules/shared.py b/modules/shared.py index 47cd9bfc5..b3244e573 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -11,7 +11,7 @@ import modules.interrogate import modules.memmon import modules.styles import modules.devices as devices -from modules import script_loading, errors, ui_components, shared_items, cmd_args +from modules import errors, ui_components, shared_items, cmd_args from modules.paths_internal import models_path, script_path, data_path, sd_configs_path, sd_default_config, sd_model_file, default_sd_model_file, extensions_dir, extensions_builtin_dir # pylint: disable=W0611 import modules.paths_internal as paths from setup import log as setup_log # pylint: disable=E0611 @@ -21,9 +21,6 @@ demo: gr.Blocks = None log = setup_log parser = cmd_args.parser -script_loading.preload_extensions(paths.extensions_dir, parser) -script_loading.preload_extensions(paths.extensions_builtin_dir, parser) - if os.environ.get('IGNORE_CMD_ARGS_ERRORS', None) is None: cmd_opts = parser.parse_args() else: diff --git a/modules/xlmr.py b/modules/xlmr.py index beab3fdf5..9da3161cc 100644 --- a/modules/xlmr.py +++ b/modules/xlmr.py @@ -1,9 +1,8 @@ -from transformers import BertPreTrainedModel,BertModel,BertConfig -import torch.nn as nn -import torch -from transformers.models.xlm_roberta.configuration_xlm_roberta import XLMRobertaConfig -from transformers import XLMRobertaModel,XLMRobertaTokenizer from typing import Optional +import torch +import torch.nn as nn +from transformers import XLMRobertaModel,XLMRobertaTokenizer, BertPreTrainedModel, BertModel, BertConfig # pylint: disable=unused-import +from transformers.models.xlm_roberta.configuration_xlm_roberta import XLMRobertaConfig class BertSeriesConfig(BertConfig): def __init__(self, vocab_size=30522, hidden_size=768, num_hidden_layers=12, num_attention_heads=12, intermediate_size=3072, hidden_act="gelu", hidden_dropout_prob=0.1, attention_probs_dropout_prob=0.1, max_position_embeddings=512, type_vocab_size=2, initializer_range=0.02, layer_norm_eps=1e-12, pad_token_id=0, position_embedding_type="absolute", use_cache=True, classifier_dropout=None,project_dim=512, pooler_fn="average",learn_encoder=False,model_type='bert',**kwargs): @@ -28,7 +27,7 @@ class BertSeriesModelWithTransformation(BertPreTrainedModel): config_class = BertSeriesConfig def __init__(self, config=None, **kargs): - # modify initialization for autoloading + # modify initialization for autoloading if config is None: config = XLMRobertaConfig() config.attention_probs_dropout_prob= 0.1 @@ -74,7 +73,7 @@ class BertSeriesModelWithTransformation(BertPreTrainedModel): text["attention_mask"] = torch.tensor( text['attention_mask']).to(device) features = self(**text) - return features['projection_state'] + return features['projection_state'] def forward( self, diff --git a/setup.py b/setup.py index f6f223cae..63b5818c2 100644 --- a/setup.py +++ b/setup.py @@ -11,6 +11,9 @@ try: except: import argparse parser = argparse.ArgumentParser(description="Stable Diffusion", formatter_class=lambda prog: argparse.HelpFormatter(prog,max_help_position=55,indent_increment=2,width=200)) +from modules.script_loading import preload_extensions +preload_extensions('extensions', parser) +preload_extensions('extensions-builtin', parser) class Dot(dict): # dot notation access to dictionary attributes