parse model preload

This commit is contained in:
Vladimir Mandic
2023-04-20 23:19:25 -04:00
parent df424d6d51
commit 7939a1649d
9 changed files with 37 additions and 38 deletions
+2 -1
View File
@@ -4,7 +4,7 @@
Stuff to be fixed...
- Fix integration with <https://github.com/deforum-art/deforum-for-automatic1111-webui>
- 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
+9 -3
View File
@@ -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
+1 -1
View File
@@ -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)
-2
View File
@@ -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}')
+2 -2
View File
@@ -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()
+13 -18
View File
@@ -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
+1 -4
View File
@@ -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:
+6 -7
View File
@@ -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,
+3
View File
@@ -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