Add untested Flux support.

This commit is contained in:
comfyanonymous
2024-08-04 16:16:59 -04:00
parent ef2d9743c2
commit c7faa24f78
3 changed files with 33 additions and 15 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.4"
version = "0.1.5"
license = "LICENSE"
dependencies = [
"tensorrt>=10.0.1",
+23 -13
View File
@@ -161,6 +161,8 @@ class TRT_MODEL_CONVERSION_BASE:
context_dim = model.model.model_config.unet_config.get("context_dim", None)
context_len = 77
context_len_min = context_len
y_dim = model.model.adm_channels
extra_input = {}
if isinstance(model.model, comfy.model_base.SD3): #SD3
context_embedder_config = model.model.model_config.unet_config.get("context_embedder_config", None)
@@ -171,6 +173,12 @@ class TRT_MODEL_CONVERSION_BASE:
context_dim = 2048
context_len_min = 256
context_len = 256
elif isinstance(model.model, comfy.model_base.Flux):
context_dim = model.model.model_config.unet_config.get("context_in_dim", None)
context_len_min = 256
context_len = 256
y_dim = model.model.model_config.unet_config.get("vec_in_dim", None)
extra_input = {"guidance": ()}
if context_dim is not None:
input_names = ["x", "timesteps", "context"]
@@ -209,17 +217,13 @@ class TRT_MODEL_CONVERSION_BASE:
context_len_min = context_len = 1
else:
class UNET(torch.nn.Module):
def forward(self, x, timesteps, context, y=None):
if y is None:
return self.unet(x, timesteps, context, transformer_options=self.transformer_options)
else:
return self.unet(
x,
timesteps,
context,
y,
transformer_options=self.transformer_options,
)
def forward(self, x, timesteps, context, *args):
extras = input_names[3:]
extra_args = {}
for i in range(len(extras)):
extra_args[extras[i]] = args[i]
return self.unet(x, timesteps, context, transformer_options=self.transformer_options, **extra_args)
_unet = UNET()
_unet.unet = unet
_unet.transformer_options = transformer_options
@@ -243,8 +247,6 @@ class TRT_MODEL_CONVERSION_BASE:
(batch_size_max, context_len * context_max, context_dim),
)
y_dim = model.model.adm_channels
if y_dim > 0:
input_names.append("y")
dynamic_axes["y"] = {0: "batch"}
@@ -252,6 +254,14 @@ class TRT_MODEL_CONVERSION_BASE:
inputs_shapes_opt += ((batch_size_opt, y_dim),)
inputs_shapes_max += ((batch_size_max, y_dim),)
for k in extra_input:
input_names.append(k)
dynamic_axes[k] = {0: "batch"}
inputs_shapes_min += ((batch_size_min,) + extra_input[k],)
inputs_shapes_opt += ((batch_size_opt,) + extra_input[k],)
inputs_shapes_max += ((batch_size_max,) + extra_input[k],)
inputs = ()
for shape in inputs_shapes_opt:
inputs += (
+9 -1
View File
@@ -54,6 +54,10 @@ class TrTUnet:
if y is not None:
model_inputs["y"] = y
for i in range(len(model_inputs), self.engine.num_io_tensors - 1):
name = self.engine.get_tensor_name(i)
model_inputs[name] = kwargs[name]
batch_size = x.shape[0]
dims = self.engine.get_tensor_profile_shape(self.engine.get_tensor_name(0), 0)
min_batch = dims[0][0]
@@ -110,7 +114,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", "sd3", "auraflow"], ),
"model_type": (["sdxl_base", "sdxl_refiner", "sd1.x", "sd2.x-768v", "svd", "sd3", "auraflow", "flux"], ),
}}
RETURN_TYPES = ("MODEL",)
FUNCTION = "load_unet"
@@ -150,6 +154,10 @@ class TensorRTLoader:
conf = comfy.supported_models.AuraFlow({})
conf.unet_config["disable_unet_model_creation"] = True
model = conf.get_model({})
elif model_type == "flux":
conf = comfy.supported_models.Flux({})
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