From f626ff6a89eaa72dc711a058560cfcdae20ae4f0 Mon Sep 17 00:00:00 2001 From: blepping Date: Tue, 24 Dec 2024 21:55:56 -0700 Subject: [PATCH] Improved approach to external nodepack integration --- README.md | 2 +- changelog.md | 4 ++ py/external.py | 135 ++++++++++++++++++++++++++++++----------- py/nodes.py | 7 +-- py/sampler.py | 2 +- py/schedule.py | 9 ++- py/tensor_image_ops.py | 9 ++- py/vae.py | 9 ++- 8 files changed, 129 insertions(+), 48 deletions(-) diff --git a/README.md b/README.md index 3b3db6d..f2c46d2 100644 --- a/README.md +++ b/README.md @@ -248,7 +248,7 @@ iteration_override: `schedule_name` and `steps` are required, `denoise` is optional and defaults to `1.0` (not recommended for actual use). You may also specify additional parameters if the scheduler node supports them. For example, `karras` supports `sigma_min`, `sigma_max` and `rho`. `sigma_min` and `sigma_max` will default to the model's values which may be different from the node. -Supported schedules: `alignyoursteps`, `beta`, `ddim_uniform`, `exponential`, `gits`, `karras`, `laplace`, `normal`, `polyexponential`, `sgm_uniform`, `simple`, `vp` +Supported schedules: `alignyoursteps`, `beta`, `ddim_uniform`, `exponential`, `gits`, `karras`, `laplace`, `normal`, `polyexponential`, `sgm_uniform`, `simple`, `vp`, `kl_optimal` (once/if support is merged into ComfyUI) diff --git a/changelog.md b/changelog.md index b8e08ec..3380915 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,10 @@ Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top. +## 20241224 + +Reworked approach to integrating with external node packs. This _shouldn't_ cause any visible changes from a user perspective but please create an issue if you notice anything weird. + ## 20241116 * Added the ability to mask both guidance and global changes. See the README section on masks. diff --git a/py/external.py b/py/external.py index bdbe384..05bc268 100644 --- a/py/external.py +++ b/py/external.py @@ -1,46 +1,113 @@ +from __future__ import annotations + import contextlib import importlib.util import sys +from functools import partial +from typing import TYPE_CHECKING, Callable, NamedTuple -EXTERNAL = {} -INITIALIZED = False +if TYPE_CHECKING: + from types import ModuleType -def get_custom_node(name): - module_key = f"custom_nodes.{name}" - try: - spec = importlib.util.find_spec(module_key) - if spec is None: - raise ModuleNotFoundError(module_key) - module = next( - v - for v in sys.modules.copy().values() - if hasattr(v, "__spec__") - and v.__spec__ is not None - and v.__spec__.origin == spec.origin - ) - except StopIteration: - raise ModuleNotFoundError(module_key) from None - return module +class Integrations: + class Integration(NamedTuple): # noqa: D106 + key: str + module_name: str + handler: Callable | None = None + + def __init__(self): + self.initialized = False + self.modules = {} + self.init_handlers = [] + self.handlers = [] + + def __getitem__(self, key): + return self.modules[key] + + def __contains__(self, key): + return key in self.modules + + def __getattr__(self, key): + return self.modules.get(key) + + @staticmethod + def get_custom_node(name: str) -> ModuleType | None: + module_key = f"custom_nodes.{name}" + with contextlib.suppress(StopIteration): + spec = importlib.util.find_spec(module_key) + if spec is None: + return None + return next( + v + for v in sys.modules.copy().values() + if hasattr(v, "__spec__") + and v.__spec__ is not None + and v.__spec__.origin == spec.origin + ) + return None + + def register_init_handler(self, handler): + self.init_handlers.append(handler) + + def register_integration(self, key: str, module_name: str, handler=None) -> None: + if self.initialized: + raise ValueError( + "Internal error: Cannot register integration after initialization", + ) + if any(item[0] == key or item[1] == module_name for item in self.handlers): + errstr = ( + f"Module {module_name} ({key}) already in integration handlers list!" + ) + raise ValueError(errstr) + self.handlers.append(self.Integration(key, module_name, handler)) + + def initialize(self) -> None: + if self.initialized: + return + self.initialized = True + for ih in self.handlers: + module = self.get_custom_node(ih.module_name) + if module is None: + continue + if ih.handler is not None: + module = ih.handler(module) + if module is not None: + self.modules[ih.key] = module + + for init_handler in self.init_handlers: + init_handler(self) -def init_integrations() -> None: - global INITIALIZED # noqa: PLW0603 - if INITIALIZED: - return - INITIALIZED = True +class JDHIntegrations(Integrations): + def __init__(self, *args: list, **kwargs: dict): + super().__init__(*args, **kwargs) + self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration) + self.register_integration("tiled_diffusion", "ComfyUI-TiledDiffusion") - with contextlib.suppress(ModuleNotFoundError): - EXTERNAL["tiled_diffusion"] = get_custom_node("ComfyUI-TiledDiffusion") - - with contextlib.suppress(ModuleNotFoundError, NotImplementedError): - bleh = get_custom_node("ComfyUI-bleh") - bleh_version = getattr(bleh, "BLEH_VERSION", -1) + @classmethod + def bleh_integration(cls, module: ModuleType) -> ModuleType | None: + bleh_version = getattr(module, "BLEH_VERSION", -1) if bleh_version < 1: - raise NotImplementedError - EXTERNAL["bleh"] = bleh.py + return None + return module.py - from . import tensor_image_ops, vae # noqa: PLC0415 + @classmethod + def sonar_integration(cls, module: ModuleType) -> ModuleType | None: + return module.py - tensor_image_ops.init_integrations() - vae.init_integrations() + +MODULES = JDHIntegrations() + + +class IntegratedNode(type): + @staticmethod + def wrap_INPUT_TYPES(orig_method: Callable, *args: list, **kwargs: dict) -> dict: + MODULES.initialize() + return orig_method(*args, **kwargs) + + def __new__(cls: type, name: str, bases: tuple, attrs: dict) -> object: + obj = type.__new__(cls, name, bases, attrs) + if hasattr(obj, "INPUT_TYPES"): + obj.INPUT_TYPES = partial(cls.wrap_INPUT_TYPES, obj.INPUT_TYPES) + return obj diff --git a/py/nodes.py b/py/nodes.py index d4f3d6f..c17586b 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -5,7 +5,7 @@ import yaml from comfy.samplers import KSAMPLER from .config import ParamGroup -from .external import init_integrations +from .external import IntegratedNode from .sampler import diffusehigh_sampler from .vae import VAEMode @@ -53,7 +53,7 @@ else: WILDCARD_PARAM = ",".join(PARAM_TYPES) -class DiffuseHighSamplerNode: +class DiffuseHighSamplerNode(metaclass=IntegratedNode): DESCRIPTION = "Jank DiffuseHigh sampler node, used for generating directly to resolutions higher than what the model was trained for. Can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input." OUTPUT_TOOLTIPS = ( "SAMPLER that can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input.", @@ -199,7 +199,6 @@ class DiffuseHighSamplerNode: yaml_parameters: str | None = None, **kwargs: dict, ) -> tuple[KSAMPLER]: - init_integrations() if yaml_parameters: extra_params = yaml.safe_load(yaml_parameters) if extra_params is None: @@ -248,7 +247,7 @@ class DiffuseHighSamplerNode: ) -class DiffuseHighParamNode: +class DiffuseHighParamNode(metaclass=IntegratedNode): RETURN_TYPES = ("DIFFUSEHIGH_PARAMS",) CATEGORY = "sampling/custom_sampling/JankDiffuseHigh" DESCRIPTION = "Jank DiffuseHigh parameter definition node. Used to set parameters like custom noise types that require an input." diff --git a/py/sampler.py b/py/sampler.py index 1a9ea16..9289a1d 100644 --- a/py/sampler.py +++ b/py/sampler.py @@ -425,7 +425,7 @@ class DiffuseHighSampler: self.seed_offset += 1 return custom_noise.make_noise_sampler( x, - sigmas.min(), + sigmas[sigmas > 0].min(), sigmas.max(), **custom_noise_params, ) diff --git a/py/schedule.py b/py/schedule.py index 67922f9..09eaa71 100644 --- a/py/schedule.py +++ b/py/schedule.py @@ -56,12 +56,12 @@ class Schedule: self.total_steps = int(steps / denoise) self.schedule_name = schedule_name self.schedule_kwargs = self.get_schedule_kwargs(schedule_name, kwargs) - self._sigmas: None | torch.Tensor = None + self._sigmas: torch.Tensor | None = None def get_schedule_kwargs( self, schedule_name: str, - schedule_kwargs: None | dict[str, Any] = None, + schedule_kwargs: dict[str, Any] | None = None, ) -> dict[str, int | float | str]: schedule_default_kwargs = self.schedule_default_kwargs.get( schedule_name.lower(), @@ -166,6 +166,11 @@ class Schedule: f=kds.get_sigmas_vp, ), } + if hasattr(samplers, "kl_optimal_scheduler"): + _make_sigmas_handlers["kl_optimal"] = partial( + make_sigmas_no_model_sampling, + f=samplers.kl_optimal_scheduler, + ) def make_sigmas(self) -> torch.Tensor: handler = self._make_sigmas_handlers.get(self.schedule_name) diff --git a/py/tensor_image_ops.py b/py/tensor_image_ops.py index 5ae39a8..01700ba 100644 --- a/py/tensor_image_ops.py +++ b/py/tensor_image_ops.py @@ -8,7 +8,7 @@ import torch import torchvision from PIL import Image as PILImage -from .external import EXTERNAL +from .external import MODULES as EXT if TYPE_CHECKING: from collections.abc import Sequence @@ -20,14 +20,17 @@ BLENDING_MODES = { } -def init_integrations(): +def init_integrations(integrations): global BLENDING_MODES # noqa: PLW0603 - ext_bleh = EXTERNAL.get("bleh") + ext_bleh = integrations.bleh if ext_bleh is not None: BLENDING_MODES.clear() BLENDING_MODES |= ext_bleh.latent_utils.BLENDING_MODES +EXT.register_init_handler(init_integrations) + + class SharpenMode(Enum): GAUSSIAN = auto() CONTRAST_ADAPTIVE = auto() diff --git a/py/vae.py b/py/vae.py index 88ba109..3703dfc 100644 --- a/py/vae.py +++ b/py/vae.py @@ -8,7 +8,7 @@ from comfy.model_management import device_supports_non_blocking from comfy.taesd import taesd from tqdm import tqdm -from .external import EXTERNAL +from .external import MODULES as EXT from .utils import fallback if TYPE_CHECKING: @@ -17,9 +17,12 @@ if TYPE_CHECKING: tiled_diffusion = None -def init_integrations(): +def init_integrations(integrations): global tiled_diffusion # noqa: PLW0603 - tiled_diffusion = EXTERNAL.get("tiled_diffusion") + tiled_diffusion = integrations.tiled_diffusion + + +EXT.register_init_handler(init_integrations) class VAEMode(Enum):