From c7faa24f78445a3f74f7021f9808623e4c022514 Mon Sep 17 00:00:00 2001 From: comfyanonymous Date: Sun, 4 Aug 2024 16:16:39 -0400 Subject: [PATCH] Add untested Flux support. --- pyproject.toml | 2 +- tensorrt_convert.py | 36 +++++++++++++++++++++++------------- tensorrt_loader.py | 10 +++++++++- 3 files changed, 33 insertions(+), 15 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 79d7276..7dcc0bb 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.4" +version = "0.1.5" license = "LICENSE" dependencies = [ "tensorrt>=10.0.1", diff --git a/tensorrt_convert.py b/tensorrt_convert.py index e7e3e92..513b348 100644 --- a/tensorrt_convert.py +++ b/tensorrt_convert.py @@ -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 += ( diff --git a/tensorrt_loader.py b/tensorrt_loader.py index 8ceff29..dce734e 100644 --- a/tensorrt_loader.py +++ b/tensorrt_loader.py @@ -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