diff --git a/comfy_api/torch_helpers/__init__.py b/comfy_api/torch_helpers/__init__.py new file mode 100644 index 000000000..be7ae7a61 --- /dev/null +++ b/comfy_api/torch_helpers/__init__.py @@ -0,0 +1,5 @@ +from .torch_compile import set_torch_compile_wrapper + +__all__ = [ + "set_torch_compile_wrapper", +] diff --git a/comfy_api/torch_helpers/torch_compile.py b/comfy_api/torch_helpers/torch_compile.py new file mode 100644 index 000000000..09429087e --- /dev/null +++ b/comfy_api/torch_helpers/torch_compile.py @@ -0,0 +1,39 @@ +from __future__ import annotations +import torch + +from comfy.patcher_extension import WrappersMP +from typing import TYPE_CHECKING, Callable, Optional +if TYPE_CHECKING: + from comfy.model_patcher import ModelPatcher + from comfy.patcher_extension import WrapperExecutor + + +COMPILE_KEY = "torch.compile" + + +def apply_torch_compile_factory(compiled_diffusion_model: Callable) -> Callable: + ''' + Create a wrapper that will refer to the compiled_diffusion_model + ''' + def apply_torch_compile_wrapper(executor: WrapperExecutor, *args, **kwargs): + try: + orig_diffusion_model = executor.class_obj.diffusion_model + executor.class_obj.diffusion_model = compiled_diffusion_model + return executor(*args, **kwargs) + finally: + executor.class_obj.diffusion_model = orig_diffusion_model + return apply_torch_compile_wrapper + + +def set_torch_compile_wrapper(model: ModelPatcher, backend: str, options: Optional[dict[str,str]]=None, *args, **kwargs): + # clear out any other torch.compile wrappers + model.remove_wrappers_with_key(WrappersMP.APPLY_MODEL, COMPILE_KEY) + # add torch.compile wrapper + wrapper_func = apply_torch_compile_factory( + torch.compile( + model=model.get_model_object("diffusion_model"), + backend=backend, + options=options, + ), + ) + model.add_wrapper_with_key(WrappersMP.APPLY_MODEL, COMPILE_KEY, wrapper_func) diff --git a/comfy_extras/nodes_torch_compile.py b/comfy_extras/nodes_torch_compile.py index 9e46a1e8e..605536678 100644 --- a/comfy_extras/nodes_torch_compile.py +++ b/comfy_extras/nodes_torch_compile.py @@ -1,38 +1,4 @@ -from __future__ import annotations -import torch - -from comfy.patcher_extension import WrappersMP, WrapperExecutor -from typing import TYPE_CHECKING, Callable, Optional -if TYPE_CHECKING: - from comfy.model_patcher import ModelPatcher - - -COMPILE_KEY = "torch.compile" - - -def apply_torch_compile_factory(compiled_diffusion_model: Callable) -> Callable: - def apply_torch_compile_wrapper(executor: WrapperExecutor, *args, **kwargs): - try: - orig_diffusion_model = executor.class_obj.diffusion_model - executor.class_obj.diffusion_model = compiled_diffusion_model - return executor(*args, **kwargs) - finally: - executor.class_obj.diffusion_model = orig_diffusion_model - return apply_torch_compile_wrapper - - -def set_torch_compile_wrapper(model: ModelPatcher, backend: str, options: Optional[dict[str,str]]=None, *args, **kwargs): - # clear out any other torch.compile wrappers - model.remove_wrappers_with_key(WrappersMP.APPLY_MODEL, COMPILE_KEY) - # add torch.compile wrapper - wrapper_func = apply_torch_compile_factory( - torch.compile( - model=model.get_model_object("diffusion_model"), - backend=backend, - options=options, - ), - ) - model.add_wrapper_with_key(WrappersMP.APPLY_MODEL, COMPILE_KEY, wrapper_func) +from comfy_api.torch_helpers import set_torch_compile_wrapper class TorchCompileModel: