Olive SDXL optimization.

This commit is contained in:
Seunghoon Lee
2023-12-17 03:13:46 +09:00
parent 86b56d2d2c
commit f10b257dd9
4 changed files with 486 additions and 153 deletions
+92 -37
View File
@@ -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)
+392 -114
View File
@@ -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'
+1 -1
View File
@@ -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