4 Commits
Author SHA1 Message Date
blepping 78eb9a8447 Update changelog 2024-12-24 21:42:23 -07:00
blepping a54e89efa5 Remove unused import 2024-12-24 21:39:51 -07:00
blepping f6449b0ab0 Integration refactor part 2 2024-12-22 10:37:28 -07:00
blepping 0ad10230cf Different approach to integrating external modules
Other internal cleanups
2024-12-20 07:01:39 -07:00
5 changed files with 209 additions and 70 deletions
+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.
## 20241014
_Note_: Advanced MSW-MSA Attention node parameters changed. May break workflows.
+22 -11
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"
@@ -455,7 +466,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 +488,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 +539,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 +565,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 +608,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)
+41 -32
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.",
@@ -602,7 +611,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 +639,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 +825,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 +861,7 @@ class ApplyRAUNetSimple:
"upscale_mode": (
(
"default",
*UPSCALE_METHODS,
*utils.UPSCALE_METHODS,
),
{
"tooltip": "Method used when upscaling latents in output Upsample blocks.",
@@ -861,7 +870,7 @@ class ApplyRAUNetSimple:
"ca_upscale_mode": (
(
"default",
*UPSCALE_METHODS,
*utils.UPSCALE_METHODS,
),
{
"tooltip": "Method used when upscaling latents in cross attention blocks.",
@@ -898,7 +907,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]