From 0f44332e5cf4aa48749380e46c40e2f8a0e81183 Mon Sep 17 00:00:00 2001 From: Seunghoon Lee Date: Fri, 28 Jul 2023 23:11:57 +0900 Subject: [PATCH] Make sequential CPU offload available for non-CUDA Add settings override for DirectML. Move `devices.set_cuda_params()` to correct line. --- modules/devices.py | 6 ++---- modules/dml/__init__.py | 29 +++++++++++++++++++++++++++-- modules/dml/hijack/torch.py | 12 ------------ modules/sd_models.py | 4 ++-- modules/shared.py | 3 +-- modules/ui.py | 3 +++ 6 files changed, 35 insertions(+), 22 deletions(-) diff --git a/modules/devices.py b/modules/devices.py index 8c57b54af..8a30e2ee4 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -3,7 +3,6 @@ import sys import contextlib import torch from modules import cmd_args, shared, memstats -from modules.dml import directml_init if sys.platform == "darwin": from modules import mac_specific # pylint: disable=ungrouped-imports @@ -172,6 +171,8 @@ if args.use_ipex or (hasattr(torch, 'xpu') and torch.xpu.is_available()): ipex_init() elif args.use_directml: backend = 'directml' + from modules.dml import directml_init + directml_init() elif torch.cuda.is_available() and torch.version.cuda: backend = 'cuda' elif torch.cuda.is_available() and torch.version.hip: @@ -181,9 +182,6 @@ elif sys.platform == 'darwin': else: backend = 'cpu' -if backend == "directml": - directml_init() - cuda_ok = torch.cuda.is_available() and not backend == 'ipex' cpu = torch.device("cpu") device = device_interrogate = device_gfpgan = device_esrgan = device_codeformer = None diff --git a/modules/dml/__init__.py b/modules/dml/__init__.py index a656b970f..ae3d2a7b2 100644 --- a/modules/dml/__init__.py +++ b/modules/dml/__init__.py @@ -1,5 +1,6 @@ import os import torch +from typing import NamedTuple, Callable from modules.sd_hijack_utils import CondFunc @@ -41,7 +42,31 @@ def directml_do_hijack(): lambda orig_func, *args, **kwargs: orig_func(args[0].astype('float32')), lambda *args, **kwargs: args[1].dtype == float) +class OverrideItem(NamedTuple): + value: str + condition: Callable | None + message: str | None + +opts_override_table = { + "diffusers_generator_device": OverrideItem("cpu", None, "DirectML does not support torch Generator API."), + "diffusers_model_cpu_offload": OverrideItem(False, None, "Diffusers' model CPU offloading does not support DirectML devices."), + "diffusers_seq_cpu_offload": OverrideItem(False, lambda opts: opts.diffusers_pipeline != "Stable Diffusion XL", "Diffusers' sequential CPU offloading is available only on StableDiffusionXLPipeline with DirectML devices."), +} + def directml_override_opts(): from modules import shared - if shared.backend == shared.Backend.DIFFUSERS: - shared.opts.diffusers_generator_device = "cpu" # DirectML does not support torch.Generator API. + + if shared.cmd_opts.experimental: + return + + count = 0 + for key in opts_override_table: + count += 1 + item = opts_override_table[key] + if getattr(shared.opts, key) != item.value and (item.condition is None or item.condition(shared.opts)): + setattr(shared.opts, key, item.value) + shared.log.warning(item.message) + shared.log.warning(f'{key} is automatically overriden to {item.value}.') + + if count > 0: + shared.log.info(f'{count} options are automatically overriden. If you want to keep them from overriding, run with --experimental argument.') diff --git a/modules/dml/hijack/torch.py b/modules/dml/hijack/torch.py index e3d8350df..edada0c9a 100644 --- a/modules/dml/hijack/torch.py +++ b/modules/dml/hijack/torch.py @@ -2,19 +2,7 @@ import torch from modules.sd_hijack_utils import CondFunc -def to_sub(orig, self: torch.Tensor, *args, **kwargs): - def validate(device: torch.device | str): - if torch.dml.is_directml_device(torch.device(device)): - raise NotImplementedError("Cannot copy out of meta tensor; no data!") - for arg in args: - validate(arg) - if "device" in kwargs: - validate(kwargs["device"]) - return orig(self, *args, **kwargs) - CondFunc('torchsde._brownian.brownian_interval._randn', lambda _, size, dtype, device, seed: torch.randn(size, dtype=dtype, device=torch.device("cpu"), generator=torch.Generator(torch.device("cpu")).manual_seed(int(seed))).to(device), lambda _, size, dtype, device, seed: device.type == 'privateuseone') # https://github.com/microsoft/DirectML/issues/400 CondFunc('torch.Tensor.new', lambda orig, self, *args, **kwargs: orig(self.cpu(), *args, **kwargs), lambda orig, self, *args, **kwargs: torch.dml.is_directml_device(self.device)) -# https://github.com/microsoft/DirectML/issues/477 -CondFunc('torch.Tensor.to', to_sub, lambda orig, self, *args, **kwargs: self.device.type == "meta") diff --git a/modules/sd_models.py b/modules/sd_models.py index db22edf2f..803223142 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -533,6 +533,7 @@ def change_backend(): def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): # pylint: disable=unused-argument import torch # pylint: disable=reimported,redefined-outer-name + devices.set_cuda_params() if timer is None: timer = Timer() import logging @@ -564,7 +565,6 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No sd_model = None try: - devices.set_cuda_params() if shared.cmd_opts.ckpt is not None and model_data.initial: # initial load ckpt_basename = os.path.basename(shared.cmd_opts.ckpt) model_name = modelloader.find_diffuser(ckpt_basename) @@ -659,7 +659,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No sd_model.enable_model_cpu_offload() if hasattr(sd_model, "enable_sequential_cpu_offload"): if shared.opts.diffusers_seq_cpu_offload: - sd_model.enable_sequential_cpu_offload() + sd_model.enable_sequential_cpu_offload(device=devices.device) shared.log.debug(f'Diffusers {op}: enable sequential CPU offload') if hasattr(sd_model, "enable_vae_slicing"): if shared.cmd_opts.lowvram or shared.opts.diffusers_vae_slicing: diff --git a/modules/shared.py b/modules/shared.py index 029714d45..df7a23771 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -12,7 +12,7 @@ import requests import fasteners 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 -from modules.dml import directml_do_hijack, directml_override_opts +from modules.dml import directml_do_hijack import modules.interrogate import modules.memmon import modules.styles @@ -799,7 +799,6 @@ mem_mon = modules.memmon.MemUsageMonitor("MemMon", device, opts) mem_mon.start() if devices.backend == "directml": directml_do_hijack() - directml_override_opts() def reload_gradio_theme(theme_name=None): diff --git a/modules/ui.py b/modules/ui.py index f06aba6a9..0698739e1 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -14,6 +14,7 @@ from modules import sd_hijack, sd_models, script_callbacks, ui_extensions, deepb from modules.ui_components import FormRow, FormColumn, FormGroup, ToolButton, FormHTML # pylint: disable=unused-import from modules.paths import script_path, data_path from modules.shared import opts, cmd_opts, readfile +from modules.dml import directml_override_opts from modules import prompt_parser import modules.codeformer_model import modules.generation_parameters_copypaste as parameters_copypaste @@ -950,6 +951,7 @@ def create_ui(startup_timer = None): continue if opts.set(key, value): changed.append(key) + directml_override_opts() try: opts.save(modules.shared.config_filename) modules.shared.log.info(f'Settings changed: {len(changed)} {changed}') @@ -963,6 +965,7 @@ def create_ui(startup_timer = None): return gr.update(visible=True), opts.dumpjson() if not opts.set(key, value): return gr.update(value=getattr(opts, key)), opts.dumpjson() + directml_override_opts() opts.save(modules.shared.config_filename) modules.shared.log.debug(f'Setting changed: key={key}, value={value}') return get_value_for_setting(key), opts.dumpjson()