diff --git a/CHANGELOG.md b/CHANGELOG.md index 3da5402a4..a8bc842a4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,8 +1,8 @@ # Change Log for SD.Next -## Update for 2024-10-11 +## Update for 2024-10-12 -### Highlights for 2024-10-11 +### Highlights for 2024-10-12 - **Reprocess**: New workflow options that allow you to generate at lower quality and then reprocess at higher quality for select images only or generate without hires/refine and then reprocess with hires/refine @@ -19,7 +19,7 @@ And other goodies like multiple *XYZ grid* improvements, additional *Flux ControlNets*, additional *Interrogate models*, better *LoRA tags* support, and more... -### Details for 2024-10-11 +### Details for 2024-10-12 - **reprocess** - new top-level button: reprocess latent from your history of generated image(s) diff --git a/extensions-builtin/Lora/lora_patches.py b/extensions-builtin/Lora/lora_patches.py index 575ab0348..b060e433d 100644 --- a/extensions-builtin/Lora/lora_patches.py +++ b/extensions-builtin/Lora/lora_patches.py @@ -1,6 +1,6 @@ import torch import networks -from modules import patches, shared +from modules import patches, shared, model_quant class LoraPatches: @@ -24,17 +24,13 @@ class LoraPatches: def apply(self): if self.active or shared.opts.lora_force_diffusers: return - try: - import bitsandbytes - self.Linear4bit_forward = patches.patch(__name__, bitsandbytes.nn.Linear4bit, 'forward', networks.network_Linear4bit_forward) - except Exception: - pass - try: - from optimum import quanto # pylint: disable=no-name-in-module + bnb = model_quant.load_bnb(silent=True) + if bnb is not None: + self.Linear4bit_forward = patches.patch(__name__, bnb.nn.Linear4bit, 'forward', networks.network_Linear4bit_forward) + quanto = model_quant.load_quanto(silent=True) + if quanto is not None: self.QLinear_forward = patches.patch(__name__, quanto.nn.QLinear, 'forward', networks.network_QLinear_forward) self.QConv2d_forward = patches.patch(__name__, quanto.nn.QConv2d, 'forward', networks.network_QConv2d_forward) - except Exception: - pass self.Linear_forward = patches.patch(__name__, torch.nn.Linear, 'forward', networks.network_Linear_forward) self.Linear_load_state_dict = patches.patch(__name__, torch.nn.Linear, '_load_from_state_dict', networks.network_Linear_load_state_dict) self.Conv2d_forward = patches.patch(__name__, torch.nn.Conv2d, 'forward', networks.network_Conv2d_forward) @@ -53,17 +49,13 @@ class LoraPatches: def undo(self): if not self.active or shared.opts.lora_force_diffusers: return - try: - import bitsandbytes - self.Linear4bit_forward = patches.undo(__name__, bitsandbytes.nn.Linear4bit, 'forward') # pylint: disable=E1128 - except Exception: - pass - try: - from optimum import quanto # pylint: disable=no-name-in-module + bnb = model_quant.load_bnb(silent=True) + if bnb is not None: + self.Linear4bit_forward = patches.undo(__name__, bnb.nn.Linear4bit, 'forward') # pylint: disable=E1128 + quanto = model_quant.load_quanto(silent=True) + if quanto is not None: self.QLinear_forward = patches.undo(__name__, quanto.nn.QLinear, 'forward') # pylint: disable=E1128 self.QConv2d_forward = patches.undo(__name__, quanto.nn.QConv2d, 'forward') # pylint: disable=E1128 - except Exception: - pass self.Linear_forward = patches.undo(__name__, torch.nn.Linear, 'forward') # pylint: disable=E1128 self.Linear_load_state_dict = patches.undo(__name__, torch.nn.Linear, '_load_from_state_dict') # pylint: disable=E1128 self.Conv2d_forward = patches.undo(__name__, torch.nn.Conv2d, 'forward') # pylint: disable=E1128 diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index b41041f95..0373080cd 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -17,7 +17,7 @@ import network_overrides import lora_convert import torch import diffusers.models.lora -from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, files_cache +from modules import shared, devices, sd_models, sd_models_compile, errors, scripts, files_cache, model_quant debug = os.environ.get('SD_LORA_DEBUG', None) is not None @@ -299,11 +299,13 @@ def network_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Li self.weight = torch.nn.Parameter(weights_backup.to(self.weight.device, copy=True)) self.freeze() elif getattr(self, "quant_type", None) in ['nf4', 'fp4']: - import bitsandbytes - device = self.weight.device - self.weight = bitsandbytes.nn.Params4bit(weights_backup, quant_state=self.quant_state, - quant_type=self.quant_type, blocksize=self.blocksize) - self.weight.to(device) + bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + if bnb is not None: + device = self.weight.device + self.weight = bnb.nn.Params4bit(weights_backup, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) + self.weight.to(device) + else: + self.weight.copy_(weights_backup) else: self.weight.copy_(weights_backup) if bias_backup is not None: @@ -339,16 +341,15 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn if isinstance(self, torch.nn.MultiheadAttention): weights_backup = (self.in_proj_weight.clone().to(devices.cpu), self.out_proj.weight.clone().to(devices.cpu)) elif getattr(self.weight, "quant_type", None) in ['nf4', 'fp4']: - import bitsandbytes - with devices.inference_context(): - weights_backup = bitsandbytes.functional.dequantize_4bit(self.weight, - quant_state=self.weight.quant_state, - quant_type=self.weight.quant_type, - blocksize=self.weight.blocksize, - ).to(devices.cpu) - self.quant_state = self.weight.quant_state - self.quant_type = self.weight.quant_type - self.blocksize = self.weight.blocksize + bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + if bnb is not None: + with devices.inference_context(): + weights_backup = bnb.functional.dequantize_4bit(self.weight, quant_state=self.weight.quant_state, quant_type=self.weight.quant_type, blocksize=self.weight.blocksize,).to(devices.cpu) + self.quant_state = self.weight.quant_state + self.quant_type = self.weight.quant_type + self.blocksize = self.weight.blocksize + else: + weights_backup = self.weight.clone().to(devices.cpu) else: weights_backup = self.weight.clone().to(devices.cpu) self.network_weights_backup = weights_backup @@ -376,16 +377,14 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn # inpainting model. zero pad updown to make channel[1] 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if getattr(self.weight, "quant_type", None) in ['nf4', 'fp4'] or self.weight.numel() != updown.numel(): - import bitsandbytes - device = self.weight.device - weight = bitsandbytes.functional.dequantize_4bit(self.weight, - quant_state=self.weight.quant_state, - quant_type=self.weight.quant_type, - blocksize=self.weight.blocksize) - self.weight = bitsandbytes.nn.Params4bit(weight + updown, quant_state=self.quant_state, - quant_type=shared.opts.lora_quant.lower(), - blocksize=self.blocksize) - self.weight.to(device) + bnb = model_quant.load_bnb('Load network: type=LoRA', silent=True) + if bnb is not None: + device = self.weight.device + weight = bnb.functional.dequantize_4bit(self.weight, quant_state=self.weight.quant_state, quant_type=self.weight.quant_type, blocksize=self.weight.blocksize) + self.weight = bnb.nn.Params4bit(weight + updown, quant_state=self.quant_state, quant_type=shared.opts.lora_quant.lower(), blocksize=self.blocksize) + self.weight.to(device) + else: + self.weight = torch.nn.Parameter(weight + updown) else: self.weight = torch.nn.Parameter(weight + updown) if hasattr(self, "qweight") and hasattr(self, "freeze"): diff --git a/installer.py b/installer.py index 4c51a0bfa..3e1e8843b 100644 --- a/installer.py +++ b/installer.py @@ -828,11 +828,6 @@ def install_packages(): # tensorflow_package = os.environ.get('TENSORFLOW_PACKAGE', None) # if tensorflow_package is not None: # install(tensorflow_package, 'tensorflow-rocm' if 'rocm' in tensorflow_package else 'tensorflow', ignore=True, quiet=True) - # bitsandbytes_package = os.environ.get('BITSANDBYTES_PACKAGE', None) - # if bitsandbytes_package is not None: - # install(bitsandbytes_package, 'bitsandbytes', ignore=True, quiet=True) - # elif not args.experimental: - # uninstall('bitsandbytes') if args.profile: pr.disable( ) print_profile(pr, 'Packages') diff --git a/modules/control/units/controlnet.py b/modules/control/units/controlnet.py index c1bf9aff8..d8ed7fe57 100644 --- a/modules/control/units/controlnet.py +++ b/modules/control/units/controlnet.py @@ -4,7 +4,7 @@ from typing import Union from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, FluxPipeline, ControlNetModel from modules.control.units import detect from modules.shared import log, opts, listdir -from modules import errors, sd_models, devices +from modules import errors, sd_models, devices, model_quant what = 'ControlNet' @@ -229,8 +229,7 @@ class ControlNet(): elif "ControlNet" in opts.optimum_quanto_weights: try: log.debug(f'Control {what} model Optimum Quanto: id="{model_id}"') - from installer import install - install('optimum-quanto', quiet=True) + model_quant.load_quanto('Load model: type=ControlNet') from modules.sd_models_compile import optimum_quanto_model self.model = optimum_quanto_model(self.model) except Exception as e: diff --git a/modules/devices.py b/modules/devices.py index cd20de418..253bd10a7 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -225,7 +225,8 @@ def torch_gc(force=False, fast=False): after = { 'gpu': mem.get('gpu', {}).get('used', 0), 'ram': mem.get('ram', {}).get('used', 0), 'retries': mem.get('retries', 0), 'oom': mem.get('oom', 0) } utilization = { 'gpu': used_gpu, 'ram': used_ram, 'threshold': threshold } results = { 'collected': collected, 'saved': saved } - log.debug(f'GC: utilization={utilization} gc={results} before={before} after={after} device={torch.device(get_optimal_device_name())} fn={sys._getframe(1).f_code.co_name} time={round(t1 - t0, 2)}') # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + log.debug(f'GC: utilization={utilization} gc={results} before={before} after={after} device={torch.device(get_optimal_device_name())} fn={fn} time={round(t1 - t0, 2)}') # pylint: disable=protected-access def set_cuda_sync_mode(mode): diff --git a/modules/generation_parameters_copypaste.py b/modules/generation_parameters_copypaste.py index feebfe50e..eacab607d 100644 --- a/modules/generation_parameters_copypaste.py +++ b/modules/generation_parameters_copypaste.py @@ -238,9 +238,10 @@ def connect_paste(button, local_paste_fields, input_comp, override_settings_comp v = params.get(param_name, None) if v is None: continue - if shared.opts.disable_weights_auto_swap: - if setting_name == "sd_model_checkpoint" or setting_name == 'sd_model_refiner' or setting_name == 'sd_backend' or setting_name == 'sd_vae': - continue + if setting_name == 'sd_backend': + continue + if shared.opts.disable_weights_auto_swap and setting_name in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_model_dict', 'sd_vae', 'sd_unet', 'sd_text_encoder']: + continue v = shared.opts.cast_value(setting_name, v) current_value = getattr(shared.opts, setting_name, None) if v == current_value: diff --git a/modules/images.py b/modules/images.py index 6808bb176..add7c5c7d 100644 --- a/modules/images.py +++ b/modules/images.py @@ -118,7 +118,8 @@ def save_image(image, suffix='', save_to_dirs=None, ): # pylint: disable=unused-argument - debug(f'Save: fn={sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + debug(f'Save: fn={fn}') # pylint: disable=protected-access if image is None: shared.log.warning('Image is none') return None, None, None diff --git a/modules/images_resize.py b/modules/images_resize.py index cf2a13385..d86ff6f22 100644 --- a/modules/images_resize.py +++ b/modules/images_resize.py @@ -128,5 +128,6 @@ def resize_image(resize_mode, im, width, height, upscaler_name=None, output_type res = im.copy() shared.log.error(f'Invalid resize mode: {resize_mode}') t1 = time.time() - shared.log.debug(f'Image resize: input={im} width={width} height={height} mode="{shared.resize_modes[resize_mode]}" upscaler="{upscaler_name}" context="{context}" type={output_type} result={res} time={t1-t0:.2f} fn={sys._getframe(1).f_code.co_filename}:{sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + shared.log.debug(f'Image resize: input={im} width={width} height={height} mode="{shared.resize_modes[resize_mode]}" upscaler="{upscaler_name}" context="{context}" type={output_type} result={res} time={t1-t0:.2f} fn={fn}') # pylint: disable=protected-access return np.array(res) if output_type == 'np' else res diff --git a/modules/infotext.py b/modules/infotext.py index 485bcf523..05e06e600 100644 --- a/modules/infotext.py +++ b/modules/infotext.py @@ -79,10 +79,15 @@ def parse(infotext): mapping = [ + # Backend ('Backend', 'sd_backend'), + # Models ('Model hash', 'sd_model_checkpoint'), ('Refiner', 'sd_model_refiner'), ('VAE', 'sd_vae'), + ('TE', 'sd_text_encoder'), + ('Unet', 'sd_unet'), + # Other ('Parser', 'prompt_attention'), ('Color correction', 'img2img_color_correction'), # Samplers diff --git a/modules/masking.py b/modules/masking.py index dc071eead..42e8e894e 100644 --- a/modules/masking.py +++ b/modules/masking.py @@ -375,7 +375,8 @@ def outpaint(input_image: Image.Image, outpaint_type: str = 'Edge'): def run_mask(input_image: Image.Image, input_mask: Image.Image = None, return_type: str = None, mask_blur: int = None, mask_padding: int = None, segment_enable=True, invert=None): - debug(f'Run mask: fn={sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + debug(f'Run mask: fn={fn}') # pylint: disable=protected-access if input_image is None: return input_mask diff --git a/modules/model_flux.py b/modules/model_flux.py index 95c89e451..2a80e5c34 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -5,38 +5,15 @@ import diffusers import transformers from safetensors.torch import load_file from huggingface_hub import hf_hub_download -from modules import shared, devices, modelloader, sd_models, sd_unet, model_te +from modules import shared, devices, modelloader, sd_models, sd_unet, model_te, model_quant debug = shared.log.trace if os.environ.get('SD_LOAD_DEBUG', None) is not None else lambda *args, **kwargs: None -def get_quant(file_path): - if "qint8" in file_path.lower(): - return 'qint8' - if "qint4" in file_path.lower(): - return 'qint4' - if "fp8" in file_path.lower(): - return 'fp8' - if "fp4" in file_path.lower(): - return 'fp4' - if "nf4" in file_path.lower(): - return 'nf4' - if file_path.endswith('.gguf'): - return 'gguf' - return 'none' - - def load_flux_quanto(checkpoint_info): transformer, text_encoder_2 = None, None - from installer import install - install('optimum-quanto', quiet=True) - try: - from optimum import quanto # pylint: disable=no-name-in-module - from optimum.quanto import requantize # pylint: disable=no-name-in-module - except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to import optimum-quanto: {e}") - raise + quanto = model_quant.load_quanto('Load model: type=FLUX') quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs) if isinstance(checkpoint_info, str): @@ -56,7 +33,7 @@ def load_flux_quanto(checkpoint_info): dtype = state_dict['context_embedder.bias'].dtype with torch.device("meta"): transformer = diffusers.FluxTransformer2DModel.from_config(os.path.join(repo_path, "transformer", "config.json")).to(dtype=dtype) - requantize(transformer, state_dict, quantization_map, device=torch.device("cpu")) + quanto.requantize(transformer, state_dict, quantization_map, device=torch.device("cpu")) transformer.eval() if transformer.dtype != devices.dtype: try: @@ -83,7 +60,7 @@ def load_flux_quanto(checkpoint_info): dtype = state_dict['encoder.block.0.layer.0.SelfAttention.relative_attention_bias.weight'].dtype with torch.device("meta"): text_encoder_2 = transformers.T5EncoderModel(t5_config).to(dtype=dtype) - requantize(text_encoder_2, state_dict, quantization_map, device=torch.device("cpu")) + quanto.requantize(text_encoder_2, state_dict, quantization_map, device=torch.device("cpu")) text_encoder_2.eval() if text_encoder_2.dtype != devices.dtype: try: @@ -105,9 +82,8 @@ def load_flux_bnb(checkpoint_info, diffusers_load_config): # pylint: disable=unu repo_path = checkpoint_info else: repo_path = checkpoint_info.path - from installer import install - install('bitsandbytes', quiet=True) - quant = get_quant(repo_path) + model_quant.load_bnb('Load model: type=T5') + quant = model_quant.get_quant(repo_path) try: if quant == 'fp8': quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True, bnb_4bit_compute_dtype=devices.dtype) @@ -162,7 +138,7 @@ def load_transformer(file_path): # triggered by opts.sd_unet change if file_path is None or not os.path.exists(file_path): return None transformer = None - quant = get_quant(file_path) + quant = model_quant.get_quant(file_path) diffusers_load_config = { "low_cpu_mem_usage": True, "torch_dtype": devices.dtype, @@ -195,7 +171,7 @@ def load_transformer(file_path): # triggered by opts.sd_unet change def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_checkpoint change - quant = get_quant(checkpoint_info.path) + quant = model_quant.get_quant(checkpoint_info.path) repo_id = sd_models.path_to_repo(checkpoint_info.name) shared.log.debug(f'Load model: type=FLUX model="{checkpoint_info.name}" repo="{repo_id}" unet="{shared.opts.sd_unet}" t5="{shared.opts.sd_text_encoder}" vae="{shared.opts.sd_vae}" quant={quant} offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype}') debug(f'Load model: type=FLUX config={diffusers_load_config}') diff --git a/modules/model_flux_nf4.py b/modules/model_flux_nf4.py index 244ca29b9..a1b46fd54 100644 --- a/modules/model_flux_nf4.py +++ b/modules/model_flux_nf4.py @@ -11,25 +11,12 @@ from accelerate import init_empty_weights from accelerate.utils import set_module_tensor_to_device from diffusers.loaders.single_file_utils import convert_flux_transformer_checkpoint_to_diffusers import safetensors.torch -from modules import shared, devices +from modules import shared, devices, model_quant -bnb = None debug = os.environ.get('SD_LOAD_DEBUG', None) is not None -def load_bnb(): - from installer import install - install('bitsandbytes', quiet=True) - try: - import bitsandbytes - global bnb # pylint: disable=global-statement - bnb = bitsandbytes - except Exception as e: - shared.log.error(f"Load model: type=FLUX Failed to import bitsandbytes: {e}") - raise - - def _replace_with_bnb_linear( model, method="nf4", @@ -40,6 +27,7 @@ def _replace_with_bnb_linear( Returns the converted model and a boolean that indicates if the conversion has been successfull or not. """ + bnb = model_quant.load_bnb('Load model: type=FLUX') for name, module in model.named_children(): if isinstance(module, nn.Linear): with init_empty_weights(): @@ -83,6 +71,7 @@ def check_quantized_param( model, param_name: str, ) -> bool: + bnb = model_quant.load_bnb('Load model: type=FLUX') module, tensor_name = get_module_from_name(model, param_name) if isinstance(module._parameters.get(tensor_name, None), bnb.nn.Params4bit): # pylint: disable=protected-access # Add here check for loaded components' dtypes once serialization is implemented @@ -104,6 +93,7 @@ def create_quantized_param( unexpected_keys=None, pre_quantized=False ): + bnb = model_quant.load_bnb('Load model: type=FLUX') module, tensor_name = get_module_from_name(model, param_name) if tensor_name not in module._parameters: # pylint: disable=protected-access @@ -163,7 +153,6 @@ def create_quantized_param( def load_flux_nf4(checkpoint_info): - load_bnb() transformer = None text_encoder_2 = None if isinstance(checkpoint_info, str): diff --git a/modules/model_quant.py b/modules/model_quant.py new file mode 100644 index 000000000..24d673e95 --- /dev/null +++ b/modules/model_quant.py @@ -0,0 +1,60 @@ +import sys +from installer import install, log + + +bnb = None +quanto = None + + +def load_bnb(msg='', silent=False): + global bnb # pylint: disable=global-statement + if bnb is not None: + return bnb + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + log.debug(f'Quantization: module=bitsandbytes fn={fn}') # pylint: disable=protected-access + install('bitsandbytes', quiet=True) + try: + import bitsandbytes + bnb = bitsandbytes + return bnb + except Exception as e: + if len(msg) > 0: + log.error(f"{msg} failed to import bitsandbytes: {e}") + bnb = None + if not silent: + raise + + +def load_quanto(msg='', silent=False): + global quanto # pylint: disable=global-statement + if quanto is not None: + return quanto + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + log.debug(f'Quantization: module=quanto fn={fn}') # pylint: disable=protected-access + install('optimum-quanto', quiet=True) + try: + from optimum import quanto as optimum_quanto + quanto = optimum_quanto + return quanto + except Exception as e: + if len(msg) > 0: + log.error(f"{msg} failed to import optimum.quanto: {e}") + quanto = None + if not silent: + raise + + +def get_quant(name): + if "qint8" in name.lower(): + return 'qint8' + if "qint4" in name.lower(): + return 'qint4' + if "fp8" in name.lower(): + return 'fp8' + if "fp4" in name.lower(): + return 'fp4' + if "nf4" in name.lower(): + return 'nf4' + if name.endswith('.gguf'): + return 'gguf' + return 'none' diff --git a/modules/model_te.py b/modules/model_te.py index 27fa131b4..40d544e1d 100644 --- a/modules/model_te.py +++ b/modules/model_te.py @@ -3,7 +3,7 @@ import json import torch import transformers from safetensors.torch import load_file -from modules import shared, devices, files_cache, errors +from modules import shared, devices, files_cache, errors, model_quant from installer import install @@ -67,15 +67,15 @@ def load_t5(name=None, cache_dir=None): elif 'fp16' in name.lower(): t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype) elif 'fp4' in name.lower(): - install('bitsandbytes', quiet=True) + model_quant.load_bnb('Load model: type=T5') quantization_config = transformers.BitsAndBytesConfig(load_in_4bit=True) t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', quantization_config=quantization_config, cache_dir=cache_dir, torch_dtype=devices.dtype) elif 'fp8' in name.lower(): - install('bitsandbytes', quiet=True) + model_quant.load_bnb('Load model: type=T5') quantization_config = transformers.BitsAndBytesConfig(load_in_8bit=True) t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', quantization_config=quantization_config, cache_dir=cache_dir, torch_dtype=devices.dtype) elif 'qint8' in name.lower(): - install('optimum-quanto', quiet=True) + model_quant.load_quanto('Load model: type=T5') from modules.sd_models_compile import optimum_quanto_model t5 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder='text_encoder_3', cache_dir=cache_dir, torch_dtype=devices.dtype) t5 = optimum_quanto_model(t5, weights="qint8", activations="none") diff --git a/modules/modeldata.py b/modules/modeldata.py index 711d1f71b..27444aeaf 100644 --- a/modules/modeldata.py +++ b/modules/modeldata.py @@ -56,7 +56,8 @@ class Shared(sys.modules[__name__].__class__): def sd_model(self): import modules.sd_models # pylint: disable=W0621 if modules.sd_models.model_data.sd_model is None: - shared.log.debug(f'Model requested: fn={sys._getframe(1).f_code.co_filename}:{sys._getframe(1).f_code.co_name}/{sys._getframe(2).f_code.co_filename}:{sys._getframe(2).f_code.co_name}') # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + shared.log.debug(f'Model requested: fn={fn}') # pylint: disable=protected-access return modules.sd_models.model_data.get_sd_model() @sd_model.setter diff --git a/modules/processing_class.py b/modules/processing_class.py index ccb426605..9265ea3cf 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -620,7 +620,8 @@ def switch_class(p: StableDiffusionProcessing, new_class: type, dct: dict = None kwargs[k] = v if new_class == StableDiffusionProcessingTxt2Img: sd_models.clean_diffuser_pipe(shared.sd_model) - debug(f"Switching class: {p.__class__.__name__} -> {new_class.__name__} fn={sys._getframe(1).f_code.co_name}") # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + debug(f"Switching class: {p.__class__.__name__} -> {new_class.__name__} fn={fn}") # pylint: disable=protected-access p.__class__ = new_class p.__init__(**kwargs) for k, v in p.__dict__.items(): diff --git a/modules/processing_info.py b/modules/processing_info.py index b61adce8a..29513167d 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -69,13 +69,14 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No args['Grid'] = grid if shared.native: args['Pipeline'] = shared.sd_model.__class__.__name__ - args['T5'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_text_encoder is None or shared.opts.sd_text_encoder == 'None') else shared.opts.sd_text_encoder + args['TE'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_text_encoder is None or shared.opts.sd_text_encoder == 'None') else shared.opts.sd_text_encoder + args['UNet'] = None if (not shared.opts.add_model_name_to_info or shared.opts.sd_unet is None or shared.opts.sd_unet == 'None') else shared.opts.sd_unet if 'txt2img' in p.ops: args["Variation seed"] = all_subseeds[index] if p.subseed_strength > 0 else None args["Variation strength"] = p.subseed_strength if p.subseed_strength > 0 else None if 'hires' in p.ops or 'upscale' in p.ops: is_resize = p.hr_resize_mode > 0 and (p.hr_upscaler != 'None' or p.hr_resize_mode == 5) - args["Second pass"] = p.enable_hr + args["Refine"] = p.enable_hr args["Hires force"] = p.hr_force args["Hires steps"] = p.hr_second_pass_steps args["HiRes resize mode"] = p.hr_resize_mode if is_resize else None @@ -89,7 +90,7 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No args["Image CFG scale"] = p.image_cfg_scale args["CFG rescale"] = p.diffusers_guidance_rescale if 'refine' in p.ops: - args["Second pass"] = p.enable_hr + args["Refine"] = p.enable_hr args["Refiner"] = None if (not shared.opts.add_model_name_to_info) or (not shared.sd_refiner) or (not shared.sd_refiner.sd_checkpoint_info.model_name) else shared.sd_refiner.sd_checkpoint_info.model_name.replace(',', '').replace(':', '') args['Image CFG scale'] = p.image_cfg_scale args['Refiner steps'] = p.refiner_steps diff --git a/modules/sd_models.py b/modules/sd_models.py index 09fa1dc10..538e8b047 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -877,13 +877,14 @@ def move_model(model, device=None, force=False): devices.torch_gc() return + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access if getattr(model, 'vae', None) is not None and get_diffusers_task(model) != DiffusersTaskType.TEXT_2_IMAGE: if device == devices.device and model.vae.device.type != "meta": # force vae back to gpu if not in txt2img mode model.vae.to(device) if hasattr(model.vae, '_hf_hook'): - debug_move(f'Model move: to={device} class={model.vae.__class__} fn={sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access + debug_move(f'Model move: to={device} class={model.vae.__class__} fn={fn}') # pylint: disable=protected-access model.vae._hf_hook.execution_device = device # pylint: disable=protected-access - debug_move(f'Model move: device={device} class={model.__class__} accelerate={getattr(model, "has_accelerate", False)} fn={sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access + debug_move(f'Model move: device={device} class={model.__class__} accelerate={getattr(model, "has_accelerate", False)} fn={fn}') # pylint: disable=protected-access if hasattr(model, "components"): # accelerate patch for name, m in model.components.items(): if not hasattr(m, "_hf_hook"): # not accelerate hook @@ -1553,7 +1554,8 @@ def set_diffuser_pipe(pipe, new_pipe_type): new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True) if hasattr(new_pipe, "watermark"): new_pipe.watermark = NoWatermark() - shared.log.debug(f"Pipeline class change: original={pipe.__class__.__name__} target={new_pipe.__class__.__name__} device={pipe.device} fn={sys._getframe().f_back.f_code.co_name}") # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + shared.log.debug(f"Pipeline class change: original={pipe.__class__.__name__} target={new_pipe.__class__.__name__} device={pipe.device} fn={fn}") # pylint: disable=protected-access pipe = new_pipe return pipe diff --git a/modules/sd_models_compile.py b/modules/sd_models_compile.py index 9b3eb90af..1ef333d5a 100644 --- a/modules/sd_models_compile.py +++ b/modules/sd_models_compile.py @@ -2,7 +2,7 @@ import copy import time import logging import torch -from modules import shared, devices, sd_models +from modules import shared, devices, sd_models, model_quant from installer import install, setup_logging @@ -141,6 +141,7 @@ def ipex_optimize(sd_model): shared.log.warning(f"IPEX Optimize: error: {e}") return sd_model + def nncf_send_to_device(model): for child in model.children(): if child.__class__.__name__ == "WeightsDecompressor": @@ -148,6 +149,7 @@ def nncf_send_to_device(model): child.zero_point = child.zero_point.to(devices.device) nncf_send_to_device(child) + def nncf_compress_model(model, op=None, sd_model=None): import nncf global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement @@ -177,6 +179,7 @@ def nncf_compress_model(model, op=None, sd_model=None): devices.torch_gc(force=True) return model + def nncf_compress_weights(sd_model): try: t0 = time.time() @@ -201,8 +204,9 @@ def nncf_compress_weights(sd_model): shared.log.warning(f"NNCF Compress Weights: error: {e}") return sd_model + def optimum_quanto_model(model, op=None, sd_model=None, weights=None, activations=None): - from optimum import quanto # pylint: disable=no-name-in-module + quanto = model_quant.load_quanto('Compile model: type=Optimum Quanto') global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement if sd_model is not None and "Flux" in sd_model.__class__.__name__: # LayerNorm is not supported exclude_list = ["transformer_blocks.*.norm1.norm", "transformer_blocks.*.norm2", "transformer_blocks.*.norm1_context.norm", "transformer_blocks.*.norm2_context", "single_transformer_blocks.*.norm.norm", "norm_out.norm"] @@ -241,6 +245,7 @@ def optimum_quanto_model(model, op=None, sd_model=None, weights=None, activation devices.torch_gc(force=True) return model + def optimum_quanto_weights(sd_model): try: if shared.opts.diffusers_offload_mode in {"balanced", "sequential"}: @@ -249,8 +254,7 @@ def optimum_quanto_weights(sd_model): t0 = time.time() shared.log.info(f"Optimum Quanto Weights: {shared.opts.optimum_quanto_weights}") global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement - install('optimum-quanto', quiet=True) - from optimum import quanto # pylint: disable=no-name-in-module + quanto = model_quant.load_quanto() quanto.tensor.qbits.QBitsTensor.create = lambda *args, **kwargs: quanto.tensor.qbits.QBitsTensor(*args, **kwargs) sd_model = apply_compile_to_model(sd_model, optimum_quanto_model, shared.opts.optimum_quanto_weights, op="optimum-quanto") @@ -300,6 +304,7 @@ def optimum_quanto_weights(sd_model): shared.log.warning(f"Optimum Quanto Weights: error: {e}") return sd_model + def optimize_openvino(sd_model): try: from modules.intel.openvino import openvino_fx # pylint: disable=unused-import diff --git a/modules/shared_state.py b/modules/shared_state.py index 82f8d21dc..067fb21eb 100644 --- a/modules/shared_state.py +++ b/modules/shared_state.py @@ -88,7 +88,8 @@ class State: def end(self, api=None): import modules.devices if self.time_start is None: # someone called end before being - log.debug(f'Access state.end: {sys._getframe().f_back.f_code.co_name}') # pylint: disable=protected-access + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + log.debug(f'Access state.end: {fn}') # pylint: disable=protected-access self.time_start = time.time() if self.debug_output: log.debug(f'State end: {self.job} time={time.time() - self.time_start:.2f}') diff --git a/modules/ui.py b/modules/ui.py index 41ef6c6c1..8ea565cfc 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -85,8 +85,9 @@ def apply_setting(key, value): return gr.update() if shared.cmd_opts.freeze: return gr.update() - # dont allow model to be swapped when model hash exists in prompt - if key == "sd_model_checkpoint" and opts.disable_weights_auto_swap: + if key == 'sd_backend': + return gr.update() + if shared.opts.disable_weights_auto_swap and key in ['sd_model_checkpoint', 'sd_model_refiner', 'sd_model_dict', 'sd_vae', 'sd_unet', 'sd_text_encoder']: return gr.update() if key == "sd_model_checkpoint": ckpt_info = sd_models.get_closet_checkpoint_match(value) diff --git a/modules/ui_control.py b/modules/ui_control.py index d7cfcdf0a..b7c5ead96 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -621,6 +621,7 @@ def create_ui(_blocks: gr.Blocks=None): (hidiffusion, "HiDiffusion"), # second pass (enable_hr, "Second pass"), + (enable_hr, "Refine"), (hr_sampler_index, "Hires sampler"), (denoising_strength, "Denoising strength"), (hr_upscaler, "Hires upscaler"), diff --git a/modules/ui_txt2img.py b/modules/ui_txt2img.py index f61196997..1ae3a8dad 100644 --- a/modules/ui_txt2img.py +++ b/modules/ui_txt2img.py @@ -125,6 +125,7 @@ def create_ui(): (hidiffusion, "HiDiffusion"), # second pass (enable_hr, "Second pass"), + (enable_hr, "Refine"), (denoising_strength, "Denoising strength"), (hr_sampler_index, "Hires sampler"), (hr_resize_mode, "Hires resize mode"), diff --git a/webui.py b/webui.py index cf244fc44..f0f29e13a 100644 --- a/webui.py +++ b/webui.py @@ -169,10 +169,10 @@ def load_model(): timer.startup.record("checkpoint") shared.opts.onchange("sd_model_checkpoint", wrap_queued_call(lambda: modules.sd_models.reload_model_weights(op='model')), call=False) shared.opts.onchange("sd_model_refiner", wrap_queued_call(lambda: modules.sd_models.reload_model_weights(op='refiner')), call=False) - shared.opts.onchange("sd_text_encoder", wrap_queued_call(lambda: modules.sd_models.reload_text_encoder()), call=False) shared.opts.onchange("sd_model_dict", wrap_queued_call(lambda: modules.sd_models.reload_model_weights(op='dict')), call=False) shared.opts.onchange("sd_vae", wrap_queued_call(lambda: modules.sd_vae.reload_vae_weights()), call=False) shared.opts.onchange("sd_unet", wrap_queued_call(lambda: modules.sd_unet.load_unet(shared.sd_model)), call=False) + shared.opts.onchange("sd_text_encoder", wrap_queued_call(lambda: modules.sd_models.reload_text_encoder()), call=False) shared.opts.onchange("sd_backend", wrap_queued_call(lambda: modules.sd_models.change_backend()), call=False) shared.opts.onchange("temp_dir", gr_tempdir.on_tmpdir_changed) timer.startup.record("onchange")