SD3 support.

This commit is contained in:
comfyanonymous
2024-06-12 03:05:15 -04:00
parent 750450906c
commit 975ed7a603
3 changed files with 20 additions and 9 deletions
+1 -1
View File
@@ -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",
+13 -6
View File
@@ -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),
+6 -2
View File
@@ -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