diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index fdd6ec680..1e0715e01 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -243,6 +243,8 @@ def process_diffusers(p: processing.StableDiffusionProcessing): if shared.state.interrupted or shared.state.skipped: shared.sd_model = orig_pipeline return results + if shared.opts.diffusers_offload_mode == "balanced": + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) if shared.opts.diffusers_move_refiner: sd_models.move_model(shared.sd_refiner, devices.device) p.ops.append('refine') @@ -296,7 +298,9 @@ def process_diffusers(p: processing.StableDiffusionProcessing): for refiner_image in refiner_images: results.append(refiner_image) - if shared.opts.diffusers_move_refiner: + if shared.opts.diffusers_offload_mode == "balanced": + shared.sd_refiner = sd_models.apply_balanced_offload(shared.sd_refiner) + elif shared.opts.diffusers_move_refiner: shared.log.debug('Moving to CPU: model=refiner') sd_models.move_model(shared.sd_refiner, devices.cpu) shared.state.job = prev_job diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 80ce38acf..76f1939e6 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -35,7 +35,9 @@ def full_vae_decode(latents, model): t0 = time.time() if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False): base_device = sd_models.move_base(model, devices.cpu) - if not shared.opts.diffusers_offload_mode == "sequential" and hasattr(model, 'vae'): + if shared.opts.diffusers_offload_mode == "balanced": + shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model) + elif not shared.opts.diffusers_offload_mode == "sequential" and hasattr(model, 'vae'): sd_models.move_model(model.vae, devices.device) latents.to(model.vae.device) diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index ae4405a45..cf6ef497b 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -5,7 +5,7 @@ import typing import torch from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider from transformers import PreTrainedTokenizer -from modules import shared, prompt_parser, devices +from modules import shared, prompt_parser, devices, sd_models debug_enabled = os.environ.get('SD_PROMPT_DEBUG', None) @@ -173,7 +173,9 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c p.negative_embeds = [] p.negative_pooleds = [] - if hasattr(pipe, "maybe_free_model_hooks"): + if shared.opts.diffusers_offload_mode == "balanced": + pipe = sd_models.apply_balanced_offload(pipe) + elif hasattr(pipe, "maybe_free_model_hooks"): # if the last job is interrupted, model will stay in the vram and cause oom, send everything back to cpu before continuing pipe.maybe_free_model_hooks() devices.torch_gc() @@ -209,7 +211,9 @@ def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, c if debug_enabled: get_tokens('positive', prompts[0]) get_tokens('negative', negative_prompts[0]) - if hasattr(pipe, "maybe_free_model_hooks"): + if shared.opts.diffusers_offload_mode == "balanced": + pipe = sd_models.apply_balanced_offload(pipe) + elif hasattr(pipe, "maybe_free_model_hooks"): # text encoder will stay in the vram and cause oom, send everything back to cpu before continuing pipe.maybe_free_model_hooks() debug(f"Prompt encode: time={(time.time() - t0):.3f}") @@ -247,7 +251,7 @@ def get_prompts_with_weights(prompt: str): def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: - device = pipe.device if str(pipe.device) != 'meta' else devices.device + device = devices.device embeddings_providers = [] if 'StableCascade' in pipe.__class__.__name__: embedding_type = -(clip_skip) @@ -272,7 +276,7 @@ def prepare_embedding_providers(pipe, clip_skip) -> list[EmbeddingsProvider]: def pad_to_same_length(pipe, embeds, empty_embedding_providers=None): if not hasattr(pipe, 'encode_prompt') and 'StableCascade' not in pipe.__class__.__name__: return embeds - device = pipe.device if str(pipe.device) != 'meta' else devices.device + device = devices.device if shared.opts.diffusers_zeros_prompt_pad or 'StableDiffusion3' in pipe.__class__.__name__: empty_embed = [torch.zeros((1, 77, embeds[0].shape[2]), device=device, dtype=embeds[0].dtype)] else: @@ -317,7 +321,7 @@ def split_prompts(prompt, SD3 = False): def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None): - device = pipe.device if str(pipe.device) != 'meta' else devices.device + device = devices.device SD3 = hasattr(pipe, 'text_encoder_3') prompt, prompt_2, prompt_3 = split_prompts(prompt, SD3) neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(neg_prompt, SD3) @@ -418,7 +422,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c if prompt_embeds.shape[1] != negative_prompt_embeds.shape[1]: [prompt_embeds, negative_prompt_embeds] = pad_to_same_length(pipe, [prompt_embeds, negative_prompt_embeds], empty_embedding_providers=empty_embedding_providers) if SD3: - device = pipe.device if str(pipe.device) != 'meta' else devices.device + device = devices.device t5_prompt_embed = pipe._get_t5_prompt_embeds( # pylint: disable=protected-access prompt=prompt_3, num_images_per_prompt=prompt_embeds.shape[0], diff --git a/modules/sd_models.py b/modules/sd_models.py index a842c5f2e..518584cb6 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -762,6 +762,56 @@ def set_diffuser_offload(sd_model, op: str = 'model'): else: sd_model.enable_sequential_cpu_offload(device=devices.device) sd_model.has_accelerate = True + if shared.opts.diffusers_offload_mode == "balanced": + sd_model = apply_balanced_offload(sd_model) + + +def apply_balanced_offload(sd_model): + from accelerate import infer_auto_device_map, dispatch_model + from accelerate.hooks import add_hook_to_module, remove_hook_from_module, ModelHook + + class dispatch_from_cpu_hook(ModelHook): + def init_hook(self, module): + device_index = torch.device(devices.device).index + if device_index is None: + device_index = 0 + max_memory = {device_index: f"{shared.opts.diffusers_offload_max_gpu_memory}GiB", "cpu": f"{shared.opts.diffusers_offload_max_cpu_memory}GiB"} + self.device_map = infer_auto_device_map(module, max_memory=max_memory) + model_name = sd_model.sd_checkpoint_info.name if getattr(sd_model, "sd_checkpoint_info", None) is not None else None + if model_name is None: + model_name = "" + self.accelerate_offload_path = os.path.join(shared.opts.accelerate_offload_path, model_name) + return module + def pre_forward(self, module, *args, **kwargs): + if normalize_device(module.device) != normalize_device(devices.device): + module = remove_hook_from_module(module, recurse=True) + module = dispatch_model(module, device_map=self.device_map, offload_dir=self.accelerate_offload_path) + module = add_hook_to_module(module, dispatch_from_cpu_hook(), append=True) + module._hf_hook.execution_device = torch.device(devices.device) + return args, kwargs + def post_forward(self, module, output): + return output + def detach_hook(self, module): + return module + + def apply_balanced_offload_to_module(pipe): + module_names, _ = pipe._get_signature_keys(pipe) # pylint: disable=protected-access + for module in module_names: + module = getattr(pipe, module) + if isinstance(module, torch.nn.Module): + module = remove_hook_from_module(module, recurse=True) + module = module.to("cpu") + module = add_hook_to_module(module, dispatch_from_cpu_hook(), append=True) + module._hf_hook.execution_device = torch.device(devices.device) + + apply_balanced_offload_to_module(sd_model) + if hasattr(sd_model, "prior_pipe"): + apply_balanced_offload_to_module(sd_model.prior_pipe) + if hasattr(sd_model, "decoder_pipe"): + apply_balanced_offload_to_module(sd_model.decoder_pipe) + sd_model.has_accelerate = True + return sd_model + def normalize_device(device): if torch.device(device).type in {"cpu", "mps", "meta"}: @@ -926,14 +976,6 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No "requires_safety_checker": False, # "use_safetensors": True, } - if shared.opts.diffusers_offload_mode == "balanced": - diffusers_load_config['device_map'] = "balanced" - if shared.opts.diffusers_offload_max_memory > 0: - device_index = torch.device(devices.device).index - if device_index is None: - device_index = 0 - diffusers_load_config['max_memory'] = {device_index:f"{shared.opts.diffusers_offload_max_memory}GB"} - if shared.opts.diffusers_model_load_variant != 'default': diffusers_load_config['variant'] = shared.opts.diffusers_model_load_variant if shared.opts.diffusers_pipeline == 'Custom Diffusers Pipeline' and len(shared.opts.custom_diffusers_pipeline) > 0: @@ -1217,8 +1259,6 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No insert_parser_highjack(sd_model.__class__.__name__) set_diffuser_options(sd_model, vae, op, offload=False) - if shared.opts.diffusers_offload_mode == "balanced": - sd_model.has_accelerate = True if shared.opts.nncf_compress_weights and not (shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"): sd_model = sd_models_compile.nncf_compress_weights(sd_model) # run this before move model so it can be compressed in CPU if shared.opts.optimum_quanto_weights: diff --git a/modules/shared.py b/modules/shared.py index 4e8ba659d..68b9acddc 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -3,6 +3,7 @@ import os import sys import time import json +import psutil import threading import contextlib from types import SimpleNamespace @@ -354,6 +355,7 @@ def temp_disable_extensions(): gpu_memory = 0 offload_mode_default = "none" +cpu_memory = psutil.virtual_memory().total / 1024 / 1024 / 1024 mem_stat = memory_stats() if "gpu" in mem_stat: @@ -536,7 +538,8 @@ options_templates.update(options_section(('diffusers', "Diffusers Settings"), { "diffusers_extract_ema": OptionInfo(False, "Use model EMA weights when possible"), "diffusers_generator_device": OptionInfo("GPU", "Generator device", gr.Radio, {"choices": ["GPU", "CPU", "Unset"]}), "diffusers_offload_mode": OptionInfo(offload_mode_default, "Model offload mode", gr.Radio, {"choices": ['none', 'balanced', 'cpu', 'sequential']}), - "diffusers_offload_max_memory": OptionInfo(gpu_memory * 0.8, "Max memory for balanced offload mode in GB", gr.Slider, {"minimum": 0, "maximum": gpu_memory, "step": 0.1,}), + "diffusers_offload_max_gpu_memory": OptionInfo(gpu_memory * 0.8, "Max GPU memory for balanced offload mode in GB", gr.Slider, {"minimum": 0, "maximum": gpu_memory, "step": 0.1,}), + "diffusers_offload_max_cpu_memory": OptionInfo(cpu_memory * 0.8, "Max CPU memory for balanced offload mode in GB", gr.Slider, {"minimum": 0, "maximum": cpu_memory, "step": 0.1,}), "diffusers_vae_upcast": OptionInfo("default", "VAE upcasting", gr.Radio, {"choices": ['default', 'true', 'false']}), "diffusers_vae_slicing": OptionInfo(True, "VAE slicing"), "diffusers_vae_tiling": OptionInfo(cmd_opts.lowvram or cmd_opts.medvram, "VAE tiling"), @@ -585,6 +588,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), { "clip_models_path": OptionInfo(os.path.join(paths.models_path, 'CLIP'), "Folder with CLIP models", folder=True), "other_paths_sep_options": OptionInfo("

Other paths

", "", gr.HTML), "openvino_cache_path": OptionInfo('cache', "Directory for OpenVINO cache", folder=True), + "accelerate_offload_path": OptionInfo('cache/accelerate', "Directory for disk offload with Accelerate", folder=True), "onnx_cached_models_path": OptionInfo(os.path.join(paths.models_path, 'ONNX', 'cache'), "Folder with ONNX cached models", folder=True), "onnx_temp_dir": OptionInfo(os.path.join(paths.models_path, 'ONNX', 'temp'), "Directory for ONNX conversion and Olive optimization process", folder=True), "temp_dir": OptionInfo("", "Directory for temporary images; leave empty for default", folder=True),