Flux needs to be bf16.

This commit is contained in:
comfyanonymous
2024-08-04 16:40:52 -04:00
parent c7faa24f78
commit 76cdd93680
3 changed files with 10 additions and 3 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.5"
version = "0.1.6"
license = "LICENSE"
dependencies = [
"tensorrt>=10.0.1",
+8 -2
View File
@@ -163,6 +163,7 @@ class TRT_MODEL_CONVERSION_BASE:
context_len_min = context_len
y_dim = model.model.adm_channels
extra_input = {}
dtype = torch.float16
if isinstance(model.model, comfy.model_base.SD3): #SD3
context_embedder_config = model.model.model_config.unet_config.get("context_embedder_config", None)
@@ -179,6 +180,7 @@ class TRT_MODEL_CONVERSION_BASE:
context_len = 256
y_dim = model.model.model_config.unet_config.get("vec_in_dim", None)
extra_input = {"guidance": ()}
dtype = torch.bfloat16
if context_dim is not None:
input_names = ["x", "timesteps", "context"]
@@ -268,7 +270,7 @@ class TRT_MODEL_CONVERSION_BASE:
torch.zeros(
shape,
device=comfy.model_management.get_torch_device(),
dtype=torch.float16,
dtype=dtype,
),
)
@@ -325,7 +327,11 @@ class TRT_MODEL_CONVERSION_BASE:
input_names[k], encode(min_shape), encode(opt_shape), encode(max_shape)
)
config.set_flag(trt.BuilderFlag.FP16)
if dtype == torch.float16:
config.set_flag(trt.BuilderFlag.FP16)
if dtype == torch.bfloat16:
config.set_flag(trt.BuilderFlag.BF16)
config.add_optimization_profile(profile)
if is_static:
+1
View File
@@ -158,6 +158,7 @@ class TensorRTLoader:
conf = comfy.supported_models.Flux({})
conf.unet_config["disable_unet_model_creation"] = True
model = conf.get_model({})
unet.dtype = torch.bfloat16 #TODO: autodetect
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