From 975ed7a60321730eb494749e56e1d14716137574 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Wed, 12 Jun 2024 02:51:01 -0400 Subject: [PATCH] SD3 support. --- pyproject.toml | 2 +- tensorrt_convert.py | 19 +++++++++++++------ tensorrt_loader.py | 8 ++++++-- 3 files changed, 20 insertions(+), 9 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index ffb1227..0d27a6d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui_tensorrt" description = "TensorRT Node for ComfyUI\nThis node enables the best performance on NVIDIA RTX™ Graphics Cards (GPUs) for Stable Diffusion by leveraging NVIDIA TensorRT." -version = "0.1.0" +version = "0.1.1" license = "LICENSE" dependencies = [ "tensorrt>=10.0.1", diff --git a/tensorrt_convert.py b/tensorrt_convert.py index 35e73a6..ef694ee 100644 --- a/tensorrt_convert.py +++ b/tensorrt_convert.py @@ -159,7 +159,17 @@ class TRT_MODEL_CONVERSION_BASE: comfy.model_management.load_models_gpu([model], force_patch_weights=True) unet = model.model.diffusion_model - if "context_dim" in model.model.model_config.unet_config: + context_dim = model.model.model_config.unet_config.get("context_dim", None) + context_len = 77 + context_len_min = context_len + + if context_dim is None: #SD3 + context_embedder_config = model.model.model_config.unet_config.get("context_embedder_config", None) + if context_embedder_config is not None: + context_dim = context_embedder_config.get("params", {}).get("in_features", None) + context_len = 154 #NOTE: SD3 can have 77 or 154 depending on which text encoders are used, this is why context_len_min stays 77 + + if context_dim is not None: input_names = ["x", "timesteps", "context"] output_names = ["h"] @@ -170,7 +180,6 @@ class TRT_MODEL_CONVERSION_BASE: } transformer_options = model.model_options['transformer_options'].copy() - context_len = 77 if model.model.model_config.unet_config.get( "use_temporal_resblock", False ): # SVD @@ -194,7 +203,7 @@ class TRT_MODEL_CONVERSION_BASE: svd_unet.unet = unet svd_unet.transformer_options = transformer_options unet = svd_unet - context_len = 1 + context_len_min = context_len = 1 else: class UNET(torch.nn.Module): def forward(self, x, timesteps, context, y=None): @@ -212,12 +221,10 @@ class TRT_MODEL_CONVERSION_BASE: input_channels = model.model.model_config.unet_config.get("in_channels") - context_dim = model.model.model_config.unet_config.get("context_dim") - inputs_shapes_min = ( (batch_size_min, input_channels, height_min // 8, width_min // 8), (batch_size_min,), - (batch_size_min, context_len * context_min, context_dim), + (batch_size_min, context_len_min * context_min, context_dim), ) inputs_shapes_opt = ( (batch_size_opt, input_channels, height_opt // 8, width_opt // 8), diff --git a/tensorrt_loader.py b/tensorrt_loader.py index a60ba7f..3739042 100644 --- a/tensorrt_loader.py +++ b/tensorrt_loader.py @@ -110,7 +110,7 @@ class TensorRTLoader: @classmethod def INPUT_TYPES(s): return {"required": {"unet_name": (folder_paths.get_filename_list("tensorrt"), ), - "model_type": (["sdxl_base", "sdxl_refiner", "sd1.x", "sd2.x-768v", "svd"], ), + "model_type": (["sdxl_base", "sdxl_refiner", "sd1.x", "sd2.x-768v", "svd", "sd3"], ), }} RETURN_TYPES = ("MODEL",) FUNCTION = "load_unet" @@ -141,7 +141,11 @@ class TensorRTLoader: elif model_type == "svd": conf = comfy.supported_models.SVD_img2vid({}) conf.unet_config["disable_unet_model_creation"] = True - model = conf.get_model({}) + model = conf.get_model({}) + elif model_type == "sd3": + conf = comfy.supported_models.SD3({}) + conf.unet_config["disable_unet_model_creation"] = True + model = conf.get_model({}) model.diffusion_model = unet model.memory_required = lambda *args, **kwargs: 0 #always pass inputs batched up as much as possible, our TRT code will handle batch splitting