mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
parse model preload
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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}')
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user