From 16efd4c687b4126d14cc65026f0f4256aa69a9ce Mon Sep 17 00:00:00 2001 From: matt3o Date: Thu, 19 Oct 2023 16:00:39 +0200 Subject: [PATCH] add model compile --- essentials.py | 26 +++++++++++++++++++++++++- 1 file changed, 25 insertions(+), 1 deletion(-) diff --git a/essentials.py b/essentials.py index ffa34ed..100aa2d 100644 --- a/essentials.py +++ b/essentials.py @@ -39,7 +39,7 @@ class AnyType(str): return False any = AnyType("*") -EPSILON = 1e-7 +EPSILON = 1e-5 class GetImageSize: @classmethod @@ -403,6 +403,26 @@ class SimpleMath: return (round(result), result, ) +class ModelCompile(SaveImage): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "fullgraph": ("BOOLEAN", { "default": False }), + "dynamic": ("BOOLEAN", { "default": False }), + "mode": (["default", "reduce-overhead", "max-autotune", "max-autotune-no-cudagraphs"],), + }, + } + + RETURN_TYPES = ("MODEL", ) + FUNCTION = "execute" + CATEGORY = "essentials" + + def execute(self, model, fullgraph, dynamic, mode): + model.model.diffusion_model = torch.compile(model.model.diffusion_model, dynamic=dynamic, fullgraph=fullgraph, mode=mode) + return( model, ) + class ConsoleDebug: def __init__(self): pass @@ -446,6 +466,8 @@ NODE_CLASS_MAPPINGS = { "SimpleMath+": SimpleMath, "ConsoleDebug+": ConsoleDebug, + + "ModelCompile+": ModelCompile } NODE_DISPLAY_NAME_MAPPINGS = { @@ -465,4 +487,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SimpleMath+": "🔧 Simple Math", "ConsoleDebug+": "🔧 Console Debug", + + "ModelCompile": "🔧 Compile Model", } \ No newline at end of file