Support LoRA and Control Net in Stable Diffusion demo (#18593)

### Description
(1) Export onnx model with LoRA weights for both SD 1.5 and SDXL
(2) Export onnx model with Control Net for both SD 1.5 and SDXL. For SD
1.5, it is allowed to use multiple control nets. For SDXL, at most one
control net is supported right now.
(3) Add demo of LCM LoRA
(3) Add demo of control net.
This commit is contained in:
Tianlei Wu 2023-11-28 15:46:42 -08:00 committed by GitHub
parent 50e6235af1
commit f13380f3d8
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
14 changed files with 1042 additions and 210 deletions

View file

@ -83,8 +83,21 @@ For example:
If you do not provide prompt, the script will generate different image sizes for a list of prompts for demonstration.
#### Generate an image with SDXL LCM guided by a text prompt
```python3 demo_txt2img_xl.py --lcm --disable-refiner "an astronaut riding a rainbow unicorn, cinematic, dramatic"```
### Generate an image guided by a text prompt using LCM LoRA
```
python3 demo_txt2img_xl.py "Self-portrait oil painting, a beautiful cyborg with golden hair, 8k" --scheduler LCM --lora-weights latent-consistency/lcm-lora-sdxl --denoising-steps 4
```
#### Generate an image with SDXL LCM model guided by a text prompt
```
python3 demo_txt2img_xl.py --lcm --disable-refiner "an astronaut riding a rainbow unicorn, cinematic, dramatic"
```
#### Generate an image with a text prompt using a control net
```
python3 demo_txt2img.py "Stormtrooper's lecture in beautiful lecture hall" --controlnet-type depth --controlnet-scale 1.0
python3 demo_txt2img_xl.py "young Mona Lisa" --controlnet-type canny --controlnet-scale 0.5 --scheduler UniPC --disable-refiner
```
## Optimize Stable Diffusion ONNX models for Hugging Face Diffusers or Optimum
@ -482,7 +495,7 @@ Most ROCm kernel optimizations are from [composable kernel](https://github.com/R
Some kernels are enabled by MIOpen. We hereby thank for the AMD developers' collaboration.
### Future Works
* Update demo to support inpainting, LoRA Weights and Control Net.
* Update demo to support inpainting.
* Support flash attention in Windows.
* Integration with UI.
* Optimization for H100 GPU.

View file

@ -22,7 +22,16 @@
import coloredlogs
from cuda import cudart
from demo_utils import get_metadata, init_pipeline, parse_arguments, repeat_prompt
from demo_utils import (
add_controlnet_arguments,
arg_parser,
get_metadata,
init_pipeline,
max_batch,
parse_arguments,
process_controlnet_arguments,
repeat_prompt,
)
from diffusion_models import PipelineInfo
from engine_builder import EngineType, get_engine_type
from pipeline_txt2img import Txt2ImgPipeline
@ -30,7 +39,12 @@ from pipeline_txt2img import Txt2ImgPipeline
if __name__ == "__main__":
coloredlogs.install(fmt="%(funcName)20s: %(message)s")
args = parse_arguments(is_xl=False, description="Options for Stable Diffusion Demo")
parser = arg_parser("Options for Stable Diffusion Demo")
add_controlnet_arguments(parser)
args = parse_arguments(is_xl=False, parser=parser)
controlnet_images, controlnet_scale = process_controlnet_arguments(args)
prompt, negative_prompt = repeat_prompt(args)
image_height = args.height
@ -43,9 +57,7 @@ if __name__ == "__main__":
init_trt_plugins()
max_batch_size = 16
if engine_type != EngineType.ORT_CUDA and (args.build_dynamic_shape or image_height > 512 or image_width > 512):
max_batch_size = 4
max_batch_size = max_batch(args)
batch_size = len(prompt)
if batch_size > max_batch_size:
@ -58,7 +70,15 @@ if __name__ == "__main__":
# This range can cover common used shape of landscape 512x768, portrait 768x512, or square 512x512 and 768x768.
min_image_size = 512 if args.engine != "ORT_CUDA" else 256
max_image_size = 768 if args.engine != "ORT_CUDA" else 1024
pipeline_info = PipelineInfo(args.version, min_image_size=min_image_size, max_image_size=max_image_size)
pipeline_info = PipelineInfo(
args.version,
min_image_size=min_image_size,
max_image_size=max_image_size,
do_classifier_free_guidance=(args.guidance > 1.0),
controlnet=args.controlnet_type,
lora_weights=args.lora_weights,
lora_scale=args.lora_scale,
)
# Ideally, the optimized batch size and image size for TRT engine shall align with user's preference. That is to
# optimize the shape used most frequently. We can let user config it when we develop a UI plugin.
@ -99,6 +119,8 @@ if __name__ == "__main__":
denoising_steps=args.denoising_steps,
guidance=args.guidance,
seed=args.seed,
controlnet_images=controlnet_images,
controlnet_scales=controlnet_scale,
return_type="image",
)

View file

@ -22,7 +22,16 @@
import coloredlogs
from cuda import cudart
from demo_utils import get_metadata, init_pipeline, parse_arguments, repeat_prompt
from demo_utils import (
add_controlnet_arguments,
arg_parser,
get_metadata,
init_pipeline,
max_batch,
parse_arguments,
process_controlnet_arguments,
repeat_prompt,
)
from diffusion_models import PipelineInfo
from engine_builder import EngineType, get_engine_type
from pipeline_img2img_xl import Img2ImgXLPipeline
@ -37,11 +46,7 @@ def load_pipelines(args, batch_size):
init_trt_plugins()
max_batch_size = 16
if (engine_type in [EngineType.ORT_TRT, EngineType.TRT]) and (
args.build_dynamic_shape or args.height > 512 or args.width > 512
):
max_batch_size = 4
max_batch_size = max_batch(args)
if batch_size > max_batch_size:
raise ValueError(f"Batch size {batch_size} is larger than allowed {max_batch_size}.")
@ -59,6 +64,10 @@ def load_pipelines(args, batch_size):
min_image_size=min_image_size,
max_image_size=max_image_size,
use_lcm=args.lcm,
do_classifier_free_guidance=(args.guidance > 1.0),
controlnet=args.controlnet_type,
lora_weights=args.lora_weights,
lora_scale=args.lora_scale,
)
# Ideally, the optimized batch size and image size for TRT engine shall align with user's preference. That is to
@ -113,7 +122,9 @@ def load_pipelines(args, batch_size):
return base, refiner
def run_pipelines(args, base, refiner, prompt, negative_prompt, is_warm_up=False):
def run_pipelines(
args, base, refiner, prompt, negative_prompt, controlnet_image=None, controlnet_scale=None, is_warm_up=False
):
image_height = args.height
image_width = args.width
batch_size = len(prompt)
@ -131,6 +142,8 @@ def run_pipelines(args, base, refiner, prompt, negative_prompt, is_warm_up=False
denoising_steps=args.denoising_steps,
guidance=args.guidance,
seed=args.seed,
controlnet_images=controlnet_image,
controlnet_scales=controlnet_scale,
return_type="latent" if refiner else "image",
)
if refiner is None:
@ -180,9 +193,9 @@ def run_pipelines(args, base, refiner, prompt, negative_prompt, is_warm_up=False
cudart.cudaProfilerStop()
if refiner:
print("|------------|--------------|")
print("| {:^10} | {:>9.2f} ms |".format("e2e", perf_data["latency"]))
print("|------------|--------------|")
print("|----------------|--------------|")
print("| {:^14} | {:>9.2f} ms |".format("e2e", perf_data["latency"]))
print("|----------------|--------------|")
metadata = get_metadata(args, True)
metadata.update({"base." + key: val for key, val in base.metadata().items()})
@ -197,11 +210,11 @@ def run_pipelines(args, base, refiner, prompt, negative_prompt, is_warm_up=False
def run_demo(args):
"""Run Stable Diffusion XL Base + Refiner together (known as ensemble of expert denoisers) to generate an image."""
controlnet_image, controlnet_scale = process_controlnet_arguments(args)
prompt, negative_prompt = repeat_prompt(args)
batch_size = len(prompt)
base, refiner = load_pipelines(args, batch_size)
run_pipelines(args, base, refiner, prompt, negative_prompt)
run_pipelines(args, base, refiner, prompt, negative_prompt, controlnet_image, controlnet_scale)
base.teardown()
if refiner:
refiner.teardown()
@ -294,7 +307,10 @@ def run_dynamic_shape_demo(args):
if __name__ == "__main__":
coloredlogs.install(fmt="%(funcName)20s: %(message)s")
args = parse_arguments(is_xl=True, description="Options for Stable Diffusion XL Demo")
parser = arg_parser("Options for Stable Diffusion XL Demo")
add_controlnet_arguments(parser)
args = parse_arguments(is_xl=True, parser=parser)
no_prompt = isinstance(args.prompt, list) and len(args.prompt) == 1 and not args.prompt[0]
if no_prompt:
run_dynamic_shape_demo(args)

View file

@ -19,22 +19,33 @@
# See the License for the specific language governing permissions and
# limitations under the License.
# --------------------------------------------------------------------------
import argparse
from typing import Any, Dict
import os
import sys
from importlib.metadata import PackageNotFoundError, version
from io import BytesIO
from typing import Any, Dict, List
import controlnet_aux
import cv2
import numpy as np
import requests
import torch
from diffusers.utils import load_image
from diffusion_models import PipelineInfo
from engine_builder import EngineType, get_engine_paths
from PIL import Image
class RawTextArgumentDefaultsHelpFormatter(argparse.ArgumentDefaultsHelpFormatter, argparse.RawTextHelpFormatter):
pass
def parse_arguments(is_xl: bool, description: str):
parser = argparse.ArgumentParser(description=description, formatter_class=RawTextArgumentDefaultsHelpFormatter)
def arg_parser(description: str):
return argparse.ArgumentParser(description=description, formatter_class=RawTextArgumentDefaultsHelpFormatter)
def parse_arguments(is_xl: bool, parser):
engines = ["ORT_CUDA", "ORT_TRT", "TRT"]
parser.add_argument(
@ -69,7 +80,7 @@ def parse_arguments(is_xl: bool, description: str):
"--scheduler",
type=str,
default="DDIM",
choices=["DDIM", "UniPC", "LCM"] if is_xl else ["DDIM", "EulerA", "UniPC"],
choices=["DDIM", "UniPC", "LCM"] if is_xl else ["DDIM", "EulerA", "UniPC", "LCM"],
help="Scheduler for diffusion process" + " of base" if is_xl else "",
)
@ -106,6 +117,11 @@ def parse_arguments(is_xl: bool, description: str):
help="Higher guidance scale encourages to generate images that are closely linked to the text prompt.",
)
parser.add_argument(
"--lora-scale", type=float, default=1, help="Scale of LoRA weights, default 1 (must between 0 and 1)"
)
parser.add_argument("--lora-weights", type=str, default="", help="LoRA weights to apply in the base model")
if is_xl:
parser.add_argument(
"--lcm",
@ -142,6 +158,10 @@ def parse_arguments(is_xl: bool, description: str):
help="A value between 0 and 1. The higher the value less the final image similar to the seed image.",
)
parser.add_argument(
"--disable-refiner", action="store_true", help="Disable refiner and only run base for XL pipeline."
)
# ONNX export
parser.add_argument(
"--onnx-opset",
@ -182,10 +202,6 @@ def parse_arguments(is_xl: bool, description: str):
parser.add_argument("--seed", type=int, default=None, help="Seed for random generator to get consistent results.")
parser.add_argument("--disable-cuda-graph", action="store_true", help="Disable cuda graph.")
parser.add_argument(
"--disable-refiner", action="store_true", help="Disable refiner and only run base for XL pipeline."
)
group = parser.add_argument_group("Options for ORT_CUDA engine only")
group.add_argument("--enable-vae-slicing", action="store_true", help="True will feed only one image to VAE once.")
@ -228,25 +244,39 @@ def parse_arguments(is_xl: bool, description: str):
args.onnx_opset = 14 if args.engine == "ORT_CUDA" else 17
if is_xl:
if args.lcm:
if args.guidance > 1.0:
print("[I] Use --guidance=1.0 for base since LCM is used.")
args.guidance = 1.0
if args.scheduler != "LCM":
print("[I] Use --scheduler=LCM for base since LCM is used.")
args.scheduler = "LCM"
if args.denoising_steps > 16:
print("[I] Use --denoising_steps=8 (no more than 16) for base since LCM is used.")
args.denoising_steps = 8
if args.lcm and args.scheduler != "LCM":
print("[I] Use --scheduler=LCM for base since LCM is used.")
args.scheduler = "LCM"
assert args.strength > 0.0 and args.strength < 1.0
assert not (args.lcm and args.lora_weights), "it is not supported to use both lcm unet and Lora together"
if args.scheduler == "LCM":
if args.guidance > 1.0:
print("[I] Use --guidance=1.0 for base since LCM is used.")
args.guidance = 1.0
if args.denoising_steps > 16:
print("[I] Use --denoising_steps=8 (no more than 16) for base since LCM is used.")
args.denoising_steps = 8
print(args)
return args
def max_batch(args):
do_classifier_free_guidance = args.guidance > 1.0
batch_multiplier = 2 if do_classifier_free_guidance else 1
max_batch_size = 32 // batch_multiplier
if args.engine != "ORT_CUDA" and (args.build_dynamic_shape or args.height > 512 or args.width > 512):
max_batch_size = 8 // batch_multiplier
return max_batch_size
def get_metadata(args, is_xl: bool = False) -> Dict[str, Any]:
metadata = {
"command": " ".join(['"' + x + '"' if " " in x else x for x in sys.argv]),
"args.prompt": args.prompt,
"args.negative_prompt": args.negative_prompt,
"args.batch_size": args.batch_size,
@ -257,6 +287,14 @@ def get_metadata(args, is_xl: bool = False) -> Dict[str, Any]:
"engine": args.engine,
}
if args.lora_weights:
metadata["lora_weights"] = args.lora_weights
metadata["lora_scale"] = args.lora_scale
if args.controlnet_type:
metadata["controlnet_type"] = args.controlnet_type
metadata["controlnet_scale"] = args.controlnet_scale
if is_xl and not args.disable_refiner:
metadata["base.scheduler"] = args.scheduler
metadata["base.denoising_steps"] = args.denoising_steps
@ -270,6 +308,27 @@ def get_metadata(args, is_xl: bool = False) -> Dict[str, Any]:
metadata["denoising_steps"] = args.denoising_steps
metadata["guidance"] = args.guidance
# Version of installed python packages
packages = ""
for name in [
"onnxruntime-gpu",
"torch",
"tensorrt",
"transformers",
"diffusers",
"onnx",
"onnx-graphsurgeon",
"polygraphy",
"controlnet_aux",
]:
try:
packages += (" " if packages else "") + f"{name}=={version(name)}"
except PackageNotFoundError:
continue
metadata["packages"] = packages
metadata["device"] = torch.cuda.get_device_name()
metadata["torch.version.cuda"] = torch.version.cuda
return metadata
@ -318,6 +377,7 @@ def init_pipeline(
engine_dir=engine_dir,
framework_model_dir=framework_model_dir,
onnx_dir=onnx_dir,
tmp_dir=os.path.join(args.work_dir or ".", engine_type.name, pipeline_info.short_name(), "tmp"),
force_engine_rebuild=args.force_engine_build,
device_id=torch.cuda.current_device(),
)
@ -361,3 +421,248 @@ def init_pipeline(
)
return pipeline
def get_depth_image(image):
"""
Create depth map for SDXL depth control net.
"""
from transformers import DPTFeatureExtractor, DPTForDepthEstimation
depth_estimator = DPTForDepthEstimation.from_pretrained("Intel/dpt-hybrid-midas").to("cuda")
feature_extractor = DPTFeatureExtractor.from_pretrained("Intel/dpt-hybrid-midas")
image = feature_extractor(images=image, return_tensors="pt").pixel_values.to("cuda")
with torch.no_grad(), torch.autocast("cuda"):
depth_map = depth_estimator(image).predicted_depth
depth_map = torch.nn.functional.interpolate(
depth_map.unsqueeze(1),
size=(1024, 1024),
mode="bicubic",
align_corners=False,
)
depth_min = torch.amin(depth_map, dim=[1, 2, 3], keepdim=True)
depth_max = torch.amax(depth_map, dim=[1, 2, 3], keepdim=True)
depth_map = (depth_map - depth_min) / (depth_max - depth_min)
image = torch.cat([depth_map] * 3, dim=1)
image = image.permute(0, 2, 3, 1).cpu().numpy()[0]
image = Image.fromarray((image * 255.0).clip(0, 255).astype(np.uint8))
return image
def get_canny_image(image) -> Image.Image:
"""
Create canny image for SDXL control net.
"""
image = np.array(image)
image = cv2.Canny(image, 100, 200)
image = image[:, :, None]
image = np.concatenate([image, image, image], axis=2)
image = Image.fromarray(image)
return image
def process_controlnet_images_xl(args) -> List[Image.Image]:
"""
Process control image for SDXL control net.
"""
image = None
if args.controlnet_image:
image = Image.open(args.controlnet_image[0])
else:
# If no image is provided, download an image for demo purpose.
if args.controlnet_type[0] == "canny":
image = load_image(
"https://hf.co/datasets/huggingface/documentation-images/resolve/main/diffusers/input_image_vermeer.png"
)
elif args.controlnet_type[0] == "depth":
image = load_image(
"https://huggingface.co/lllyasviel/sd-controlnet-depth/resolve/main/images/stormtrooper.png"
)
controlnet_images = []
if args.controlnet_type[0] == "canny":
controlnet_images.append(get_canny_image(image))
elif args.controlnet_type[0] == "depth":
controlnet_images.append(get_depth_image(image))
else:
raise ValueError(f"The controlnet is not supported for SDXL: {args.controlnet_type}")
return controlnet_images
def add_controlnet_arguments(parser, is_xl: bool = False):
"""
Add control net related arguments.
"""
group = parser.add_argument_group("Options for ControlNet (only supports SD 1.5 or XL).")
group.add_argument(
"--controlnet-image",
nargs="*",
type=str,
default=[],
help="Path to the input regular RGB image/images for controlnet",
)
group.add_argument(
"--controlnet-type",
nargs="*",
type=str,
default=[],
choices=list(PipelineInfo.supported_controlnet("xl-1.0" if is_xl else "1.5").keys()),
help="A list of controlnet type",
)
group.add_argument(
"--controlnet-scale",
nargs="*",
type=float,
default=[],
help="The outputs of the controlnet are multiplied by `controlnet_scale` before they are added to the residual in the original unet. Default is 0.35 for SDXL, or 1.0 for SD 1.5",
)
def download_image(url) -> Image.Image:
response = requests.get(url)
return Image.open(BytesIO(response.content)).convert("RGB")
def controlnet_demo_images(controlnet_list: List[str], height, width) -> List[Image.Image]:
"""
Return demo images of control net v1.1 for Stable Diffusion 1.5.
"""
control_images = []
shape = (height, width)
for controlnet in controlnet_list:
if controlnet == "canny":
canny_image = download_image(
"https://hf.co/datasets/huggingface/documentation-images/resolve/main/diffusers/input_image_vermeer.png"
)
canny_image = controlnet_aux.CannyDetector()(canny_image)
control_images.append(canny_image.resize(shape))
elif controlnet == "normalbae":
normal_image = download_image(
"https://huggingface.co/lllyasviel/sd-controlnet-normal/resolve/main/images/toy.png"
)
normal_image = controlnet_aux.NormalBaeDetector.from_pretrained("lllyasviel/Annotators")(normal_image)
control_images.append(normal_image.resize(shape))
elif controlnet == "depth":
depth_image = download_image(
"https://huggingface.co/lllyasviel/sd-controlnet-depth/resolve/main/images/stormtrooper.png"
)
depth_image = controlnet_aux.LeresDetector.from_pretrained("lllyasviel/Annotators")(depth_image)
control_images.append(depth_image.resize(shape))
elif controlnet == "mlsd":
mlsd_image = download_image(
"https://huggingface.co/lllyasviel/sd-controlnet-mlsd/resolve/main/images/room.png"
)
mlsd_image = controlnet_aux.MLSDdetector.from_pretrained("lllyasviel/Annotators")(mlsd_image)
control_images.append(mlsd_image.resize(shape))
elif controlnet == "openpose":
openpose_image = download_image(
"https://huggingface.co/lllyasviel/sd-controlnet-openpose/resolve/main/images/pose.png"
)
openpose_image = controlnet_aux.OpenposeDetector.from_pretrained("lllyasviel/Annotators")(openpose_image)
control_images.append(openpose_image.resize(shape))
elif controlnet == "scribble":
scribble_image = download_image(
"https://huggingface.co/lllyasviel/sd-controlnet-scribble/resolve/main/images/bag.png"
)
scribble_image = controlnet_aux.HEDdetector.from_pretrained("lllyasviel/Annotators")(
scribble_image, scribble=True
)
control_images.append(scribble_image.resize(shape))
elif controlnet == "seg":
seg_image = download_image(
"https://huggingface.co/lllyasviel/sd-controlnet-seg/resolve/main/images/house.png"
)
seg_image = controlnet_aux.SamDetector.from_pretrained(
"ybelkada/segment-anything", subfolder="checkpoints"
)(seg_image)
control_images.append(seg_image.resize(shape))
else:
raise ValueError(f"There is no demo image of this controlnet: {controlnet}")
return control_images
def process_controlnet_image(controlnet_type: str, image: Image.Image, height, width):
"""
Process control images of control net v1.1 for Stable Diffusion 1.5.
"""
control_image = None
shape = (height, width)
image = image.convert("RGB")
if controlnet_type == "canny":
canny_image = controlnet_aux.CannyDetector()(image)
control_image = canny_image.resize(shape)
elif controlnet_type == "normalbae":
normal_image = controlnet_aux.NormalBaeDetector.from_pretrained("lllyasviel/Annotators")(image)
control_image = normal_image.resize(shape)
elif controlnet_type == "depth":
depth_image = controlnet_aux.LeresDetector.from_pretrained("lllyasviel/Annotators")(image)
control_image = depth_image.resize(shape)
elif controlnet_type == "mlsd":
mlsd_image = controlnet_aux.MLSDdetector.from_pretrained("lllyasviel/Annotators")(image)
control_image = mlsd_image.resize(shape)
elif controlnet_type == "openpose":
openpose_image = controlnet_aux.OpenposeDetector.from_pretrained("lllyasviel/Annotators")(image)
control_image = openpose_image.resize(shape)
elif controlnet_type == "scribble":
scribble_image = controlnet_aux.HEDdetector.from_pretrained("lllyasviel/Annotators")(image, scribble=True)
control_image = scribble_image.resize(shape)
elif controlnet_type == "seg":
seg_image = controlnet_aux.SamDetector.from_pretrained("ybelkada/segment-anything", subfolder="checkpoints")(
image
)
control_image = seg_image.resize(shape)
else:
raise ValueError(f"There is no demo image of this controlnet_type: {controlnet_type}")
return control_image
def process_controlnet_arguments(args):
"""
Process control net arguments, and returns a list of control images and a tensor of control net scales.
"""
assert isinstance(args.controlnet_type, list)
assert isinstance(args.controlnet_scale, list)
assert isinstance(args.controlnet_image, list)
if args.version not in ["1.5", "xl-1.0"]:
raise ValueError("This demo only supports ControlNet in Stable Diffusion 1.5 or XL.")
is_xl = args.version == "xl-1.0"
if is_xl and len(args.controlnet_type) > 1:
raise ValueError("This demo only support one ControlNet for Stable Diffusion XL.")
if len(args.controlnet_image) != 0 and len(args.controlnet_image) != len(args.controlnet_scale):
raise ValueError(
f"Numbers of ControlNets {len(args.controlnet_image)} should be equal to number of ControlNet scales {len(args.controlnet_scale)}."
)
if len(args.controlnet_type) == 0:
return None, None
if len(args.controlnet_scale) == 0:
args.controlnet_scale = [0.5 if is_xl else 1.0] * len(args.controlnet_type)
elif len(args.controlnet_type) != len(args.controlnet_scale):
raise ValueError(
f"Numbers of ControlNets {len(args.controlnet_type)} should be equal to number of ControlNet scales {len(args.controlnet_scale)}."
)
# Convert controlnet scales to tensor
controlnet_scale = torch.FloatTensor(args.controlnet_scale)
if is_xl:
images = process_controlnet_images_xl(args)
else:
images = []
if len(args.controlnet_image) > 0:
for i, image in enumerate(args.controlnet_image):
images.append(
process_controlnet_image(args.controlnet_type[i], Image.open(image), args.height, args.width)
)
else:
images = controlnet_demo_images(args.controlnet_type, args.height, args.width)
return images, controlnet_scale

View file

@ -29,7 +29,7 @@ from typing import Dict, List, Optional
import onnx
import onnx_graphsurgeon as gs
import torch
from diffusers.models import AutoencoderKL, UNet2DConditionModel
from diffusers.models import AutoencoderKL, ControlNetModel, UNet2DConditionModel
from onnx import GraphProto, ModelProto, shape_inference
from ort_optimizer import OrtStableDiffusionOptimizer
from polygraphy.backend.onnx.loader import fold_constants
@ -92,6 +92,10 @@ class PipelineInfo:
max_image_size=1024,
use_fp16_vae=True,
use_lcm=False,
do_classifier_free_guidance=True,
controlnet=None,
lora_weights=None,
lora_scale=1.0,
):
self.version = version
self._is_inpaint = is_inpaint
@ -101,6 +105,11 @@ class PipelineInfo:
self._max_image_size = max_image_size
self._use_fp16_vae = use_fp16_vae
self._use_lcm = use_lcm
self.do_classifier_free_guidance = do_classifier_free_guidance and not use_lcm
self.controlnet = controlnet # A list of control net type
self.lora_weights = lora_weights
self.lora_scale = lora_scale
if is_refiner:
assert not use_lcm
assert self.is_xl()
@ -224,6 +233,41 @@ class PipelineInfo:
return 768
return 512
@staticmethod
def supported_controlnet(version="1.5"):
if version == "xl-1.0":
return {
"canny": "diffusers/controlnet-canny-sdxl-1.0",
"depth": "diffusers/controlnet-depth-sdxl-1.0",
}
elif version == "1.5":
return {
"canny": "lllyasviel/control_v11p_sd15_canny",
"depth": "lllyasviel/control_v11f1p_sd15_depth",
"openpose": "lllyasviel/control_v11p_sd15_openpose",
# "tile": "lllyasviel/control_v11f1e_sd15_tile",
# "lineart": "lllyasviel/control_v11p_sd15_lineart",
# "inpaint": "lllyasviel/control_v11p_sd15_inpaint",
# "softedge": "lllyasviel/control_v11p_sd15_softedge",
"mlsd": "lllyasviel/control_v11p_sd15_mlsd",
"scribble": "lllyasviel/control_v11p_sd15_scribble",
# "ip2p": "lllyasviel/control_v11e_sd15_ip2p",
"normalbae": "lllyasviel/control_v11p_sd15_normalbae",
"seg": "lllyasviel/control_v11p_sd15_seg",
# "shuffle": "lllyasviel/control_v11e_sd15_shuffle",
# "lineart_anime": "lllyasviel/control_v11p_sd15s2_lineart_anime",
}
return None
def controlnet_name(self):
"""Return a list of controlnet name"""
if not self.controlnet:
return None
controlnet_map = PipelineInfo.supported_controlnet(self.version)
if controlnet_map is None:
return None
return [controlnet_map[controlnet] for controlnet in self.controlnet]
class BaseModel:
def __init__(
@ -254,6 +298,9 @@ class BaseModel:
self.embedding_dim = embedding_dim
self.text_maxlen = text_maxlen
def get_batch_multiplier(self):
return 2 if self.pipeline_info.do_classifier_free_guidance else 1
def get_ort_optimizer(self):
model_name_to_model_type = {
"CLIP": "clip",
@ -316,7 +363,10 @@ class BaseModel:
_,
) = self.get_minmax_dims(batch_size, image_height, image_width, static_batch, static_image_shape)
profile_id = f"_b_{batch_size}" if static_batch else f"_b_{min_batch}_{max_batch}"
if (self.name in ["UNet", "UNetXL"]) and (self.get_batch_multiplier() == 1):
profile_id = f"_b1_{batch_size}" if static_batch else f"_b1_{min_batch}_{max_batch}"
else:
profile_id = f"_b_{batch_size}" if static_batch else f"_b_{min_batch}_{max_batch}"
if self.name != "CLIP":
if static_image_shape:
@ -348,6 +398,7 @@ class BaseModel:
fp32_op_list=None,
optimize_by_ort=True,
optimize_by_fusion=True,
tmp_dir=None,
):
optimizer = self.get_ort_optimizer()
optimizer.optimize(
@ -358,6 +409,7 @@ class BaseModel:
fp32_op_list=fp32_op_list,
optimize_by_ort=optimize_by_ort,
optimize_by_fusion=optimize_by_fusion,
tmp_dir=tmp_dir,
)
def optimize_trt(self, input_onnx_path, optimized_onnx_path):
@ -525,6 +577,7 @@ class CLIP(BaseModel):
fp32_op_list=None,
optimize_by_ort=True,
optimize_by_fusion=True,
tmp_dir=None,
):
optimizer = self.get_ort_optimizer()
@ -538,6 +591,7 @@ class CLIP(BaseModel):
keep_outputs=["text_embeddings"],
optimize_by_ort=optimize_by_ort,
optimize_by_fusion=optimize_by_fusion,
tmp_dir=tmp_dir,
)
elif optimize_by_fusion:
with tempfile.TemporaryDirectory() as tmp_dir:
@ -556,6 +610,7 @@ class CLIP(BaseModel):
keep_outputs=["text_embeddings", "hidden_states"],
optimize_by_ort=optimize_by_ort,
optimize_by_fusion=optimize_by_fusion,
tmp_dir=tmp_dir,
)
else: # input is optimized model, there is no need to add hidden states.
optimizer.optimize(
@ -567,6 +622,7 @@ class CLIP(BaseModel):
keep_outputs=["text_embeddings", "hidden_states"],
optimize_by_ort=optimize_by_ort,
optimize_by_fusion=optimize_by_fusion,
tmp_dir=tmp_dir,
)
def optimize_trt(self, input_onnx_path, optimized_onnx_path):
@ -622,6 +678,100 @@ class CLIPWithProj(CLIP):
return output
class UNet2DConditionControlNetModel(torch.nn.Module):
def __init__(self, unet, controlnets: ControlNetModel):
super().__init__()
self.unet = unet
self.controlnets = controlnets
def forward(self, sample, timestep, encoder_hidden_states, controlnet_images, controlnet_scales):
for i, (controlnet_image, conditioning_scale, controlnet) in enumerate(
zip(controlnet_images, controlnet_scales, self.controlnets)
):
down_samples, mid_sample = controlnet(
sample,
timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_cond=controlnet_image,
return_dict=False,
)
down_samples = [down_sample * conditioning_scale for down_sample in down_samples]
mid_sample *= conditioning_scale
# merge samples
if i == 0:
down_block_res_samples, mid_block_res_sample = down_samples, mid_sample
else:
down_block_res_samples = [
samples_prev + samples_curr
for samples_prev, samples_curr in zip(down_block_res_samples, down_samples)
]
mid_block_res_sample += mid_sample
noise_pred = self.unet(
sample,
timestep,
encoder_hidden_states=encoder_hidden_states,
down_block_additional_residuals=down_block_res_samples,
mid_block_additional_residual=mid_block_res_sample,
)
return noise_pred[0]
# Modified from convert_stable_diffusion_controlnet_to_onnx.py in diffusers
class UNet2DConditionXLControlNetModel(torch.nn.Module):
def __init__(self, unet, controlnets: ControlNetModel):
super().__init__()
self.unet = unet
self.controlnets = controlnets
def forward(
self,
sample,
timestep,
encoder_hidden_states,
text_embeds,
time_ids,
controlnet_images,
controlnet_scales,
):
added_cond_kwargs = {"text_embeds": text_embeds, "time_ids": time_ids}
for i, (controlnet_image, conditioning_scale, controlnet) in enumerate(
zip(controlnet_images, controlnet_scales, self.controlnets)
):
down_samples, mid_sample = controlnet(
sample,
timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_cond=controlnet_image,
conditioning_scale=conditioning_scale,
added_cond_kwargs=added_cond_kwargs,
return_dict=False,
)
# merge samples
if i == 0:
down_block_res_samples, mid_block_res_sample = down_samples, mid_sample
else:
down_block_res_samples = [
samples_prev + samples_curr
for samples_prev, samples_curr in zip(down_block_res_samples, down_samples)
]
mid_block_res_sample += mid_sample
noise_pred = self.unet(
sample,
timestep,
encoder_hidden_states=encoder_hidden_states,
down_block_additional_residuals=down_block_res_samples,
mid_block_additional_residual=mid_block_res_sample,
added_cond_kwargs=added_cond_kwargs,
return_dict=False,
)
return noise_pred[0]
class UNet(BaseModel):
def __init__(
self,
@ -642,72 +792,129 @@ class UNet(BaseModel):
embedding_dim=pipeline_info.unet_embedding_dim(),
text_maxlen=text_maxlen,
)
self.unet_dim = unet_dim
self.controlnet = pipeline_info.controlnet_name()
def load_model(self, framework_model_dir, hf_token, subfolder="unet"):
options = {"variant": "fp16", "torch_dtype": torch.float16} if self.fp16 else {}
return self.from_pretrained(UNet2DConditionModel, framework_model_dir, hf_token, subfolder, **options)
model = self.from_pretrained(UNet2DConditionModel, framework_model_dir, hf_token, subfolder, **options)
if self.controlnet:
cnet_model_opts = {"torch_dtype": torch.float16} if self.fp16 else {}
controlnets = torch.nn.ModuleList(
[ControlNetModel.from_pretrained(name, **cnet_model_opts).to(self.device) for name in self.controlnet]
)
model = UNet2DConditionControlNetModel(model, controlnets)
return model
def get_input_names(self):
return ["sample", "timestep", "encoder_hidden_states"]
if not self.controlnet:
return ["sample", "timestep", "encoder_hidden_states"]
else:
return ["sample", "timestep", "encoder_hidden_states", "controlnet_images", "controlnet_scales"]
def get_output_names(self):
return ["latent"]
def get_dynamic_axes(self):
return {
"sample": {0: "2B", 2: "H", 3: "W"},
"encoder_hidden_states": {0: "2B"},
"latent": {0: "2B", 2: "H", 3: "W"},
b = "2B" if self.get_batch_multiplier() == 2 else "B"
output = {
"sample": {0: b, 2: "H", 3: "W"},
"encoder_hidden_states": {0: b},
"latent": {0: b, 2: "H", 3: "W"},
}
if self.controlnet:
output.update(
{
"controlnet_images": {1: b, 3: "8H", 4: "8W"},
}
)
return output
def get_input_profile(self, batch_size, image_height, image_width, static_batch, static_image_shape):
latent_height, latent_width = self.check_dims(batch_size, image_height, image_width)
(
min_batch,
max_batch,
_,
_,
_,
_,
min_image_height,
max_image_height,
min_image_width,
max_image_width,
min_latent_height,
max_latent_height,
min_latent_width,
max_latent_width,
) = self.get_minmax_dims(batch_size, image_height, image_width, static_batch, static_image_shape)
return {
m = self.get_batch_multiplier()
output = {
"sample": [
(2 * min_batch, self.unet_dim, min_latent_height, min_latent_width),
(2 * batch_size, self.unet_dim, latent_height, latent_width),
(2 * max_batch, self.unet_dim, max_latent_height, max_latent_width),
(m * min_batch, self.unet_dim, min_latent_height, min_latent_width),
(m * batch_size, self.unet_dim, latent_height, latent_width),
(m * max_batch, self.unet_dim, max_latent_height, max_latent_width),
],
"encoder_hidden_states": [
(2 * min_batch, self.text_maxlen, self.embedding_dim),
(2 * batch_size, self.text_maxlen, self.embedding_dim),
(2 * max_batch, self.text_maxlen, self.embedding_dim),
(m * min_batch, self.text_maxlen, self.embedding_dim),
(m * batch_size, self.text_maxlen, self.embedding_dim),
(m * max_batch, self.text_maxlen, self.embedding_dim),
],
}
if self.controlnet:
output.update(
{
"controlnet_images": [
(len(self.controlnet), m * min_batch, 3, min_image_height, min_image_width),
(len(self.controlnet), m * batch_size, 3, image_height, image_width),
(len(self.controlnet), m * max_batch, 3, max_image_height, max_image_width),
]
}
)
return output
def get_shape_dict(self, batch_size, image_height, image_width):
latent_height, latent_width = self.check_dims(batch_size, image_height, image_width)
return {
"sample": (2 * batch_size, self.unet_dim, latent_height, latent_width),
m = self.get_batch_multiplier()
output = {
"sample": (m * batch_size, self.unet_dim, latent_height, latent_width),
"timestep": [1],
"encoder_hidden_states": (2 * batch_size, self.text_maxlen, self.embedding_dim),
"latent": (2 * batch_size, 4, latent_height, latent_width),
"encoder_hidden_states": (m * batch_size, self.text_maxlen, self.embedding_dim),
"latent": (m * batch_size, 4, latent_height, latent_width),
}
if self.controlnet:
output.update(
{
"controlnet_images": (len(self.controlnet), m * batch_size, 3, image_height, image_width),
"controlnet_scales": [len(self.controlnet)],
}
)
return output
def get_sample_input(self, batch_size, image_height, image_width):
latent_height, latent_width = self.check_dims(batch_size, image_height, image_width)
dtype = torch.float16 if self.fp16 else torch.float32
return (
m = self.get_batch_multiplier()
output = (
torch.randn(
2 * batch_size, self.unet_dim, latent_height, latent_width, dtype=torch.float32, device=self.device
m * batch_size, self.unet_dim, latent_height, latent_width, dtype=torch.float32, device=self.device
),
torch.tensor([1.0], dtype=torch.float32, device=self.device),
torch.randn(2 * batch_size, self.text_maxlen, self.embedding_dim, dtype=dtype, device=self.device),
torch.randn(m * batch_size, self.text_maxlen, self.embedding_dim, dtype=dtype, device=self.device),
)
if self.controlnet:
output = (
*output,
torch.randn(
len(self.controlnet), m * batch_size, 3, image_height, image_width, dtype=dtype, device=self.device
),
torch.randn(len(self.controlnet), dtype=dtype, device=self.device),
)
return output
def fp32_input_output_names(self) -> List[str]:
return ["sample", "timestep"]
@ -737,8 +944,7 @@ class UNetXL(BaseModel):
self.time_dim = time_dim
self.custom_unet = pipeline_info.custom_unet()
self.do_classifier_free_guidance = not (self.custom_unet and "lcm" in self.custom_unet)
self.batch_multiplier = 2 if self.do_classifier_free_guidance else 1
self.controlnet = pipeline_info.controlnet_name()
def load_model(self, framework_model_dir, hf_token, subfolder="unet"):
options = {"variant": "fp16", "torch_dtype": torch.float16} if self.fp16 else {}
@ -750,49 +956,62 @@ class UNetXL(BaseModel):
unet.save_pretrained(model_dir)
else:
unet = UNet2DConditionModel.from_pretrained(model_dir, **options)
return unet.to(self.device)
model = unet.to(self.device)
else:
model = self.from_pretrained(UNet2DConditionModel, framework_model_dir, hf_token, subfolder, **options)
return self.from_pretrained(UNet2DConditionModel, framework_model_dir, hf_token, subfolder, **options)
if self.controlnet:
cnet_model_opts = {"torch_dtype": torch.float16} if self.fp16 else {}
controlnets = torch.nn.ModuleList(
[ControlNetModel.from_pretrained(path, **cnet_model_opts).to(self.device) for path in self.controlnet]
)
model = UNet2DConditionXLControlNetModel(model, controlnets)
return model
def get_input_names(self):
return ["sample", "timestep", "encoder_hidden_states", "text_embeds", "time_ids"]
input_names = ["sample", "timestep", "encoder_hidden_states", "text_embeds", "time_ids"]
if self.controlnet:
return [*input_names, "controlnet_images", "controlnet_scales"]
return input_names
def get_output_names(self):
return ["latent"]
def get_dynamic_axes(self):
if self.do_classifier_free_guidance:
return {
"sample": {0: "2B", 2: "H", 3: "W"},
"encoder_hidden_states": {0: "2B"},
"latent": {0: "2B", 2: "H", 3: "W"},
"text_embeds": {0: "2B"},
"time_ids": {0: "2B"},
}
return {
"sample": {0: "B", 2: "H", 3: "W"},
"encoder_hidden_states": {0: "B"},
"latent": {0: "B", 2: "H", 3: "W"},
"text_embeds": {0: "B"},
"time_ids": {0: "B"},
b = "2B" if self.get_batch_multiplier() == 2 else "B"
output = {
"sample": {0: b, 2: "H", 3: "W"},
"encoder_hidden_states": {0: b},
"text_embeds": {0: b},
"time_ids": {0: b},
"latent": {0: b, 2: "H", 3: "W"},
}
if self.controlnet:
output.update(
{
"controlnet_images": {1: b, 3: "8H", 4: "8W"},
}
)
return output
def get_input_profile(self, batch_size, image_height, image_width, static_batch, static_image_shape):
latent_height, latent_width = self.check_dims(batch_size, image_height, image_width)
(
min_batch,
max_batch,
_,
_,
_,
_,
min_image_height,
max_image_height,
min_image_width,
max_image_width,
min_latent_height,
max_latent_height,
min_latent_width,
max_latent_width,
) = self.get_minmax_dims(batch_size, image_height, image_width, static_batch, static_image_shape)
m = self.batch_multiplier
return {
m = self.get_batch_multiplier()
output = {
"sample": [
(m * min_batch, self.unet_dim, min_latent_height, min_latent_width),
(m * batch_size, self.unet_dim, latent_height, latent_width),
@ -811,35 +1030,72 @@ class UNetXL(BaseModel):
],
}
if self.controlnet:
output.update(
{
"controlnet_images": [
(len(self.controlnet), m * min_batch, 3, min_image_height, min_image_width),
(len(self.controlnet), m * batch_size, 3, image_height, image_width),
(len(self.controlnet), m * max_batch, 3, max_image_height, max_image_width),
],
}
)
return output
def get_shape_dict(self, batch_size, image_height, image_width):
latent_height, latent_width = self.check_dims(batch_size, image_height, image_width)
m = self.batch_multiplier
return {
m = self.get_batch_multiplier()
output = {
"sample": (m * batch_size, self.unet_dim, latent_height, latent_width),
"timestep": (1,),
"encoder_hidden_states": (m * batch_size, self.text_maxlen, self.embedding_dim),
"latent": (m * batch_size, 4, latent_height, latent_width),
"text_embeds": (m * batch_size, 1280),
"time_ids": (m * batch_size, self.time_dim),
"latent": (m * batch_size, 4, latent_height, latent_width),
}
if self.controlnet:
output.update(
{
"controlnet_images": (len(self.controlnet), m * batch_size, 3, image_height, image_width),
"controlnet_scales": [len(self.controlnet)],
}
)
return output
def get_sample_input(self, batch_size, image_height, image_width):
latent_height, latent_width = self.check_dims(batch_size, image_height, image_width)
dtype = torch.float16 if self.fp16 else torch.float32
m = self.batch_multiplier
return (
torch.randn(
m * batch_size, self.unet_dim, latent_height, latent_width, dtype=torch.float32, device=self.device
),
torch.tensor([1.0], dtype=torch.float32, device=self.device),
torch.randn(m * batch_size, self.text_maxlen, self.embedding_dim, dtype=dtype, device=self.device),
{
"added_cond_kwargs": {
"text_embeds": torch.randn(m * batch_size, 1280, dtype=dtype, device=self.device),
"time_ids": torch.randn(m * batch_size, self.time_dim, dtype=dtype, device=self.device),
}
},
)
m = self.get_batch_multiplier()
if not self.controlnet:
return (
torch.randn(
m * batch_size, self.unet_dim, latent_height, latent_width, dtype=torch.float32, device=self.device
),
torch.tensor([1.0], dtype=torch.float32, device=self.device),
torch.randn(m * batch_size, self.text_maxlen, self.embedding_dim, dtype=dtype, device=self.device),
{
"added_cond_kwargs": {
"text_embeds": torch.randn(m * batch_size, 1280, dtype=dtype, device=self.device),
"time_ids": torch.randn(m * batch_size, self.time_dim, dtype=dtype, device=self.device),
}
},
)
else:
# sample, timestep, encoder_hidden_states, text_embeds, time_ids, controlnet_images, controlnet_scales,
return (
torch.randn(
m * batch_size, self.unet_dim, latent_height, latent_width, dtype=torch.float32, device=self.device
),
torch.tensor([1.0], dtype=torch.float32, device=self.device),
torch.randn(m * batch_size, self.text_maxlen, self.embedding_dim, dtype=dtype, device=self.device),
torch.randn(m * batch_size, 1280, dtype=dtype, device=self.device),
torch.randn(m * batch_size, self.time_dim, dtype=dtype, device=self.device),
torch.randn(
len(self.controlnet), m * batch_size, 3, image_height, image_width, dtype=dtype, device=self.device
),
torch.randn(len(self.controlnet), dtype=dtype, device=self.device),
)
def fp32_input_output_names(self) -> List[str]:
return ["sample", "timestep"]

View file

@ -2,6 +2,7 @@
# Copyright (c) Microsoft Corporation. All rights reserved.
# Licensed under the MIT License.
# --------------------------------------------------------------------------
import hashlib
import os
from enum import Enum
@ -68,18 +69,46 @@ class EngineBuilder:
self.torch_models = {}
self.use_vae_slicing = False
self.torch_sdpa = getattr(torch.nn.functional, "scaled_dot_product_attention", None)
def enable_vae_slicing(self):
self.use_vae_slicing = True
def disable_torch_spda(self):
if hasattr(torch.nn.functional, "scaled_dot_product_attention"):
delattr(torch.nn.functional, "scaled_dot_product_attention")
def enable_torch_spda(self):
if (not hasattr(torch.nn.functional, "scaled_dot_product_attention")) and self.torch_sdpa:
torch.nn.functional.scaled_dot_product_attention = self.torch_sdpa
def teardown(self):
for engine in self.engines.values():
del engine
self.engines = {}
def get_cached_model_name(self, model_name):
hash_source = []
if model_name in ["clip", "clip2", "unet", "unetxl"] and self.pipeline_info.lora_weights:
if self.pipeline_info.lora_weights in [
"latent-consistency/lcm-lora-sdxl",
"latent-consistency/lcm-lora-sdv1-5",
]:
if model_name in ["unet", "unetxl"]:
model_name = model_name + "_lcm-lora"
else:
model_name = model_name + "_lora"
hash_source.append(self.pipeline_info.lora_weights)
# TODO(tianleiwu): save custom model to a directory named by its original model.
if model_name == "unetxl" and self.pipeline_info.custom_unet():
model_name = "lcm_" + model_name
model_name = model_name + "_lcm"
if model_name in ["unet", "unetxl"] and self.pipeline_info.controlnet:
model_name = model_name + "_" + "_".join(self.pipeline_info.controlnet)
if hash_source:
model_name += "_" + hashlib.md5("\t".join(hash_source).encode("utf-8")).digest().hex()[:8]
# TODO: When we support original VAE, we shall save custom VAE to another directory.
@ -87,22 +116,54 @@ class EngineBuilder:
model_name += "_inpaint"
return model_name
def get_onnx_path(self, model_name, onnx_dir, opt=True, suffix=""):
def get_model_dir(self, model_name, root_dir, opt=True, suffix="", create=True):
engine_name = self.engine_type.name.lower()
directory_name = self.get_cached_model_name(model_name) + (f".{engine_name}" if opt else "") + suffix
onnx_model_dir = os.path.join(onnx_dir, directory_name)
os.makedirs(onnx_model_dir, exist_ok=True)
onnx_model_dir = os.path.join(root_dir, directory_name)
if create:
os.makedirs(onnx_model_dir, exist_ok=True)
return onnx_model_dir
def get_onnx_path(self, model_name, onnx_dir, opt=True, suffix=""):
onnx_model_dir = self.get_model_dir(model_name, onnx_dir, opt=opt, suffix=suffix)
return os.path.join(onnx_model_dir, "model.onnx")
def get_engine_path(self, engine_dir, model_name, profile_id):
return os.path.join(engine_dir, self.get_cached_model_name(model_name) + profile_id)
def load_models(self, framework_model_dir: str):
# Disable torch SDPA since torch 2.0.* cannot export it to ONNX
# TODO(tianleiwu): Test and remove it if this is not needed in Torch 2.1.
if hasattr(torch.nn.functional, "scaled_dot_product_attention"):
delattr(torch.nn.functional, "scaled_dot_product_attention")
def load_pipeline_with_lora(self):
"""Load text encoders and UNet with diffusers pipeline"""
from diffusers import DiffusionPipeline
pipeline = DiffusionPipeline.from_pretrained(
self.pipeline_info.name(),
variant="fp16",
torch_dtype=torch.float16,
)
pipeline.load_lora_weights(self.pipeline_info.lora_weights)
pipeline.fuse_lora(lora_scale=self.pipeline_info.lora_scale)
del pipeline.vae
pipeline.vae = None
return pipeline
def get_or_load_model(self, pipeline, model_name, model_obj, framework_model_dir):
if model_name in ["clip", "clip2", "unet", "unetxl"] and pipeline:
if model_name == "clip":
model = pipeline.text_encoder
pipeline.text_encoder = None
elif model_name == "clip2":
model = pipeline.text_encoder_2
pipeline.text_encoder_2 = None
else:
model = pipeline.unet
pipeline.unet = None
else:
model = model_obj.load_model(framework_model_dir, self.hf_token)
return model.to(self.torch_device)
def load_models(self, framework_model_dir: str):
# For TRT or ORT_TRT, we will export fp16 torch model for UNet.
# For ORT_CUDA, we export fp32 model first, then optimize to fp16.
export_fp16_unet = self.engine_type in [EngineType.ORT_TRT, EngineType.TRT]
@ -198,6 +259,7 @@ def get_engine_paths(work_dir: str, pipeline_info: PipelineInfo, engine_type: En
onnx_dir = os.path.join(root_dir, engine_type.name, short_name, "onnx")
engine_dir = os.path.join(root_dir, engine_type.name, short_name, "engine")
output_dir = os.path.join(root_dir, engine_type.name, short_name, "output")
timing_cache = os.path.join(root_dir, engine_type.name, "timing_cache")
framework_model_dir = os.path.join(root_dir, engine_type.name, "torch_model")

View file

@ -158,6 +158,7 @@ class OrtCudaEngineBuilder(EngineBuilder):
engine_dir: str,
framework_model_dir: str,
onnx_dir: str,
tmp_dir: Optional[str] = None,
onnx_opset_version: int = 17,
force_engine_rebuild: bool = False,
device_id: int = 0,
@ -187,22 +188,39 @@ class OrtCudaEngineBuilder(EngineBuilder):
if model_name not in self.model_config:
self.model_config[model_name] = _ModelConfig(onnx_opset_version, self.use_cuda_graph)
# Load lora only when we need export text encoder or UNet to ONNX.
load_lora = False
if self.pipeline_info.lora_weights:
for model_name in self.models:
if model_name not in ["clip", "clip2", "unet", "unetxl"]:
continue
onnx_path = self.get_onnx_path(model_name, onnx_dir, opt=False)
suffix = ".fp16" if self.model_config[model_name].fp16 else ".fp32"
onnx_opt_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=suffix)
if not os.path.exists(onnx_opt_path):
if not os.path.exists(onnx_path):
load_lora = True
break
# Export models to ONNX
self.disable_torch_spda()
pipe = self.load_pipeline_with_lora() if load_lora else None
for model_name, model_obj in self.models.items():
if model_name == "vae" and self.vae_torch_fallback:
continue
onnx_path = self.get_onnx_path(model_name, onnx_dir, opt=False)
onnx_fp32_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp32")
onnx_fp16_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp16")
onnx_opt_path = onnx_fp16_path if self.model_config[model_name].fp16 else onnx_fp32_path
suffix = ".fp16" if self.model_config[model_name].fp16 else ".fp32"
onnx_opt_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=suffix)
if not os.path.exists(onnx_opt_path):
if not os.path.exists(onnx_path):
print("----")
logger.info("Exporting model: %s", onnx_path)
model = model_obj.load_model(framework_model_dir, self.hf_token)
if model_name == "vae":
model.to(torch.float32)
model = self.get_or_load_model(pipe, model_name, model_obj, framework_model_dir)
model = model.to(torch.float32)
with torch.inference_mode():
# For CUDA EP, export FP32 onnx since some graph fusion only supports fp32 graph pattern.
@ -230,18 +248,19 @@ class OrtCudaEngineBuilder(EngineBuilder):
# If final target is fp16 model, we save fp32 optimized model so that it is easy to tune
# fp16 conversion. That could save a lot of time in developing.
use_fp32_intermediate = save_fp32_intermediate_model and self.model_config[model_name].fp16
onnx_fp32_path = onnx_path
if use_fp32_intermediate:
onnx_fp32_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp32")
if not os.path.exists(onnx_fp32_path):
print("------")
logger.info("Generating optimized model: %s", onnx_fp32_path)
# There is risk that some ORT fused ops fp32 only. So far, we have not encountered such issue.
model_obj.optimize_ort(
onnx_path,
onnx_fp32_path,
to_fp16=False,
fp32_op_list=self.model_config[model_name].force_fp32_ops,
optimize_by_ort=self.model_config[model_name].optimize_by_ort,
tmp_dir=self.get_model_dir(model_name, tmp_dir, opt=False, suffix=".fp32", create=False),
)
else:
logger.info("Found cached optimized model: %s", onnx_fp32_path)
@ -255,24 +274,25 @@ class OrtCudaEngineBuilder(EngineBuilder):
optimize_by_ort = False if use_fp32_intermediate else self.model_config[model_name].optimize_by_ort
model_obj.optimize_ort(
onnx_fp32_path if use_fp32_intermediate else onnx_path,
onnx_fp32_path,
onnx_opt_path,
to_fp16=self.model_config[model_name].fp16,
fp32_op_list=self.model_config[model_name].force_fp32_ops,
optimize_by_ort=optimize_by_ort,
optimize_by_fusion=not use_fp32_intermediate,
tmp_dir=self.get_model_dir(model_name, tmp_dir, opt=False, suffix=".fp16", create=False),
)
else:
logger.info("Found cached optimized model: %s", onnx_opt_path)
self.enable_torch_spda()
built_engines = {}
for model_name in self.models:
if model_name == "vae" and self.vae_torch_fallback:
continue
onnx_fp32_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp32")
onnx_fp16_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=".fp16")
onnx_opt_path = onnx_fp16_path if self.model_config[model_name].fp16 else onnx_fp32_path
suffix = ".fp16" if self.model_config[model_name].fp16 else ".fp32"
onnx_opt_path = self.get_onnx_path(model_name, engine_dir, opt=True, suffix=suffix)
use_cuda_graph = self.model_config[model_name].use_cuda_graph

View file

@ -189,7 +189,28 @@ class OrtTensorrtEngineBuilder(EngineBuilder):
if not os.path.isdir(onnx_dir):
os.makedirs(onnx_dir)
# Load lora only when we need export text encoder or UNet to ONNX.
load_lora = False
if self.pipeline_info.lora_weights:
for model_name, model_obj in self.models.items():
if model_name not in ["clip", "clip2", "unet", "unetxl"]:
continue
profile_id = model_obj.get_profile_id(
opt_batch_size, opt_image_height, opt_image_width, static_batch, static_image_shape
)
engine_path = self.get_engine_path(engine_dir, model_name, profile_id)
if not self.has_engine_file(engine_path):
onnx_path = self.get_onnx_path(model_name, onnx_dir, opt=False)
onnx_opt_path = self.get_onnx_path(model_name, onnx_dir, opt=True)
if not os.path.exists(onnx_opt_path):
if not os.path.exists(onnx_path):
load_lora = True
break
# Export models to ONNX
self.disable_torch_spda()
pipe = self.load_pipeline_with_lora() if load_lora else None
for model_name, model_obj in self.models.items():
if model_name == "vae" and self.vae_torch_fallback:
continue
@ -204,7 +225,8 @@ class OrtTensorrtEngineBuilder(EngineBuilder):
if not os.path.exists(onnx_opt_path):
if not os.path.exists(onnx_path):
logger.info(f"Exporting model: {onnx_path}")
model = model_obj.load_model(framework_model_dir, self.hf_token)
model = self.get_or_load_model(pipe, model_name, model_obj, framework_model_dir)
with torch.inference_mode(), torch.autocast("cuda"):
inputs = model_obj.get_sample_input(opt_batch_size, opt_image_height, opt_image_width)
torch.onnx.export(
@ -230,6 +252,7 @@ class OrtTensorrtEngineBuilder(EngineBuilder):
model_obj.optimize_trt(onnx_path, onnx_opt_path)
else:
logger.info("Found cached optimized model: %s", onnx_opt_path)
self.enable_torch_spda()
built_engines = {}
for model_name, model_obj in self.models.items():

View file

@ -407,11 +407,32 @@ class TensorrtEngineBuilder(EngineBuilder):
self.load_models(framework_model_dir)
# Load lora only when we need export text encoder or UNet to ONNX.
load_lora = False
if self.pipeline_info.lora_weights:
for model_name, model_obj in self.models.items():
if model_name not in ["clip", "clip2", "unet", "unetxl"]:
continue
profile_id = model_obj.get_profile_id(
opt_batch_size, opt_image_height, opt_image_width, static_batch, static_shape
)
engine_path = self.get_engine_path(engine_dir, model_name, profile_id)
if force_export or force_build or not os.path.exists(engine_path):
onnx_path = self.get_onnx_path(model_name, onnx_dir, opt=False)
onnx_opt_path = self.get_onnx_path(model_name, onnx_dir, opt=True)
if force_export or not os.path.exists(onnx_opt_path):
if force_export or not os.path.exists(onnx_path):
load_lora = True
break
# Export models to ONNX
for model_name, obj in self.models.items():
self.disable_torch_spda()
pipe = self.load_pipeline_with_lora() if load_lora else None
for model_name, model_obj in self.models.items():
if model_name == "vae" and self.vae_torch_fallback:
continue
profile_id = obj.get_profile_id(
profile_id = model_obj.get_profile_id(
opt_batch_size, opt_image_height, opt_image_width, static_batch, static_shape
)
engine_path = self.get_engine_path(engine_dir, model_name, profile_id)
@ -421,9 +442,10 @@ class TensorrtEngineBuilder(EngineBuilder):
if force_export or not os.path.exists(onnx_opt_path):
if force_export or not os.path.exists(onnx_path):
print(f"Exporting model: {onnx_path}")
model = obj.load_model(framework_model_dir, self.hf_token)
model = self.get_or_load_model(pipe, model_name, model_obj, framework_model_dir)
with torch.inference_mode(), torch.autocast("cuda"):
inputs = obj.get_sample_input(1, opt_image_height, opt_image_width)
inputs = model_obj.get_sample_input(1, opt_image_height, opt_image_width)
torch.onnx.export(
model,
inputs,
@ -431,9 +453,9 @@ class TensorrtEngineBuilder(EngineBuilder):
export_params=True,
opset_version=onnx_opset,
do_constant_folding=True,
input_names=obj.get_input_names(),
output_names=obj.get_output_names(),
dynamic_axes=obj.get_dynamic_axes(),
input_names=model_obj.get_input_names(),
output_names=model_obj.get_output_names(),
dynamic_axes=model_obj.get_dynamic_axes(),
)
del model
torch.cuda.empty_cache()
@ -444,15 +466,16 @@ class TensorrtEngineBuilder(EngineBuilder):
# Optimize onnx
if force_optimize or not os.path.exists(onnx_opt_path):
print(f"Generating optimizing model: {onnx_opt_path}")
obj.optimize_trt(onnx_path, onnx_opt_path)
model_obj.optimize_trt(onnx_path, onnx_opt_path)
else:
print(f"Found cached optimized model: {onnx_opt_path} ")
self.enable_torch_spda()
# Build TensorRT engines
for model_name, obj in self.models.items():
for model_name, model_obj in self.models.items():
if model_name == "vae" and self.vae_torch_fallback:
continue
profile_id = obj.get_profile_id(
profile_id = model_obj.get_profile_id(
opt_batch_size, opt_image_height, opt_image_width, static_batch, static_shape
)
engine_path = self.get_engine_path(engine_dir, model_name, profile_id)
@ -463,7 +486,7 @@ class TensorrtEngineBuilder(EngineBuilder):
engine.build(
onnx_opt_path,
fp16=True,
input_profile=obj.get_input_profile(
input_profile=model_obj.get_input_profile(
opt_batch_size,
opt_image_height,
opt_image_width,

View file

@ -8,6 +8,8 @@ ONNX Model Optimizer for Stable Diffusion
"""
import logging
import os
import shutil
import tempfile
from pathlib import Path
@ -33,23 +35,32 @@ class OrtStableDiffusionOptimizer:
"clip": ClipOnnxModel,
}
def optimize_by_ort(self, onnx_model, use_external_data_format=False):
def _optimize_by_ort(self, onnx_model, use_external_data_format, tmp_dir):
# Save to a temporary file so that we can load it with Onnx Runtime.
logger.info("Saving a temporary model to run OnnxRuntime graph optimizations...")
tmp_model_path = Path(tmp_dir) / "model.onnx"
onnx_model.save_model_to_file(str(tmp_model_path), use_external_data_format=use_external_data_format)
ort_optimized_model_path = Path(tmp_dir) / "optimized.onnx"
optimize_by_onnxruntime(
str(tmp_model_path),
use_gpu=True,
optimized_model_path=str(ort_optimized_model_path),
save_as_external_data=use_external_data_format,
external_data_filename="optimized.onnx_data",
)
model = onnx.load(str(ort_optimized_model_path), load_external_data=True)
return self.model_type_class_mapping[self.model_type](model)
def optimize_by_ort(self, onnx_model, use_external_data_format=False, tmp_dir=None):
# Use this step to see the final graph that executed by Onnx Runtime.
with tempfile.TemporaryDirectory() as tmp_dir:
# Save to a temporary file so that we can load it with Onnx Runtime.
logger.info("Saving a temporary model to run OnnxRuntime graph optimizations...")
tmp_model_path = Path(tmp_dir) / "model.onnx"
onnx_model.save_model_to_file(str(tmp_model_path), use_external_data_format=use_external_data_format)
ort_optimized_model_path = Path(tmp_dir) / "optimized.onnx"
optimize_by_onnxruntime(
str(tmp_model_path),
use_gpu=True,
optimized_model_path=str(ort_optimized_model_path),
save_as_external_data=use_external_data_format,
external_data_filename="optimized.onnx_data",
)
model = onnx.load(str(ort_optimized_model_path), load_external_data=True)
return self.model_type_class_mapping[self.model_type](model)
if tmp_dir is None:
with tempfile.TemporaryDirectory() as temp_dir:
return self._optimize_by_ort(onnx_model, use_external_data_format, temp_dir)
else:
os.makedirs(tmp_dir, exist_ok=True)
model = self._optimize_by_ort(onnx_model, use_external_data_format, tmp_dir)
shutil.rmtree(tmp_dir)
return model
def optimize(
self,
@ -62,6 +73,7 @@ class OrtStableDiffusionOptimizer:
optimize_by_ort=True,
optimize_by_fusion=True,
final_target_float16=True,
tmp_dir=None,
):
"""Optimize onnx model using ONNX Runtime transformers optimizer"""
logger.info(f"Optimize {input_fp32_onnx_path}...")
@ -104,7 +116,7 @@ class OrtStableDiffusionOptimizer:
from onnxruntime import __version__ as ort_version
if optimize_by_ort and (version.parse(ort_version) >= version.parse("1.16.0") or not use_external_data_format):
m = self.optimize_by_ort(m, use_external_data_format=use_external_data_format)
m = self.optimize_by_ort(m, use_external_data_format=use_external_data_format, tmp_dir=tmp_dir)
if float16:
logger.info("Convert to float16 ...")

View file

@ -25,6 +25,7 @@ import pathlib
import random
from typing import Any, Dict, List
import numpy as np
import nvtx
import torch
from cuda import cudart
@ -103,8 +104,6 @@ class StableDiffusionPipeline:
self.verbose = verbose
self.nvtx_profile = nvtx_profile
self.stages = pipeline_info.stages()
self.use_cuda_graph = use_cuda_graph
self.tokenizer = None
@ -138,11 +137,20 @@ class StableDiffusionPipeline:
self.pipeline_info, self.framework_model_dir, self.hf_token, subfolder="tokenizer_2"
)
self.control_image_processor = None
if self.pipeline_info.is_xl() and self.pipeline_info.controlnet:
from diffusers.image_processor import VaeImageProcessor
self.control_image_processor = VaeImageProcessor(
vae_scale_factor=8, do_convert_rgb=True, do_normalize=False
)
# Create CUDA events
self.events = {}
for stage in ["clip", "denoise", "vae", "vae_encoder"]:
for marker in ["start", "stop"]:
self.events[stage + "-" + marker] = cudart.cudaEventCreate()[1]
self.markers = {}
def is_backend_tensorrt(self):
return self.engine_type == EngineType.TRT
@ -219,19 +227,63 @@ class StableDiffusionPipeline:
timesteps = self.scheduler.timesteps[t_start:].to(self.device)
return timesteps, t_start
def preprocess_images(self, batch_size, images=()):
def start_profile(self, name, color="blue"):
if self.nvtx_profile:
nvtx_image_preprocess = nvtx.start_range(message="image_preprocess", color="pink")
self.markers[name] = nvtx.start_range(message=name, color=color)
event_name = name + "-start"
if event_name in self.events:
cudart.cudaEventRecord(self.events[event_name], 0)
def stop_profile(self, name):
event_name = name + "-stop"
if event_name in self.events:
cudart.cudaEventRecord(self.events[event_name], 0)
if self.nvtx_profile:
nvtx.end_range(self.markers[name])
def preprocess_images(self, batch_size, images=()):
self.start_profile("preprocess", color="pink")
init_images = []
for i in images:
image = i.to(self.device).float()
if image.shape[0] != batch_size:
image = image.repeat(batch_size, 1, 1, 1)
init_images.append(image)
if self.nvtx_profile:
nvtx.end_range(nvtx_image_preprocess)
self.stop_profile("preprocess")
return tuple(init_images)
def preprocess_controlnet_images(
self, batch_size, images=None, do_classifier_free_guidance=True, height=1024, width=1024
):
"""
Process a list of PIL.Image.Image as control images, and return a torch tensor.
"""
if images is None:
return None
self.start_profile("preprocess", color="pink")
if not self.pipeline_info.is_xl():
images = [
(np.array(i.convert("RGB")).astype(np.float32) / 255.0)[..., None]
.transpose(3, 2, 0, 1)
.repeat(batch_size, axis=0)
for i in images
]
if do_classifier_free_guidance:
images = [torch.cat([torch.from_numpy(i).to(self.device).float()] * 2) for i in images]
else:
images = [torch.from_numpy(i).to(self.device).float() for i in images]
images = torch.cat([image[None, ...] for image in images], dim=0)
images = images.to(dtype=torch.float16)
else:
images = self.control_image_processor.preprocess(images, height=height, width=width).to(dtype=torch.float32)
images = images.repeat_interleave(batch_size, dim=0)
images = images.to(device=self.device, dtype=torch.float16)
if do_classifier_free_guidance:
images = torch.cat([images] * 2)
self.stop_profile("preprocess")
return images
def encode_prompt(
self,
prompt,
@ -246,9 +298,7 @@ class StableDiffusionPipeline:
if tokenizer is None:
tokenizer = self.tokenizer
if self.nvtx_profile:
nvtx_clip = nvtx.start_range(message="clip", color="green")
cudart.cudaEventRecord(self.events["clip-start"], 0)
self.start_profile("clip", color="green")
# Tokenize prompt
text_input_ids = (
@ -308,9 +358,7 @@ class StableDiffusionPipeline:
else:
text_embeddings = hidden_states.to(dtype=torch.float16)
cudart.cudaEventRecord(self.events["clip-stop"], 0)
if self.nvtx_profile:
nvtx.end_range(nvtx_clip)
self.stop_profile("clip")
if pooled_outputs:
return text_embeddings, pooled_output
@ -330,14 +378,12 @@ class StableDiffusionPipeline:
):
do_classifier_free_guidance = guidance > 1.0
cudart.cudaEventRecord(self.events["denoise-start"], 0)
self.start_profile("denoise", color="blue")
if not isinstance(timesteps, torch.Tensor):
timesteps = self.scheduler.timesteps
for step_index, timestep in enumerate(timesteps):
if self.nvtx_profile:
nvtx_latent_scale = nvtx.start_range(message="latent_scale", color="pink")
# Expand the latents if we are doing classifier free guidance
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
@ -347,8 +393,6 @@ class StableDiffusionPipeline:
if isinstance(mask, torch.Tensor):
latent_model_input = torch.cat([latent_model_input, mask, masked_image_latents], dim=1)
if self.nvtx_profile:
nvtx.end_range(nvtx_latent_scale)
# Predict the noise residual
if self.nvtx_profile:
@ -361,6 +405,7 @@ class StableDiffusionPipeline:
"timestep": timestep_float,
"encoder_hidden_states": text_embeddings,
}
if add_kwargs:
params.update(add_kwargs)
@ -369,9 +414,6 @@ class StableDiffusionPipeline:
if self.nvtx_profile:
nvtx.end_range(nvtx_unet)
if self.nvtx_profile:
nvtx_latent_step = nvtx.start_range(message="latent_step", color="pink")
# perform guidance
if do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
@ -384,36 +426,23 @@ class StableDiffusionPipeline:
else:
latents = self.scheduler.step(noise_pred, latents, step_offset + step_index, timestep)
if self.nvtx_profile:
nvtx.end_range(nvtx_latent_step)
cudart.cudaEventRecord(self.events["denoise-stop"], 0)
# The actual number of steps. It might be different from denoising_steps.
self.actual_steps = len(timesteps)
self.stop_profile("denoise")
return latents
def encode_image(self, init_image):
if self.nvtx_profile:
nvtx_vae = nvtx.start_range(message="vae_encoder", color="red")
cudart.cudaEventRecord(self.events["vae_encoder-start"], 0)
self.start_profile("vae_encoder", color="red")
init_latents = self.run_engine("vae_encoder", {"images": init_image})["latent"]
cudart.cudaEventRecord(self.events["vae_encoder-stop"], 0)
if self.nvtx_profile:
nvtx.end_range(nvtx_vae)
init_latents = self.vae_scaling_factor * init_latents
self.stop_profile("vae_encoder")
return init_latents
def decode_latent(self, latents):
if self.nvtx_profile:
nvtx_vae = nvtx.start_range(message="vae", color="red")
cudart.cudaEventRecord(self.events["vae-start"], 0)
self.start_profile("vae", color="red")
images = self.backend.vae_decode(latents)
cudart.cudaEventRecord(self.events["vae-stop"], 0)
if self.nvtx_profile:
nvtx.end_range(nvtx_vae)
self.stop_profile("vae")
return images
def print_summary(self, tic, toc, batch_size, vae_enc=False) -> Dict[str, Any]:
@ -428,18 +457,23 @@ class StableDiffusionPipeline:
)
latency = (toc - tic) * 1000.0
print("|------------|--------------|")
print("| {:^10} | {:^12} |".format("Module", "Latency"))
print("|------------|--------------|")
print("|----------------|--------------|")
print("| {:^14} | {:^12} |".format("Module", "Latency"))
print("|----------------|--------------|")
if vae_enc:
print("| {:^10} | {:>9.2f} ms |".format("VAE-Enc", latency_vae_encoder))
print("| {:^10} | {:>9.2f} ms |".format("CLIP", latency_clip))
print("| {:^10} | {:>9.2f} ms |".format("UNet x " + str(self.actual_steps), latency_unet))
print("| {:^10} | {:>9.2f} ms |".format("VAE-Dec", latency_vae))
print("| {:^14} | {:>9.2f} ms |".format("VAE-Enc", latency_vae_encoder))
print("| {:^14} | {:>9.2f} ms |".format("CLIP", latency_clip))
print(
"| {:^14} | {:>9.2f} ms |".format(
"UNet" + ("+CNet" if self.pipeline_info.controlnet else "") + " x " + str(self.actual_steps),
latency_unet,
)
)
print("| {:^14} | {:>9.2f} ms |".format("VAE-Dec", latency_vae))
print("|------------|--------------|")
print("| {:^10} | {:>9.2f} ms |".format("Pipeline", latency))
print("|------------|--------------|")
print("|----------------|--------------|")
print("| {:^14} | {:>9.2f} ms |".format("Pipeline", latency))
print("|----------------|--------------|")
print(f"Throughput: {throughput:.2f} image/s")
perf_data = {

View file

@ -51,6 +51,8 @@ class Txt2ImgPipeline(StableDiffusionPipeline):
denoising_steps=50,
guidance=7.5,
seed=None,
controlnet_images=None,
controlnet_scales=None,
warmup=False,
return_type="latent",
):
@ -73,10 +75,25 @@ class Txt2ImgPipeline(StableDiffusionPipeline):
e2e_tic = time.perf_counter()
# CLIP text encoder
text_embeddings = self.encode_prompt(prompt, negative_prompt)
do_classifier_free_guidance = guidance > 1.0
text_embeddings = self.encode_prompt(
prompt,
negative_prompt,
do_classifier_free_guidance=do_classifier_free_guidance,
)
add_kwargs = None
if self.pipeline_info.controlnet:
controlnet_images = self.preprocess_controlnet_images(
latents.shape[0], controlnet_images, do_classifier_free_guidance=do_classifier_free_guidance
)
add_kwargs = {
"controlnet_images": controlnet_images,
"controlnet_scales": controlnet_scales.to(controlnet_images.dtype).to(controlnet_images.device),
}
# UNet denoiser
latents = self.denoise_latent(latents, text_embeddings, guidance=guidance)
latents = self.denoise_latent(latents, text_embeddings, guidance=guidance, add_kwargs=add_kwargs)
# VAE decode latent
images = self.decode_latent(latents / self.vae_scaling_factor)
@ -99,6 +116,8 @@ class Txt2ImgPipeline(StableDiffusionPipeline):
denoising_steps=30,
guidance=7.5,
seed=None,
controlnet_images=None,
controlnet_scales=None,
warmup=False,
return_type="image",
):
@ -138,6 +157,8 @@ class Txt2ImgPipeline(StableDiffusionPipeline):
denoising_steps=denoising_steps,
guidance=guidance,
seed=seed,
controlnet_images=controlnet_images,
controlnet_scales=controlnet_scales,
warmup=warmup,
return_type=return_type,
)
@ -150,6 +171,8 @@ class Txt2ImgPipeline(StableDiffusionPipeline):
denoising_steps=denoising_steps,
guidance=guidance,
seed=seed,
controlnet_images=controlnet_images,
controlnet_scales=controlnet_scales,
warmup=warmup,
return_type=return_type,
)

View file

@ -58,6 +58,8 @@ class Txt2ImgXLPipeline(StableDiffusionPipeline):
denoising_steps=30,
guidance=5.0,
seed=None,
controlnet_images=None,
controlnet_scales=None,
warmup=False,
return_type="image",
):
@ -117,6 +119,20 @@ class Txt2ImgXLPipeline(StableDiffusionPipeline):
add_time_ids = torch.cat([add_time_ids, add_time_ids], dim=0)
add_kwargs = {"text_embeds": pooled_embeddings2, "time_ids": add_time_ids.to(self.device)}
if self.pipeline_info.controlnet:
controlnet_images = self.preprocess_controlnet_images(
latents.shape[0],
controlnet_images,
do_classifier_free_guidance=do_classifier_free_guidance,
height=image_height,
width=image_width,
)
add_kwargs.update(
{
"controlnet_images": controlnet_images,
"controlnet_scales": controlnet_scales.to(controlnet_images.dtype).to(controlnet_images.device),
}
)
# UNet denoiser
latents = self.denoise_latent(
@ -152,6 +168,8 @@ class Txt2ImgXLPipeline(StableDiffusionPipeline):
denoising_steps=30,
guidance=5.0,
seed=None,
controlnet_images=None,
controlnet_scales=None,
warmup=False,
return_type="image",
):
@ -192,6 +210,8 @@ class Txt2ImgXLPipeline(StableDiffusionPipeline):
denoising_steps=denoising_steps,
guidance=guidance,
seed=seed,
controlnet_images=controlnet_images,
controlnet_scales=controlnet_scales,
warmup=warmup,
return_type=return_type,
)
@ -204,6 +224,8 @@ class Txt2ImgXLPipeline(StableDiffusionPipeline):
denoising_steps=denoising_steps,
guidance=guidance,
seed=seed,
controlnet_images=controlnet_images,
controlnet_scales=controlnet_scales,
warmup=warmup,
return_type=return_type,
)

View file

@ -9,6 +9,7 @@ packaging
protobuf==3.20.3
psutil
sympy
controlnet_aux
# The following are for SDXL
optimum==1.13.1
safetensors