Improved approach to external nodepack integration

This commit is contained in:
blepping
2024-12-24 21:55:56 -07:00
parent 7e173e1611
commit f626ff6a89
8 changed files with 129 additions and 48 deletions
+1 -1
View File
@@ -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)
</details>
+4
View File
@@ -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.
+101 -34
View File
@@ -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
+3 -4
View File
@@ -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."
+1 -1
View File
@@ -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,
)
+7 -2
View File
@@ -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)
+6 -3
View File
@@ -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()
+6 -3
View File
@@ -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):