diff --git a/configs/olive/sdxl_vae_encoder.json b/configs/olive/sdxl_vae_encoder.json index 2ab771f6f..fd74a5c98 100644 --- a/configs/olive/sdxl_vae_encoder.json +++ b/configs/olive/sdxl_vae_encoder.json @@ -51,7 +51,7 @@ } }, "passes": { - "optimize": { + "optimize_DmlExecutionProvider": { "type": "OrtTransformersOptimization", "disable_search": true, "config": { diff --git a/modules/olive.py b/modules/olive.py index 2ea9df851..1faf7742e 100644 --- a/modules/olive.py +++ b/modules/olive.py @@ -1,16 +1,77 @@ import os import torch import diffusers +from typing import Type, Callable, Dict, Any from transformers.models.clip.modeling_clip import CLIPTextModel, CLIPTextModelWithProjection -is_sdxl = False +class ENVStore: + __DESERIALIZER: Dict[Type, Callable[[str,], Any]] = { + bool: lambda x: bool(int(x)), + int: int, + } + __SERIALIZER: Dict[Type, Callable[[Any,], str]] = { + bool: lambda x: str(int(x)), + int: str, + } -width = 512 -height = 512 -batch_size = 1 -cross_attention_dim = 768 -time_ids_size = 5 + def __getattr__(self, name: str): + value = os.environ.get(f"SDNEXT_OLIVE_{name}", None) + if value is None: + return + ty = self.__class__.__annotations__[name] + deserialize = self.__DESERIALIZER[ty] + return deserialize(value) + + def __setattr__(self, name: str, value) -> None: + if name not in self.__class__.__annotations__: + return + ty = self.__class__.__annotations__[name] + serialize = self.__SERIALIZER[ty] + os.environ[f"SDNEXT_OLIVE_{name}"] = serialize(value) + + def __delattr__(self, name: str) -> None: + if name not in self.__class__.__annotations__: + return + os.environ.pop(f"SDNEXT_OLIVE_{name}") + + +class OliveOptimizerConfig(ENVStore): + from_huggingface_cache: bool + + is_sdxl: bool + + width: int + height: int + batch_size: int + cross_attention_dim: int + time_ids_size: int + + +config = OliveOptimizerConfig() + + +def get_variant(): + from modules.shared import opts + if opts.diffusers_model_load_variant == 'default': + from modules import devices + if devices.dtype == torch.float16: + return 'fp16' + elif opts.diffusers_model_load_variant == 'fp32': + return None + else: + return opts.diffusers_model_load_variant + + +def get_loader_arguments(): + if config.from_huggingface_cache: + from modules.shared import opts + return { + "cache_dir": opts.diffusers_dir, + "variant": get_variant(), + } + + return {} # ------------------------------------------------------------------------- @@ -36,15 +97,15 @@ class RandomDataLoader: def text_encoder_inputs(_, torch_dtype): - input_ids = torch.zeros((batch_size, 77), dtype=torch_dtype) + input_ids = torch.zeros((config.batch_size, 77), dtype=torch_dtype) return { "input_ids": input_ids, "output_hidden_states": True, - } if is_sdxl else input_ids + } if config.is_sdxl else input_ids def text_encoder_load(model_name): - model = CLIPTextModel.from_pretrained(model_name, subfolder="text_encoder") + model = CLIPTextModel.from_pretrained(model_name, subfolder="text_encoder", **get_loader_arguments()) return model @@ -53,7 +114,7 @@ def text_encoder_conversion_inputs(model): def text_encoder_data_loader(data_dir, _, *args, **kwargs): - return RandomDataLoader(text_encoder_inputs, batch_size, torch.int32) + return RandomDataLoader(text_encoder_inputs, config.batch_size, torch.int32) # ----------------------------------------------------------------------------- @@ -63,13 +124,13 @@ def text_encoder_data_loader(data_dir, _, *args, **kwargs): def text_encoder_2_inputs(_, torch_dtype): return { - "input_ids": torch.zeros((batch_size, 77), dtype=torch_dtype), + "input_ids": torch.zeros((config.batch_size, 77), dtype=torch_dtype), "output_hidden_states": True, } def text_encoder_2_load(model_name): - model = CLIPTextModelWithProjection.from_pretrained(model_name, subfolder="text_encoder_2") + model = CLIPTextModelWithProjection.from_pretrained(model_name, subfolder="text_encoder_2", **get_loader_arguments()) return model @@ -78,7 +139,7 @@ def text_encoder_2_conversion_inputs(model): def text_encoder_2_data_loader(data_dir, _, *args, **kwargs): - return RandomDataLoader(text_encoder_2_inputs, batch_size, torch.int64) + return RandomDataLoader(text_encoder_2_inputs, config.batch_size, torch.int64) # ----------------------------------------------------------------------------- @@ -87,28 +148,28 @@ def text_encoder_2_data_loader(data_dir, _, *args, **kwargs): def unet_inputs(_, torch_dtype, is_conversion_inputs=False): - if is_sdxl: + if config.is_sdxl: inputs = { - "sample": torch.rand((2 * batch_size, 4, height // 8, width // 8), dtype=torch_dtype), + "sample": torch.rand((2 * config.batch_size, 4, config.height // 8, config.width // 8), dtype=torch_dtype), "timestep": torch.rand((1,), dtype=torch_dtype), - "encoder_hidden_states": torch.rand((2 * batch_size, 77, cross_attention_dim), dtype=torch_dtype), + "encoder_hidden_states": torch.rand((2 * config.batch_size, 77, config.cross_attention_dim), dtype=torch_dtype), } if is_conversion_inputs: inputs["additional_inputs"] = { "added_cond_kwargs": { - "text_embeds": torch.rand((2 * batch_size, 1280), dtype=torch_dtype), - "time_ids": torch.rand((2 * batch_size, time_ids_size), dtype=torch_dtype), + "text_embeds": torch.rand((2 * config.batch_size, 1280), dtype=torch_dtype), + "time_ids": torch.rand((2 * config.batch_size, config.time_ids_size), dtype=torch_dtype), } } else: - inputs["text_embeds"] = torch.rand((2 * batch_size, 1280), dtype=torch_dtype) - inputs["time_ids"] = torch.rand((2 * batch_size, time_ids_size), dtype=torch_dtype) + inputs["text_embeds"] = torch.rand((2 * config.batch_size, 1280), dtype=torch_dtype) + inputs["time_ids"] = torch.rand((2 * config.batch_size, config.time_ids_size), dtype=torch_dtype) else: inputs = { - "sample": torch.rand((batch_size, 4, height // 8, width // 8), dtype=torch_dtype), - "timestep": torch.rand((batch_size,), dtype=torch_dtype), - "encoder_hidden_states": torch.rand((batch_size, 77, cross_attention_dim), dtype=torch_dtype), + "sample": torch.rand((config.batch_size, 4, config.height // 8, config.width // 8), dtype=torch_dtype), + "timestep": torch.rand((config.batch_size,), dtype=torch_dtype), + "encoder_hidden_states": torch.rand((config.batch_size, 77, config.cross_attention_dim), dtype=torch_dtype), } # use as kwargs since they won't be in the correct position if passed along with the tuple of inputs @@ -132,7 +193,7 @@ def unet_inputs(_, torch_dtype, is_conversion_inputs=False): def unet_load(model_name): - model = diffusers.UNet2DConditionModel.from_pretrained(model_name, subfolder="unet") + model = diffusers.UNet2DConditionModel.from_pretrained(model_name, subfolder="unet", **get_loader_arguments()) return model @@ -141,7 +202,7 @@ def unet_conversion_inputs(model): def unet_data_loader(data_dir, _, *args, **kwargs): - return RandomDataLoader(unet_inputs, batch_size, torch.float16) + return RandomDataLoader(unet_inputs, config.batch_size, torch.float16) # ----------------------------------------------------------------------------- @@ -151,16 +212,13 @@ def unet_data_loader(data_dir, _, *args, **kwargs): def vae_encoder_inputs(_, torch_dtype): return { - "sample": torch.rand((batch_size, 3, height, width), dtype=torch_dtype), + "sample": torch.rand((config.batch_size, 3, config.height, config.width), dtype=torch_dtype), "return_dict": False, } def vae_encoder_load(model_name): - source = os.path.join(model_name, "vae") - if not os.path.isdir(source): - source += "_encoder" - model = diffusers.AutoencoderKL.from_pretrained(source) + model = diffusers.AutoencoderKL.from_pretrained(model_name, subfolder="vae_encoder" if os.path.isdir(os.path.join(model_name, "vae_encoder")) else "vae", **get_loader_arguments()) model.forward = lambda sample, return_dict: model.encode(sample, return_dict)[0].sample() return model @@ -170,7 +228,7 @@ def vae_encoder_conversion_inputs(model): def vae_encoder_data_loader(data_dir, _, *args, **kwargs): - return RandomDataLoader(vae_encoder_inputs, batch_size, torch.float16) + return RandomDataLoader(vae_encoder_inputs, config.batch_size, torch.float16) # ----------------------------------------------------------------------------- @@ -180,16 +238,13 @@ def vae_encoder_data_loader(data_dir, _, *args, **kwargs): def vae_decoder_inputs(_, torch_dtype): return { - "latent_sample": torch.rand((batch_size, 4, height // 8, width // 8), dtype=torch_dtype), + "latent_sample": torch.rand((config.batch_size, 4, config.height // 8, config.width // 8), dtype=torch_dtype), "return_dict": False, } def vae_decoder_load(model_name): - source = os.path.join(model_name, "vae") - if not os.path.isdir(source): - source += "_decoder" - model = diffusers.AutoencoderKL.from_pretrained(source) + model = diffusers.AutoencoderKL.from_pretrained(model_name, subfolder="vae_decoder" if os.path.isdir(os.path.join(model_name, "vae_decoder")) else "vae", **get_loader_arguments()) model.forward = model.decode return model @@ -199,4 +254,4 @@ def vae_decoder_conversion_inputs(model): def vae_decoder_data_loader(data_dir, _, *args, **kwargs): - return RandomDataLoader(vae_decoder_inputs, batch_size, torch.float16) + return RandomDataLoader(vae_decoder_inputs, config.batch_size, torch.float16) diff --git a/modules/onnx.py b/modules/onnx.py index b78ef9640..ed8e3dbc5 100644 --- a/modules/onnx.py +++ b/modules/onnx.py @@ -14,11 +14,14 @@ from abc import ABCMeta from typing import Union, Optional, Callable, Type, Tuple, List, Any, Dict from diffusers.pipelines.onnx_utils import ORT_TO_NP_TYPE from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput +from diffusers.pipelines.stable_diffusion_xl import StableDiffusionXLPipelineOutput from diffusers.image_processor import VaeImageProcessor, PipelineImageInput +from optimum.pipelines.diffusers.pipeline_stable_diffusion_xl import rescale_noise_cfg from installer import log -from modules import shared, olive +from modules import shared from modules.paths import sd_configs_path from modules.sd_models import CheckpointInfo +from modules.olive import config class ExecutionProvider(str, Enum): CPU = "CPUExecutionProvider" @@ -78,6 +81,28 @@ def get_execution_provider_options(): return execution_provider_options +def get_provider() -> Tuple: + return (shared.opts.onnx_execution_provider, get_execution_provider_options(),) + + +def get_sess_options(batch_size: int, height: int, width: int, is_sdxl: bool) -> ort.SessionOptions: + sess_options = ort.SessionOptions() + sess_options.enable_mem_pattern = False + sess_options.add_free_dimension_override_by_name("unet_sample_batch", batch_size * 2) + sess_options.add_free_dimension_override_by_name("unet_sample_channels", 4) + sess_options.add_free_dimension_override_by_name("unet_sample_height", height // 8) + sess_options.add_free_dimension_override_by_name("unet_sample_width", width // 8) + sess_options.add_free_dimension_override_by_name("unet_time_batch", 1) + sess_options.add_free_dimension_override_by_name("unet_hidden_batch", batch_size * 2) + sess_options.add_free_dimension_override_by_name("unet_hidden_sequence", 77) + if is_sdxl: + sess_options.add_free_dimension_override_by_name("unet_text_embeds_batch", batch_size * 2) + sess_options.add_free_dimension_override_by_name("unet_text_embeds_size", 1280) + sess_options.add_free_dimension_override_by_name("unet_time_ids_batch", batch_size * 2) + sess_options.add_free_dimension_override_by_name("unet_time_ids_size", 6) + return sess_options + + class OnnxFakeModule: device = torch.device("cpu") dtype = torch.float32 @@ -117,30 +142,36 @@ def load_init_dict(cls: Type[diffusers.DiffusionPipeline], path: os.PathLike): return R -def load_submodel(path: os.PathLike, submodel_name: str, item: List[Union[str, None]], **kwargs): +def check_pipeline_sdxl(cls: Type[diffusers.DiffusionPipeline]) -> bool: + return 'XL' in cls.__name__ + + +def load_submodel(path: os.PathLike, is_sdxl: bool, submodel_name: str, item: List[Union[str, None]], **kwargs_ort): lib, atr = item if lib is None or atr is None: return None library = importlib.import_module(lib) attribute = getattr(library, atr) - if not issubclass(attribute, diffusers.OnnxRuntimeModel): - kwargs.clear() - return attribute.from_pretrained( - os.path.join(path, submodel_name), - **kwargs, - ) + path = os.path.join(path, submodel_name) + if issubclass(attribute, diffusers.OnnxRuntimeModel): + return diffusers.OnnxRuntimeModel.load_model( + os.path.join(path, "model.onnx"), + **kwargs_ort, + ) if is_sdxl else diffusers.OnnxRuntimeModel.from_pretrained( + path, + **kwargs_ort, + ) + return attribute.from_pretrained(path) -def load_submodels(path: os.PathLike, init_dict: Dict[str, Type], **kwargs): - sess_options = kwargs.get("sess_options", ort.SessionOptions()) - provider = kwargs.get("provider", (shared.opts.onnx_execution_provider, get_execution_provider_options(),)) +def load_submodels(path: os.PathLike, is_sdxl: bool, init_dict: Dict[str, Type], **kwargs_ort): loaded = {} for k, v in init_dict.items(): if not isinstance(v, list): loaded[k] = v continue try: - loaded[k] = load_submodel(path, k, v, sess_options=sess_options, provider=provider) + loaded[k] = load_submodel(path, is_sdxl, k, v, **kwargs_ort) except Exception: pass return loaded @@ -150,13 +181,15 @@ def patch_kwargs(cls: Type[diffusers.DiffusionPipeline], kwargs: Dict) -> Dict: if cls == OnnxStableDiffusionPipeline or cls == OnnxStableDiffusionImg2ImgPipeline or cls == OnnxStableDiffusionInpaintPipeline: kwargs["safety_checker"] = None kwargs["requires_safety_checker"] = False + if cls == OnnxStableDiffusionXLPipeline or cls == OnnxStableDiffusionXLImg2ImgPipeline: + kwargs["config"] = {} return kwargs -def load_pipeline(cls: Type[diffusers.DiffusionPipeline], path: os.PathLike): +def load_pipeline(cls: Type[diffusers.DiffusionPipeline], path: os.PathLike, **kwargs_ort): if os.path.isdir(path): - return cls(**patch_kwargs(cls, load_submodels(path, load_init_dict(cls, path)))) + return cls(**patch_kwargs(cls, load_submodels(path, check_pipeline_sdxl(cls), load_init_dict(cls, path), **kwargs_ort))) else: return cls.from_single_file(path) @@ -192,28 +225,28 @@ class OnnxPipelineBase(OnnxFakeModule, diffusers.DiffusionPipeline, metaclass=AB class OnnxRawPipeline(OnnxPipelineBase): config = {} _is_sdxl: bool + from_huggingface_cache: bool path: os.PathLike original_filename: str constructor: Type[OnnxPipelineBase] submodels: List[str] - load_runtime_model: Callable init_dict: Dict[str, Tuple[str]] = {} scheduler: Any = None # for Img2Img def __init__(self, constructor: Type[OnnxPipelineBase], path: os.PathLike): self.model_type = constructor.__name__ - self._is_sdxl = 'XL' in self.model_type + self._is_sdxl = check_pipeline_sdxl(constructor) + self.from_huggingface_cache = shared.opts.diffusers_dir in os.path.abspath(path) self.path = path self.original_filename = os.path.basename(path) self.constructor = constructor self.submodels = submodels_sdxl if self._is_sdxl else submodels_sd - self.load_runtime_model = diffusers.OnnxRuntimeModel.load_model if self._is_sdxl else diffusers.OnnxRuntimeModel.from_pretrained if os.path.isdir(path): self.init_dict = load_init_dict(constructor, path) - self.scheduler = load_submodel(self.path, "scheduler", self.init_dict["scheduler"]) + self.scheduler = load_submodel(self.path, None, "scheduler", self.init_dict["scheduler"]) else: try: cls = None @@ -242,6 +275,9 @@ class OnnxRawPipeline(OnnxPipelineBase): return pipeline def convert(self, in_dir: os.PathLike): + if not shared.cmd_opts.debug: + ort.set_default_logger_severity(3) + out_dir = os.path.join(shared.opts.onnx_cached_models_path, self.original_filename) if os.path.isdir(out_dir): # already converted (cached) return out_dir @@ -285,23 +321,6 @@ class OnnxRawPipeline(OnnxPipelineBase): log.info(f"Converted {submodel}") - kwargs = {} - - init_dict = self.init_dict.copy() - for submodel in self.submodels: - kwargs[submodel] = self.load_runtime_model( - os.path.dirname(converted_model_paths[submodel]), - provider=(shared.opts.onnx_execution_provider, get_execution_provider_options(),), - ) - if submodel in init_dict: - del init_dict[submodel] # already loaded as OnnxRuntimeModel. - kwargs.update(load_submodels(in_dir, init_dict)) # load others. - kwargs = patch_kwargs(self.constructor, kwargs) - - pipeline = self.constructor(**kwargs) - pipeline.to_json_file(os.path.join(out_dir, "model_index.json")) - del pipeline - for submodel in self.submodels: src_path = converted_model_paths[submodel] src_parent = os.path.dirname(src_path) @@ -311,10 +330,44 @@ class OnnxRawPipeline(OnnxPipelineBase): os.mkdir(dst_parent) shutil.copyfile(src_path, dst_path) + data_src_path = os.path.join(src_parent, (os.path.basename(src_path) + ".data")) + if os.path.isfile(data_src_path): + data_dst_path = os.path.join(dst_parent, (os.path.basename(dst_path) + ".data")) + shutil.copyfile(data_src_path, data_dst_path) + weights_src_path = os.path.join(src_parent, "weights.pb") if os.path.isfile(weights_src_path): weights_dst_path = os.path.join(dst_parent, "weights.pb") shutil.copyfile(weights_src_path, weights_dst_path) + del converted_model_paths + + kwargs = {} + + init_dict = self.init_dict.copy() + for submodel in self.submodels: + kwargs[submodel] = diffusers.OnnxRuntimeModel.load_model( + os.path.join(out_dir, submodel, "model.onnx"), + provider=get_provider(), + ) if self._is_sdxl else diffusers.OnnxRuntimeModel.from_pretrained( + os.path.join(out_dir, submodel), + provider=get_provider(), + ) + if submodel in init_dict: + del init_dict[submodel] # already loaded as OnnxRuntimeModel. + kwargs.update(load_submodels(in_dir, self._is_sdxl, init_dict)) # load others. + kwargs = patch_kwargs(self.constructor, kwargs) + + pipeline = self.constructor(**kwargs) + model_index = json.loads(pipeline.to_json_string()) + del pipeline + + for k, v in init_dict.items(): # copy missing submodels. (ORTStableDiffusionXLPipeline) + if k not in model_index: + model_index[k] = v + + with open(os.path.join(out_dir, "model_index.json"), 'w') as file: + json.dump(model_index, file) + return out_dir except Exception as e: log.error(f"Failed to convert model '{self.original_filename}'.") @@ -324,21 +377,10 @@ class OnnxRawPipeline(OnnxPipelineBase): return None def optimize(self, in_dir: os.PathLike): - sess_options = ort.SessionOptions() - sess_options.add_free_dimension_override_by_name("unet_sample_batch", olive.batch_size * 2) - sess_options.add_free_dimension_override_by_name("unet_sample_channels", 4) - sess_options.add_free_dimension_override_by_name("unet_sample_height", olive.height // 8) - sess_options.add_free_dimension_override_by_name("unet_sample_width", olive.width // 8) - sess_options.add_free_dimension_override_by_name("unet_time_batch", 1) - sess_options.add_free_dimension_override_by_name("unet_hidden_batch", olive.batch_size * 2) - sess_options.add_free_dimension_override_by_name("unet_hidden_sequence", 77) - if olive.is_sdxl: - sess_options.add_free_dimension_override_by_name("unet_text_embeds_batch", olive.batch_size * 2) - sess_options.add_free_dimension_override_by_name("unet_text_embeds_size", 1280) - sess_options.add_free_dimension_override_by_name("unet_time_ids_batch", olive.batch_size * 2) - sess_options.add_free_dimension_override_by_name("unet_time_ids_size", 6) + if not shared.cmd_opts.debug: + ort.set_default_logger_severity(4) - out_dir = os.path.join(shared.opts.onnx_cached_models_path, f"{self.original_filename}-{olive.width}w-{olive.height}h") + out_dir = os.path.join(shared.opts.onnx_cached_models_path, f"{self.original_filename}-{config.width}w-{config.height}h") if os.path.isdir(out_dir): # already optimized (cached) return out_dir @@ -389,23 +431,6 @@ class OnnxRawPipeline(OnnxPipelineBase): log.info(f"Optimized {submodel}") - kwargs = {} - - init_dict = self.init_dict.copy() - for submodel in self.submodels: - kwargs[submodel] = self.load_runtime_model( - os.path.dirname(optimized_model_paths[submodel]), - sess_options=sess_options, - provider=(shared.opts.onnx_execution_provider, get_execution_provider_options(),), - ) - if submodel in init_dict: - del init_dict[submodel] # already loaded as OnnxRuntimeModel. - kwargs.update(load_submodels(in_dir, init_dict)) # load others. - kwargs = patch_kwargs(self.constructor, kwargs) - - pipeline = self.constructor(**kwargs) - pipeline.to_json_file(os.path.join(out_dir, "model_index.json")) - for submodel in self.submodels: src_path = optimized_model_paths[submodel] src_parent = os.path.dirname(src_path) @@ -415,10 +440,44 @@ class OnnxRawPipeline(OnnxPipelineBase): os.mkdir(dst_parent) shutil.copyfile(src_path, dst_path) - weights_src_path = os.path.join(src_parent, (os.path.basename(src_path) + ".data")) + data_src_path = os.path.join(src_parent, (os.path.basename(src_path) + ".data")) + if os.path.isfile(data_src_path): + data_dst_path = os.path.join(dst_parent, (os.path.basename(dst_path) + ".data")) + shutil.copyfile(data_src_path, data_dst_path) + + weights_src_path = os.path.join(src_parent, "weights.pb") if os.path.isfile(weights_src_path): - weights_dst_path = os.path.join(dst_parent, (os.path.basename(dst_path) + ".data")) + weights_dst_path = os.path.join(dst_parent, "weights.pb") shutil.copyfile(weights_src_path, weights_dst_path) + del optimized_model_paths + + kwargs = {} + + init_dict = self.init_dict.copy() + for submodel in self.submodels: + kwargs[submodel] = diffusers.OnnxRuntimeModel.load_model( + os.path.join(out_dir, submodel, "model.onnx"), + provider=get_provider(), + ) if self._is_sdxl else diffusers.OnnxRuntimeModel.from_pretrained( + os.path.join(out_dir, submodel), + provider=get_provider(), + ) + if submodel in init_dict: + del init_dict[submodel] # already loaded as OnnxRuntimeModel. + kwargs.update(load_submodels(in_dir, self._is_sdxl, init_dict)) # load others. + kwargs = patch_kwargs(self.constructor, kwargs) + + pipeline = self.constructor(**kwargs) + model_index = json.loads(pipeline.to_json_string()) + del pipeline + + for k, v in init_dict.items(): # copy missing submodels. (ORTStableDiffusionXLPipeline) + if k not in model_index: + model_index[k] = v + + with open(os.path.join(out_dir, "model_index.json"), 'w') as file: + json.dump(model_index, file) + return out_dir except Exception as e: log.error(f"Failed to optimize model '{self.original_filename}'.") @@ -427,25 +486,31 @@ class OnnxRawPipeline(OnnxPipelineBase): shutil.rmtree(out_dir, ignore_errors=True) return None - def preprocess(self, width: int, height: int, batch_size: int): - if not shared.cmd_opts.debug: - ort.set_default_logger_severity(3) - olive.width = width - olive.height = height - olive.batch_size = batch_size + def preprocess(self, batch_size: int, height: int, width: int): + config.from_huggingface_cache = self.from_huggingface_cache - olive.is_sdxl = self._is_sdxl - if olive.is_sdxl: - olive.cross_attention_dim = 2048 - olive.time_ids_size = 6 + config.is_sdxl = self._is_sdxl + + config.width = width + config.height = height + config.batch_size = batch_size + + if self._is_sdxl: + config.cross_attention_dim = 2048 + config.time_ids_size = 6 else: - olive.cross_attention_dim = height + 256 - olive.time_ids_size = 5 + config.cross_attention_dim = height + 256 + config.time_ids_size = 5 + + kwargs = { + "provider": get_provider(), + "sess_options": get_sess_options(batch_size, height, width, self._is_sdxl), + } converted_dir = self.convert(self.path if os.path.isdir(self.path) else shared.opts.onnx_temp_dir) if converted_dir is None: log.error('Failed to convert model. The generation will fall back to unconverted one.') - return self.derive_properties(load_pipeline(diffusers.StableDiffusionXLPipeline if self._is_sdxl else diffusers.StableDiffusionPipeline, self.path)) + return self.derive_properties(load_pipeline(diffusers.StableDiffusionXLPipeline if self._is_sdxl else diffusers.StableDiffusionPipeline, self.path, **kwargs)) out_dir = converted_dir if shared.opts.onnx_enable_olive: @@ -455,18 +520,52 @@ class OnnxRawPipeline(OnnxPipelineBase): optimized_dir = self.optimize(converted_dir) if optimized_dir is None: log.error('Failed to optimize pipeline. The generation will fall back to unoptimized one.') - return self.derive_properties(load_pipeline(diffusers.OnnxStableDiffusionXLPipeline if self._is_sdxl else diffusers.OnnxStableDiffusionPipeline, converted_dir)) + return self.derive_properties(load_pipeline(diffusers.OnnxStableDiffusionXLPipeline if self._is_sdxl else diffusers.OnnxStableDiffusionPipeline, converted_dir, **kwargs)) out_dir = optimized_dir - pipeline = self.derive_properties(load_pipeline(diffusers.OnnxStableDiffusionXLPipeline if self._is_sdxl else diffusers.OnnxStableDiffusionPipeline, out_dir)) + pipeline = self.derive_properties(load_pipeline(diffusers.OnnxStableDiffusionXLPipeline if self._is_sdxl else diffusers.OnnxStableDiffusionPipeline, out_dir, **kwargs)) if not shared.opts.onnx_cache_converted: shutil.rmtree(converted_dir) - shutil.rmtree(shared.opts.onnx_temp_dir) + shutil.rmtree(shared.opts.onnx_temp_dir, ignore_errors=True) return pipeline +def prepare_latents( + scheduler, + batch_size: int, + height: int, + width: int, + dtype: torch.dtype, + generator: Union[torch.Generator, List[torch.Generator]], + latents: Union[np.ndarray, None]=None, + num_channels_latents=4, + vae_scale_factor=8, +): + shape = (batch_size, num_channels_latents, height // vae_scale_factor, width // vae_scale_factor) + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + if latents is None: + if isinstance(generator, list): + generator = [g.seed() for g in generator] + if len(generator) == 1: + generator = generator[0] + + latents = np.random.default_rng(generator).standard_normal(shape).astype(dtype) + elif latents.shape != shape: + raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}") + + # scale the initial noise by the standard deviation required by the scheduler + latents = latents * np.float64(scheduler.init_noise_sigma) + + return latents + + class OnnxStableDiffusionPipeline(diffusers.OnnxStableDiffusionPipeline, OnnxPipelineBase): def __init__( self, @@ -514,6 +613,9 @@ class OnnxStableDiffusionPipeline(diffusers.OnnxStableDiffusionPipeline, OnnxPip else: batch_size = prompt_embeds.shape[0] + if generator is None: + generator = torch.Generator("cpu") + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` # corresponds to doing no classifier free guidance. @@ -529,23 +631,19 @@ class OnnxStableDiffusionPipeline(diffusers.OnnxStableDiffusionPipeline, OnnxPip ) # get the initial random noise unless the user supplied it - latents_dtype = prompt_embeds.dtype - latents_shape = (batch_size * num_images_per_prompt, 4, height // 8, width // 8) - if latents is None: - if isinstance(generator, list): - generator = [g.seed() for g in generator] - if len(generator) == 1: - generator = generator[0] - - latents = np.random.default_rng(generator).standard_normal(latents_shape).astype(latents_dtype) - elif latents.shape != latents_shape: - raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {latents_shape}") + latents = prepare_latents( + self.scheduler, + batch_size * num_images_per_prompt, + height, + width, + prompt_embeds.dtype, + generator, + latents + ) # set timesteps self.scheduler.set_timesteps(num_inference_steps) - latents = latents * np.float64(self.scheduler.init_noise_sigma) - # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers. # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502 @@ -586,7 +684,7 @@ class OnnxStableDiffusionPipeline(diffusers.OnnxStableDiffusionPipeline, OnnxPip if callback is not None and i % callback_steps == 0: callback(i, t, torch.from_numpy(latents)) - latents = 1 / 0.18215 * latents + latents /= self.vae_decoder.config.get("scaling_factor", 0.18215) has_nsfw_concept = None @@ -676,6 +774,9 @@ class OnnxStableDiffusionImg2ImgPipeline(diffusers.OnnxStableDiffusionImg2ImgPip else: batch_size = prompt_embeds.shape[0] + if generator is None: + generator = torch.Generator("cpu") + if strength < 0 or strength > 1: raise ValueError(f"The value of strength should in [0.0, 1.0] but is {strength}") @@ -698,11 +799,13 @@ class OnnxStableDiffusionImg2ImgPipeline(diffusers.OnnxStableDiffusionImg2ImgPip negative_prompt_embeds=negative_prompt_embeds, ) + scaling_factor = self.vae_decoder.config.get("scaling_factor", 0.18215) + latents_dtype = prompt_embeds.dtype image = image.astype(latents_dtype) # encode the init image into latents and scale the latents init_latents = self.vae_encoder(sample=image)[0] - init_latents = 0.18215 * init_latents + init_latents = scaling_factor * init_latents if isinstance(prompt, str): prompt = [prompt] @@ -775,7 +878,7 @@ class OnnxStableDiffusionImg2ImgPipeline(diffusers.OnnxStableDiffusionImg2ImgPip if callback is not None and i % callback_steps == 0: callback(i, t, torch.from_numpy(latents)) - latents = 1 / 0.18215 * latents + latents /= scaling_factor has_nsfw_concept = None @@ -874,6 +977,9 @@ class OnnxStableDiffusionInpaintPipeline(diffusers.OnnxStableDiffusionInpaintPip else: batch_size = prompt_embeds.shape[0] + if generator is None: + generator = torch.Generator("cpu") + # set timesteps self.scheduler.set_timesteps(num_inference_steps) @@ -905,13 +1011,15 @@ class OnnxStableDiffusionInpaintPipeline(diffusers.OnnxStableDiffusionInpaintPip if latents.shape != latents_shape: raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {latents_shape}") + scaling_factor = self.vae_decoder.config.get("scaling_factor", 0.18215) + # prepare mask and masked_image mask, masked_image = diffusers.pipelines.stable_diffusion.pipeline_onnx_stable_diffusion_inpaint.prepare_mask_and_masked_image(image[0], mask_image, latents_shape[-2:]) mask = mask.astype(latents.dtype) masked_image = masked_image.astype(latents.dtype) masked_image_latents = self.vae_encoder(sample=masked_image)[0] - masked_image_latents = 0.18215 * masked_image_latents + masked_image_latents = scaling_factor * masked_image_latents # duplicate mask and masked_image_latents for each generation per prompt mask = mask.repeat(batch_size * num_images_per_prompt, 0) @@ -985,7 +1093,7 @@ class OnnxStableDiffusionInpaintPipeline(diffusers.OnnxStableDiffusionInpaintPip step_idx = i // getattr(self.scheduler, "order", 1) callback(step_idx, t, torch.from_numpy(latents)) - latents = 1 / 0.18215 * latents + latents /= scaling_factor has_nsfw_concept = None @@ -1033,21 +1141,191 @@ diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["onnx-stable-di class OnnxStableDiffusionXLPipeline(OnnxPipelineBase, optimum.onnxruntime.ORTStableDiffusionXLPipeline): def __init__( self, - vae_decoder_session, - text_encoder_session, - unet_session, + vae_decoder, + text_encoder, + unet, config: Dict[str, Any], tokenizer, scheduler, feature_extractor = None, - vae_encoder_session = None, - text_encoder_2_session = None, + vae_encoder = None, + text_encoder_2 = None, tokenizer_2 = None, use_io_binding: bool | None = None, model_save_dir = None, add_watermarker: bool | None = None ): - super(optimum.onnxruntime.ORTStableDiffusionXLPipeline, self).__init__(vae_decoder_session, text_encoder_session, unet_session, config, tokenizer, scheduler, feature_extractor, vae_encoder_session, text_encoder_2_session, tokenizer_2, use_io_binding, model_save_dir, add_watermarker) + super(optimum.onnxruntime.ORTStableDiffusionXLPipeline, self).__init__(vae_decoder, text_encoder, unet, config, tokenizer, scheduler, feature_extractor, vae_encoder, text_encoder_2, tokenizer_2, use_io_binding, model_save_dir, add_watermarker) + + # Adapted from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl.StableDiffusionXLPipeline.__call__ + def __call__( + self, + prompt: Optional[Union[str, List[str]]] = None, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 50, + guidance_scale: float = 5.0, + negative_prompt: Optional[Union[str, List[str]]] = None, + num_images_per_prompt: int = 1, + eta: float = 0.0, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[np.ndarray] = None, + prompt_embeds: Optional[np.ndarray] = None, + negative_prompt_embeds: Optional[np.ndarray] = None, + pooled_prompt_embeds: Optional[np.ndarray] = None, + negative_pooled_prompt_embeds: Optional[np.ndarray] = None, + output_type: str = "pil", + return_dict: bool = True, + callback: Optional[Callable[[int, int, np.ndarray], None]] = None, + callback_steps: int = 1, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + guidance_rescale: float = 0.0, + original_size: Optional[Tuple[int, int]] = None, + crops_coords_top_left: Tuple[int, int] = (0, 0), + target_size: Optional[Tuple[int, int]] = None, + ): + # 0. Default height and width to unet + height = height or self.unet.config["sample_size"] * self.vae_scale_factor + width = width or self.unet.config["sample_size"] * self.vae_scale_factor + + original_size = original_size or (height, width) + target_size = target_size or (height, width) + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + height, + width, + callback_steps, + negative_prompt, + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + negative_pooled_prompt_embeds, + ) + + # 2. Define call parameters + if isinstance(prompt, str): + batch_size = 1 + elif isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + if generator is None: + generator = torch.Generator("cpu") + + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) + # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` + # corresponds to doing no classifier free guidance. + do_classifier_free_guidance = guidance_scale > 1.0 + + # 3. Encode input prompt + ( + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + negative_pooled_prompt_embeds, + ) = self._encode_prompt( + prompt, + num_images_per_prompt, + do_classifier_free_guidance, + negative_prompt, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + negative_pooled_prompt_embeds=negative_pooled_prompt_embeds, + ) + + # 4. Prepare timesteps + self.scheduler.set_timesteps(num_inference_steps) + timesteps = self.scheduler.timesteps + + # 5. Prepare latent variables + latents = prepare_latents( + self.scheduler, + batch_size * num_images_per_prompt, + height, + width, + prompt_embeds.dtype, + generator, + latents, + self.unet.config.get("in_channels", 4), + self.vae_scale_factor, + ) + + # 6. Prepare extra step kwargs + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + + # 7. Prepare added time ids & embeddings + add_text_embeds = pooled_prompt_embeds + add_time_ids = (original_size + crops_coords_top_left + target_size,) + add_time_ids = np.array(add_time_ids, dtype=prompt_embeds.dtype) + + if do_classifier_free_guidance: + prompt_embeds = np.concatenate((negative_prompt_embeds, prompt_embeds), axis=0) + add_text_embeds = np.concatenate((negative_pooled_prompt_embeds, add_text_embeds), axis=0) + add_time_ids = np.concatenate((add_time_ids, add_time_ids), axis=0) + add_time_ids = np.repeat(add_time_ids, batch_size * num_images_per_prompt, axis=0) + + # Adapted from diffusers to extend it for other runtimes than ORT + timestep_dtype = self.unet.input_dtype.get("timestep", np.float32) + + # 8. Denoising loop + num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order + for i, t in enumerate(self.progress_bar(timesteps)): + # expand the latents if we are doing classifier free guidance + latent_model_input = np.concatenate([latents] * 2) if do_classifier_free_guidance else latents + latent_model_input = self.scheduler.scale_model_input(torch.from_numpy(latent_model_input), t) + latent_model_input = latent_model_input.cpu().numpy() + + # predict the noise residual + timestep = np.array([t], dtype=timestep_dtype) + noise_pred = self.unet( + sample=latent_model_input, + timestep=timestep, + encoder_hidden_states=prompt_embeds, + text_embeds=add_text_embeds, + time_ids=add_time_ids, + ) + noise_pred = noise_pred[0] + + # perform guidance + if do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = np.split(noise_pred, 2) + noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond) + if guidance_rescale > 0.0: + # Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf + noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=guidance_rescale) + + # compute the previous noisy sample x_t -> x_t-1 + scheduler_output = self.scheduler.step( + torch.from_numpy(noise_pred), t, torch.from_numpy(latents), **extra_step_kwargs + ) + latents = scheduler_output.prev_sample.numpy() + + # call the callback, if provided + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + if callback is not None and i % callback_steps == 0: + callback(i, t, torch.from_numpy(latents)) + + if output_type == "latent": + image = latents + else: + latents /= self.vae_decoder.config.get("scaling_factor", 0.18215) + # it seems likes there is a strange result for using half-precision vae decoder if batchsize>1 + image = np.concatenate( + [self.vae_decoder(latent_sample=latents[i : i + 1])[0] for i in range(latents.shape[0])] + ) + # apply watermark if available + if self.watermark is not None: + image = self.watermark.apply_watermark(image) + image = self.image_processor.postprocess(image, output_type=output_type) + + if not return_dict: + return (image,) + + return StableDiffusionXLPipelineOutput(images=image) OnnxStableDiffusionXLPipeline.__module__ = 'optimum.onnxruntime.modeling_diffusion' @@ -1059,21 +1337,21 @@ diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["onnx-stable class OnnxStableDiffusionXLImg2ImgPipeline(OnnxPipelineBase, optimum.onnxruntime.ORTStableDiffusionXLImg2ImgPipeline): def __init__( self, - vae_decoder_session, - text_encoder_session, - unet_session, + vae_decoder, + text_encoder, + unet, config: Dict[str, Any], tokenizer, scheduler, feature_extractor = None, - vae_encoder_session = None, - text_encoder_2_session = None, + vae_encoder = None, + text_encoder_2 = None, tokenizer_2 = None, use_io_binding: bool | None = None, model_save_dir = None, add_watermarker: bool | None = None ): - super(optimum.onnxruntime.ORTStableDiffusionXLImg2ImgPipeline, self).__init__(vae_decoder_session, text_encoder_session, unet_session, config, tokenizer, scheduler, feature_extractor, vae_encoder_session, text_encoder_2_session, tokenizer_2, use_io_binding, model_save_dir, add_watermarker) + super(optimum.onnxruntime.ORTStableDiffusionXLImg2ImgPipeline, self).__init__(vae_decoder, text_encoder, unet, config, tokenizer, scheduler, feature_extractor, vae_encoder, text_encoder_2, tokenizer_2, use_io_binding, model_save_dir, add_watermarker) OnnxStableDiffusionXLImg2ImgPipeline.__module__ = 'optimum.onnxruntime.modeling_diffusion' diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 74caefd2d..9dcddaf98 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -21,7 +21,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing): results = [] if hasattr(shared.sd_model, 'preprocess'): - shared.sd_model = shared.sd_model.preprocess(p.width, p.height, p.batch_size) + shared.sd_model = shared.sd_model.preprocess(p.batch_size, p.height, p.width) def is_txt2img(): return sd_models.get_diffusers_task(shared.sd_model) == sd_models.DiffusersTaskType.TEXT_2_IMAGE