Improved approach to external nodepack integration
This commit is contained in:
@@ -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>
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user