From 8e71286b6dc05af33cfd3069c08208537548d0bc Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 5 Aug 2025 01:31:31 +0300 Subject: [PATCH] Allow torch.compile VAE decoder Slight speedup especially for 5B... --- nodes_model_loading.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 4635a40..f62c079 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1312,6 +1312,7 @@ class WanVideoVAELoader: "precision": (["fp16", "fp32", "bf16"], {"default": "bf16"} ), + "compile_args": ("WANCOMPILEARGS", ), } } @@ -1321,7 +1322,7 @@ class WanVideoVAELoader: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'" - def loadmodel(self, model_name, precision): + def loadmodel(self, model_name, precision, compile_args=None): from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38 dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] @@ -1342,7 +1343,9 @@ class WanVideoVAELoader: vae.load_state_dict(vae_sd) vae.eval() vae.to(device = offload_device, dtype = dtype) - + if compile_args is not None: + vae.model.decoder = torch.compile(vae.model.decoder, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + print(vae) return (vae,)