Author SHA1 Message Date
blepping 8de02bc769 Workaround for recent ComfyUI frontend breakage 2025-05-06 04:33:59 -06:00
blepping 54d60e0d18 Different approach to integrating external modules (#28)
* Different approach to integrating external modules
* Other internal cleanups
2024-12-24 21:47:01 -07:00
blepping 4e66c60163 Merge pull request #26 from blepping/attention_improvements
October update
2024-10-14 04:34:14 -06:00
5 changed files with 218 additions and 74 deletions
+11
View File
@@ -2,6 +2,17 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 20250506
ComfyUI in its infinite wisdom decided to make it so you can no longer have parameters that default to inputs but can be converted to widgets and widgets take up the full space even if you're using an input now. Because of this, the YAML parameters in the node will take up a lot more space and can't be hidden. I'd make it so the parameter was just always an input but then it would be impossible to acces the parameters in any workflows that had previously converted the input to a widget.
* Work around breakage caused by recent ComfyUI frontend versions.
## 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.
## 20241014
_Note_: Advanced MSW-MSA Attention node parameters changed. May break workflows.
+23 -13
View File
@@ -1,15 +1,15 @@
from __future__ import annotations
import itertools
import logging
import math
from time import time
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
from . import utils
from .utils import (
UPSCALE_METHODS,
IntegratedNode,
ModelType,
StrEnum,
TimeMode,
@@ -18,6 +18,7 @@ from .utils import (
convert_time,
get_sigma,
guess_model_type,
logger,
parse_blocks,
rescale_size,
scale_samples,
@@ -28,9 +29,19 @@ F = torch.nn.functional
if TYPE_CHECKING:
import comfy
SCALE_METHODS = ()
REVERSE_SCALE_METHODS = ()
def init_integrations(_integrations) -> None:
global scale_samples, SCALE_METHODS, REVERSE_SCALE_METHODS # noqa: PLW0603
SCALE_METHODS = ("disabled", "skip", *utils.UPSCALE_METHODS)
REVERSE_SCALE_METHODS = utils.UPSCALE_METHODS
scale_samples = utils.scale_samples
utils.MODULES.register_init_handler(init_integrations)
SCALE_METHODS = ("disabled", "skip", *UPSCALE_METHODS)
REVERSE_SCALE_METHODS = UPSCALE_METHODS
DEFAULT_WARN_INTERVAL = 60
@@ -193,7 +204,7 @@ class State:
or self.last_warned is None
or now - self.last_warned >= DEFAULT_WARN_INTERVAL
):
logging.warning(
logger.warning(
f"** jankhidiffusion: MSW-MSA attention({self.pretty_last_block}): {s}",
)
self.last_warned = now
@@ -202,7 +213,7 @@ class State:
return f"<MSWMSAAttentionState:last_sigma={self.last_sigma}, last_block={self.pretty_last_block}, last_shift={self.last_shift}, last_shifts={self.last_shifts}>"
class ApplyMSWMSAAttention:
class ApplyMSWMSAAttention(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
FUNCTION = "patch"
@@ -274,10 +285,9 @@ class ApplyMSWMSAAttention:
"yaml_parameters": (
"STRING",
{
"tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. This input can be converted into a multiline text widget. See main README for possible options. Note: When specifying paramaters this way, there is very little error checking.",
"tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. See main README for possible options. Note: When specifying paramaters this way, there is very little error checking.",
"dynamicPrompts": False,
"multiline": True,
"defaultInput": True,
},
),
},
@@ -455,7 +465,7 @@ class ApplyMSWMSAAttention:
cls,
*,
model: comfy.model_patcher.ModelPatcher,
yaml_parameters: None | str = None,
yaml_parameters: str | None = None,
**kwargs: dict[str, Any],
) -> tuple[comfy.model_patcher.ModelPatcher]:
if yaml_parameters:
@@ -477,7 +487,7 @@ class ApplyMSWMSAAttention:
if not config.use_blocks:
return (model,)
if config.verbose:
logging.info(
logger.info(
f"** jankhidiffusion: MSW-MSA Attention: Using config: {config}",
)
@@ -528,7 +538,7 @@ class ApplyMSWMSAAttention:
for idx, tensor in enumerate(attn_parts)
)
except (RuntimeError, ValueError) as exc:
logging.warning(
logger.warning(
f"** jankhidiffusion: Exception applying MSW-MSA attention: Incompatible model patches or bad resolution. Try using resolutions that are multiples of 64 or set scale/reverse_scale modes to something other than disabled. Original exception: {exc}",
)
state.window_args = None
@@ -554,7 +564,7 @@ class ApplyMSWMSAAttention:
return (model,)
class ApplyMSWMSAAttentionSimple:
class ApplyMSWMSAAttentionSimple(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
FUNCTION = "go"
@@ -597,7 +607,7 @@ class ApplyMSWMSAAttentionSimple:
if preset is None:
errstr = f"Unknown model type {model_type!s}"
raise ValueError(errstr)
logging.info(
logger.info(
f"** ApplyMSWMSAAttentionSimple: Using preset {model_type!s}: in/mid/out blocks [{preset.pretty_blocks}], start/end percent {preset.start_time:.2}/{preset.end_time:.2}",
)
return ApplyMSWMSAAttention.patch(model=model, **preset.as_dict)
+42 -34
View File
@@ -1,17 +1,16 @@
from __future__ import annotations
import itertools
import logging
import os
import sys
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
from comfy.ldm.modules.diffusionmodules import openaimodel
from . import utils
from .utils import (
UPSCALE_METHODS,
IntegratedNode,
ModelType,
TimeMode,
check_time,
@@ -19,6 +18,7 @@ from .utils import (
fade_scale,
get_sigma,
guess_model_type,
logger,
parse_blocks,
scale_samples,
sigma_to_pct,
@@ -29,11 +29,20 @@ if TYPE_CHECKING:
F = torch.nn.functional
CA_DOWNSCALE_METHODS = (
("avg_pool2d", "adaptive_avg_pool2d", *UPSCALE_METHODS)
if "adaptive_avg_pool2d" not in UPSCALE_METHODS
else ("avg_pool2d", *UPSCALE_METHODS)
)
CA_DOWNSCALE_METHODS = ()
def init_integrations(_integrations) -> None:
global scale_samples, CA_DOWNSCALE_METHODS # noqa: PLW0603
CA_DOWNSCALE_METHODS = (
("avg_pool2d", "adaptive_avg_pool2d", *utils.UPSCALE_METHODS)
if "adaptive_avg_pool2d" not in utils.UPSCALE_METHODS
else ("avg_pool2d", *utils.UPSCALE_METHODS)
)
scale_samples = utils.scale_samples
utils.MODULES.register_init_handler(init_integrations)
class Preset(NamedTuple):
@@ -138,7 +147,7 @@ class Config:
ca_upscale_mode: str = "bicubic"
ca_downscale_mode: str = "adaptive_avg_pool2d"
ca_downscale_factor: float = 2.0
ca_downscale_factor_w: None | float = None
ca_downscale_factor_w: float | None = None
# Patches the input block after the skip connection.
ca_input_after_skip_mode: bool = False
ca_avg_pool2d_ceil_mode: bool = True
@@ -154,12 +163,12 @@ class Config:
ca_pre_downscale_multiplier: float = 1.0
ca_post_downscale_multiplier: float = 1.0
# Allows fading out the scale effect starting from this time.
ca_fadeout_start_sigma: None | float = None
ca_fadeout_start_sigma: float | None = None
# Maximum fadeout, as a percentage of the total scale effect.
ca_fadeout_cap: float = 0.0
ca_latent_pixel_increment: int | float = 8
verbose: int = 0
curr_sigma: None | float = None
curr_sigma: float | None = None
@classmethod
def build(
@@ -175,7 +184,7 @@ class Config:
ca_end_time: float,
ca_input_blocks: str | list[int],
ca_output_blocks: str | list[int],
ca_fadeout_start_time: None | float = None,
ca_fadeout_start_time: float | None = None,
**kwargs: dict,
) -> object:
time_mode: TimeMode = TimeMode(time_mode)
@@ -254,7 +263,7 @@ class State:
def hd_apply_control(
self,
h: torch.Tensor,
control: None | dict,
control: dict | None,
name: str,
) -> torch.Tensor:
ctrls = control.get(name) if control is not None else None
@@ -264,7 +273,7 @@ class State:
if ctrl is None:
return h
if ctrl.shape[-2:] != h.shape[-2:]:
logging.info(
logger.info(
f"* jankhidiffusion: Scaling controlnet conditioning: {ctrl.shape[-2:]} -> {h.shape[-2:]}",
)
ctrl = F.interpolate(ctrl, size=h.shape[-2:], **self.controlnet_scale_args)
@@ -279,7 +288,7 @@ class State:
return
self.orig_apply_control = openaimodel.apply_control
openaimodel.apply_control = self.hd_apply_control
logging.info("** jankhidiffusion: Patched openaimodel.apply_control")
logger.info("** jankhidiffusion: Patched openaimodel.apply_control")
# Try to be compatible with FreeU Advanced.
def try_patch_freeu_advanced(self) -> None:
@@ -288,13 +297,13 @@ class State:
# We only try one time.
self.patched_freeu_advanced = True
fua_nodes = sys.modules.get("FreeU_Advanced.nodes")
fua_nodes = getattr(utils.MODULES.freeu_advanced, "nodes", None)
if not fua_nodes:
return
self.orig_fua_apply_control = fua_nodes.apply_control
fua_nodes.apply_control = self.hd_apply_control
logging.info("** jankhidiffusion: Patched FreeU_Advanced")
logger.info("** jankhidiffusion: Patched FreeU_Advanced")
def apply_patches(self) -> None:
self.try_patch_apply_control()
@@ -303,18 +312,18 @@ class State:
def revert_patches(self) -> None:
if openaimodel.apply_control == self.hd_apply_control:
openaimodel.apply_control = self.orig_apply_control
logging.info("** jankhidiffusion: Reverted openaimodel.apply_control patch")
logger.info("** jankhidiffusion: Reverted openaimodel.apply_control patch")
if not self.patched_freeu_advanced:
return
fua_nodes = sys.modules.get("FreeU_Advanced.nodes")
fua_nodes = getattr(utils.MODULES.freeu_advanced, "nodes", None)
if not fua_nodes:
logging.warning(
logger.warning(
"** jankhidiffusion: Unexpectedly could not revert FreeU_Advanced patches",
)
return
fua_nodes.apply_control = self.orig_fua_apply_control
self.patched_freeu_advanced = False
logging.info("** jankhidiffusion: Reverted FreeU_Advanced patch")
logger.info("** jankhidiffusion: Reverted FreeU_Advanced patch")
GLOBAL_STATE: State = State()
@@ -353,7 +362,7 @@ class HDForward:
def forward_upsample(
self,
x: torch.Tensor,
output_shape: None | tuple = None,
output_shape: tuple | None = None,
) -> torch.Tensor:
config = self.config
orig_block = self.orig_block
@@ -451,7 +460,7 @@ class HDForward:
)
class ApplyRAUNet:
class ApplyRAUNet(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",)
FUNCTION = "patch"
@@ -512,7 +521,7 @@ class ApplyRAUNet:
},
),
"upscale_mode": (
UPSCALE_METHODS,
utils.UPSCALE_METHODS,
{
"tooltip": "Method used when upscaling latents in output Upscale blocks.",
},
@@ -554,7 +563,7 @@ class ApplyRAUNet:
},
),
"ca_upscale_mode": (
UPSCALE_METHODS,
utils.UPSCALE_METHODS,
{
"tooltip": "Mode used when upscaling latents in output cross-attention blocks.",
},
@@ -577,7 +586,7 @@ class ApplyRAUNet:
},
),
"two_stage_upscale_mode": (
("disabled", *UPSCALE_METHODS),
("disabled", *utils.UPSCALE_METHODS),
{
"default": "disabled",
"tooltip": "When upscaling in output Upscale blocks (non-NA), do half the upscale with this mode and half with the normal upscale mode. May produce a different effect, isn't necessarily better.",
@@ -588,10 +597,9 @@ class ApplyRAUNet:
"yaml_parameters": (
"STRING",
{
"tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. This input can be converted into a multiline text widget. See main README for possible options. Note: When specifying paramaters this way, there is very little error checking.",
"tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. See main README for possible options. Note: When specifying paramaters this way, there is very little error checking.",
"dynamicPrompts": False,
"multiline": True,
"defaultInput": True,
},
),
},
@@ -602,7 +610,7 @@ class ApplyRAUNet:
cls,
*,
model: ModelPatcher,
yaml_parameters: None | str = None,
yaml_parameters: str | None = None,
**kwargs: dict[str, Any],
) -> tuple[ModelPatcher]:
if yaml_parameters:
@@ -630,7 +638,7 @@ class ApplyRAUNet:
"avg_pool2d downscale mode can only be used with integer downscale factors",
)
if config.verbose:
logging.info(f"** jankhidiffusion: RAUNet: Using config: {config}")
logger.info(f"** jankhidiffusion: RAUNet: Using config: {config}")
have_ca_output_blocks = any(bt == "output" for (bt, _) in config.ca_use_blocks)
model = model.clone()
@@ -816,7 +824,7 @@ class ApplyRAUNet:
return (model,)
class ApplyRAUNetSimple:
class ApplyRAUNetSimple(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",)
FUNCTION = "patch"
@@ -852,7 +860,7 @@ class ApplyRAUNetSimple:
"upscale_mode": (
(
"default",
*UPSCALE_METHODS,
*utils.UPSCALE_METHODS,
),
{
"tooltip": "Method used when upscaling latents in output Upsample blocks.",
@@ -861,7 +869,7 @@ class ApplyRAUNetSimple:
"ca_upscale_mode": (
(
"default",
*UPSCALE_METHODS,
*utils.UPSCALE_METHODS,
),
{
"tooltip": "Method used when upscaling latents in cross attention blocks.",
@@ -898,7 +906,7 @@ class ApplyRAUNetSimple:
upscale_mode=upscale_mode,
ca_upscale_mode=ca_upscale_mode,
)
logging.info(
logger.info(
f"** ApplyRAUNetSimple: Using preset {model_type!s} {res}: upscale {upscale_mode}, in/out blocks [{preset.pretty_blocks}], start/end percent {preset.start_time:.2}/{preset.end_time:.2} | CA upscale {preset.ca_upscale_mode}, CA in/out blocks [{preset.ca_pretty_blocks}], CA start/end percent {preset.ca_start_time:.2}/{preset.ca_end_time:.2}",
)
return ApplyRAUNet.patch(model=model, **preset.as_dict)
+141 -26
View File
@@ -1,14 +1,22 @@
from __future__ import annotations
import contextlib
import importlib
import itertools
import logging
import math
from typing import Sequence
import sys
from functools import partial
from typing import TYPE_CHECKING, Callable, NamedTuple
import torch.nn.functional as torchf
from comfy import latent_formats
from comfy.utils import bislerp
if TYPE_CHECKING:
from collections.abc import Sequence
from types import ModuleType
try:
from enum import StrEnum
except ImportError:
@@ -24,6 +32,8 @@ except ImportError:
return str(self.value)
logger = logging.getLogger(__name__)
UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "nearest", "area")
@@ -77,7 +87,7 @@ def convert_time(
raise ValueError("invalid time mode")
def get_sigma(options: dict, key: str = "sigmas") -> None | float:
def get_sigma(options: dict, key: str = "sigmas") -> float | None:
if not isinstance(options, dict):
return None
sigmas = options.get(key)
@@ -148,7 +158,7 @@ def rescale_size(
raise ValueError(msg)
def guess_model_type(model: object) -> None | ModelType:
def guess_model_type(model: object) -> ModelType | None:
latent_format = model.get_model_object("latent_format")
if isinstance(latent_format, latent_formats.SD15):
return ModelType.SD15
@@ -179,33 +189,138 @@ def fade_scale(
return max(fade_cap, scaling_pct)
try:
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
bleh_latentutils = getattr(bleh.py, "latent_utils", None)
def scale_samples(
samples,
width,
height,
mode="bicubic",
sigma=None, # noqa: ARG001
):
if mode == "bislerp":
return bislerp(samples, width, height)
return torchf.interpolate(samples, size=(height, width), mode=mode)
class Integrations:
class Integration(NamedTuple):
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)
class JHDIntegrations(Integrations):
def __init__(self, *args: list, **kwargs: dict):
super().__init__(*args, **kwargs)
self.register_integration("bleh", "ComfyUI-bleh", self.bleh_integration)
self.register_integration("freeu_advanced", "FreeU_Advanced")
@classmethod
def bleh_integration(cls, bleh: ModuleType) -> ModuleType | None:
bleh_version = getattr(bleh, "BLEH_VERSION", -1)
if bleh_version < 0:
return None
return bleh
MODULES = JHDIntegrations()
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
def init_integrations(integrations) -> None:
global scale_samples, UPSCALE_METHODS # noqa: PLW0603
ext_bleh = integrations.bleh
if ext_bleh is None:
return
bleh_latentutils = getattr(ext_bleh.py, "latent_utils", None)
if bleh_latentutils is None:
raise ImportError # noqa: TRY301
bleh_version = getattr(bleh, "BLEH_VERSION", -1)
if bleh_version < 0:
def scale_samples(*args: list, sigma=None, **kwargs: dict): # noqa: ARG001
return bleh_latentutils.scale_samples(*args, **kwargs)
else:
scale_samples = bleh_latentutils.scale_samples
return
bleh_version = getattr(ext_bleh, "BLEH_VERSION", -1)
UPSCALE_METHODS = bleh_latentutils.UPSCALE_METHODS
except (ImportError, NotImplementedError):
if bleh_version >= 0:
scale_samples = bleh_latentutils.scale_samples
return
def scale_samples(
samples,
width,
height,
mode="bicubic",
sigma=None, # noqa: ARG001
):
if mode == "bislerp":
return bislerp(samples, width, height)
return torchf.interpolate(samples, size=(height, width), mode=mode)
def scale_samples_wrapped(*args: list, sigma=None, **kwargs: dict): # noqa: ARG001
return bleh_latentutils.scale_samples(*args, **kwargs)
scale_samples = scale_samples_wrapped
MODULES.register_init_handler(init_integrations)
__all__ = (
"UPSCALE_METHODS",
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui_jankhidiffusion"
description = "Janky implementation of HiDiffusion for ComfyUI. Enables generating at resolutions higher than what the model was trained for. Only supports SD 1.x (maybe 2.x) and SDXL."
version = "0.8.4"
version = "0.8.5"
license = { file = "LICENSE" }
[project.urls]