From 666a2f562e5a8dfc2356d9a258072dfd109b1071 Mon Sep 17 00:00:00 2001 From: Your Name Date: Wed, 19 Nov 2025 10:44:05 +0800 Subject: [PATCH] fix the issue of recompile --- node_openvino.py | 230 ++++++++++++++++++++++++++++++++++++++--------- pyproject.toml | 2 +- 2 files changed, 189 insertions(+), 43 deletions(-) diff --git a/node_openvino.py b/node_openvino.py index 54fa67b..5eb61de 100644 --- a/node_openvino.py +++ b/node_openvino.py @@ -1,10 +1,158 @@ import torch import openvino as ov from typing_extensions import override +from typing import Optional import openvino.frontend.pytorch.torchdynamo.execute as ov_ex from comfy_api.latest import ComfyExtension, io from comfy_api.torch_helpers import set_torch_compile_wrapper +TORCH_COMPILE_KWARGS_VAE = "torch_compile_kwargs_vae" + + +class VAECompileWrapper: + """ + VAE compiler wrapper that mirrors set_torch_compile_wrapper + Dynamically swaps modules during forward instead of using setattr directly + """ + def __init__(self, vae): + self.vae = vae + self.first_stage = vae.first_stage_model + self.compiled_modules = {} + self.compile_kwargs = {} + self.is_active = False + + # Store original forward methods + self.original_encode = None + self.original_decode = None + + def compile(self, backend: str, options: Optional[dict] = None, + mode: Optional[str] = None, fullgraph=False, dynamic: Optional[bool] = None, + keys: Optional[list[str]] = None): + """Compile specified VAE modules""" + + # Clean previous compilation + if self.is_active: + self.remove() + + # Determine keys to compile + if keys is None: + keys = [] + if hasattr(self.first_stage, "taesd_encoder"): + keys = ["taesd_encoder", "taesd_decoder"] + else: + keys = ["encoder", "decoder"] + + # Compile arguments + compile_kwargs = { + "backend": backend, + "options": options, + "mode": mode, + "fullgraph": fullgraph, + "dynamic": dynamic, + } + compile_kwargs = {k: v for k, v in compile_kwargs.items() if v is not None} + + # Compile each module + for key in keys: + if not hasattr(self.first_stage, key): + continue + + try: + original_module = getattr(self.first_stage, key) + # ✅ Only compile module without setattr + compiled_module = torch.compile(original_module, **compile_kwargs) + self.compiled_modules[key] = compiled_module + print(f"✅ Successfully compiled VAE.{key}") + except Exception as e: + print(f"❌ Failed to compile VAE.{key}: {e}") + + if self.compiled_modules: + self.compile_kwargs = compile_kwargs + self._wrap_forward_methods() + self.is_active = True + + # Store into vae_options + if not hasattr(self.vae, 'vae_options'): + self.vae.vae_options = {} + self.vae.vae_options[TORCH_COMPILE_KWARGS_VAE] = compile_kwargs + + def _wrap_forward_methods(self): + """Wrap encode/decode to use compiled modules at runtime""" + + # Save original methods + if hasattr(self.first_stage, 'encode'): + self.original_encode = self.first_stage.encode + self.first_stage.encode = self._create_encode_wrapper() + + if hasattr(self.first_stage, 'decode'): + self.original_decode = self.first_stage.decode + self.first_stage.decode = self._create_decode_wrapper() + + def _create_encode_wrapper(self): + """Create encode wrapper""" + def encode_wrapper(x): + # Determine which encoder to use + encoder_key = "taesd_encoder" if "taesd_encoder" in self.compiled_modules else "encoder" + + if encoder_key in self.compiled_modules: + # Temporarily replace encoder + original_encoder = getattr(self.first_stage, encoder_key) + try: + # ✅ Use compiled encoder + compiled_encoder = self.compiled_modules[encoder_key] + return compiled_encoder(x) + except Exception as e: + print(f"Compiled encoder execution failed, falling back to original: {e}") + return original_encoder(x) + else: + # Use original method + return self.original_encode(x) + + return encode_wrapper + + def _create_decode_wrapper(self): + """Create decode wrapper""" + def decode_wrapper(z): + # Determine which decoder to use + decoder_key = "taesd_decoder" if "taesd_decoder" in self.compiled_modules else "decoder" + + if decoder_key in self.compiled_modules: + # Temporarily replace decoder + original_decoder = getattr(self.first_stage, decoder_key) + try: + # ✅ Use compiled decoder + compiled_decoder = self.compiled_modules[decoder_key] + return compiled_decoder(z) + except Exception as e: + print(f"Compiled decoder execution failed, falling back to original: {e}") + return original_decoder(z) + else: + # Use original method + return self.original_decode(z) + + return decode_wrapper + + def remove(self): + """Remove compilation wrapper""" + if not self.is_active: + return + + # Restore original methods + if self.original_encode is not None: + self.first_stage.encode = self.original_encode + if self.original_decode is not None: + self.first_stage.decode = self.original_decode + + # Clean up + self.compiled_modules.clear() + self.compile_kwargs.clear() + self.is_active = False + + if hasattr(self.vae, 'vae_options') and TORCH_COMPILE_KWARGS_VAE in self.vae.vae_options: + del self.vae.vae_options[TORCH_COMPILE_KWARGS_VAE] + + print("✅ VAE compilation removed") + class TorchCompileDiffusionOpenVINO(io.ComfyNode): @classmethod @@ -16,10 +164,7 @@ class TorchCompileDiffusionOpenVINO(io.ComfyNode): category="OpenVINO", inputs=[ io.Model.Input("model"), - io.Combo.Input( - "device", - options=available_devices, - ), + io.Combo.Input("device", options=available_devices), ], outputs=[io.Model.Output()], is_experimental=True, @@ -31,9 +176,12 @@ class TorchCompileDiffusionOpenVINO(io.ComfyNode): ov_ex.compiled_cache.clear() ov_ex.req_cache.clear() ov_ex.partitioned_modules.clear() + m = model.clone() set_torch_compile_wrapper( - model=m, backend="openvino", options={"device": device} + model=m, + backend="openvino", + options={"device": device} ) return io.NodeOutput(m) @@ -48,64 +196,62 @@ class TorchCompileVAEOpenVINO(io.ComfyNode): category="OpenVINO", inputs=[ io.Vae.Input("vae"), - io.Combo.Input( - "device", - options=available_devices, - ), - io.Boolean.Input( - "compile_encoder", - default=True, - ), - io.Boolean.Input( - "compile_decoder", - default=True, - ), + io.Combo.Input("device", options=available_devices), + io.Boolean.Input("compile_encoder", default=True), + io.Boolean.Input("compile_decoder", default=True), + io.Boolean.Input("remove_compile", default=False, + tooltip="Remove VAE compilation"), ], outputs=[io.Vae.Output()], is_experimental=True, ) @classmethod - def execute(cls, vae, device, compile_encoder, compile_decoder) -> io.NodeOutput: + def execute(cls, vae, device, compile_encoder, compile_decoder, remove_compile) -> io.NodeOutput: torch._dynamo.reset() ov_ex.compiled_cache.clear() ov_ex.req_cache.clear() ov_ex.partitioned_modules.clear() + + # Get or create wrapper + if not hasattr(vae, '_compile_wrapper'): + vae._compile_wrapper = VAECompileWrapper(vae) + + wrapper = vae._compile_wrapper + + # Remove compilation if requested + if remove_compile: + wrapper.remove() + return io.NodeOutput(vae) + + # Otherwise compile as requested + keys = [] + first_stage = vae.first_stage_model + has_taesd = hasattr(first_stage, "taesd_encoder") + if compile_encoder: - encoder_name = "encoder" - if hasattr(vae.first_stage_model, "taesd_encoder"): - encoder_name = "taesd_encoder" + keys.append("taesd_encoder" if has_taesd else "encoder") - setattr( - vae.first_stage_model, - encoder_name, - torch.compile( - getattr(vae.first_stage_model, encoder_name), - backend="openvino", - options={"device": device}, - ), - ) if compile_decoder: - decoder_name = "decoder" - if hasattr(vae.first_stage_model, "taesd_decoder"): - decoder_name = "taesd_decoder" + keys.append("taesd_decoder" if has_taesd else "decoder") - setattr( - vae.first_stage_model, - decoder_name, - torch.compile( - getattr(vae.first_stage_model, decoder_name), - backend="openvino", - options={"device": device}, - ), + if keys: + wrapper.compile( + backend="openvino", + options={"device": device}, + keys=keys, ) + return io.NodeOutput(vae) class OpenVINOTorchCompileExtension(ComfyExtension): @override async def get_node_list(self) -> list[type[io.ComfyNode]]: - return [TorchCompileDiffusionOpenVINO, TorchCompileVAEOpenVINO] + return [ + TorchCompileDiffusionOpenVINO, + TorchCompileVAEOpenVINO, + ] async def comfy_entrypoint() -> OpenVINOTorchCompileExtension: diff --git a/pyproject.toml b/pyproject.toml index 736d709..c3535c6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-openvino" description = "OpenVINO node is designed for optimizing the performance of model inference in ComfyUI by leveraging Intel OpenVINO toolkits. It can support running model on Intel CPU, GPU and NPU device." -version = "1.1.1" +version = "1.1.2" license = {file = "LICENSE"} [project.urls]