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. 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 ## 20241014
_Note_: Advanced MSW-MSA Attention node parameters changed. May break workflows. _Note_: Advanced MSW-MSA Attention node parameters changed. May break workflows.
+22 -11
View File
@@ -1,15 +1,15 @@
from __future__ import annotations from __future__ import annotations
import itertools import itertools
import logging
import math import math
from time import time from time import time
from typing import TYPE_CHECKING, Any, NamedTuple from typing import TYPE_CHECKING, Any, NamedTuple
import torch import torch
from . import utils
from .utils import ( from .utils import (
UPSCALE_METHODS, IntegratedNode,
ModelType, ModelType,
StrEnum, StrEnum,
TimeMode, TimeMode,
@@ -18,6 +18,7 @@ from .utils import (
convert_time, convert_time,
get_sigma, get_sigma,
guess_model_type, guess_model_type,
logger,
parse_blocks, parse_blocks,
rescale_size, rescale_size,
scale_samples, scale_samples,
@@ -28,9 +29,19 @@ F = torch.nn.functional
if TYPE_CHECKING: if TYPE_CHECKING:
import comfy 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 DEFAULT_WARN_INTERVAL = 60
@@ -193,7 +204,7 @@ class State:
or self.last_warned is None or self.last_warned is None
or now - self.last_warned >= DEFAULT_WARN_INTERVAL or now - self.last_warned >= DEFAULT_WARN_INTERVAL
): ):
logging.warning( logger.warning(
f"** jankhidiffusion: MSW-MSA attention({self.pretty_last_block}): {s}", f"** jankhidiffusion: MSW-MSA attention({self.pretty_last_block}): {s}",
) )
self.last_warned = now 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}>" 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",) RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",) OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
FUNCTION = "patch" FUNCTION = "patch"
@@ -455,7 +466,7 @@ class ApplyMSWMSAAttention:
cls, cls,
*, *,
model: comfy.model_patcher.ModelPatcher, model: comfy.model_patcher.ModelPatcher,
yaml_parameters: None | str = None, yaml_parameters: str | None = None,
**kwargs: dict[str, Any], **kwargs: dict[str, Any],
) -> tuple[comfy.model_patcher.ModelPatcher]: ) -> tuple[comfy.model_patcher.ModelPatcher]:
if yaml_parameters: if yaml_parameters:
@@ -477,7 +488,7 @@ class ApplyMSWMSAAttention:
if not config.use_blocks: if not config.use_blocks:
return (model,) return (model,)
if config.verbose: if config.verbose:
logging.info( logger.info(
f"** jankhidiffusion: MSW-MSA Attention: Using config: {config}", f"** jankhidiffusion: MSW-MSA Attention: Using config: {config}",
) )
@@ -528,7 +539,7 @@ class ApplyMSWMSAAttention:
for idx, tensor in enumerate(attn_parts) for idx, tensor in enumerate(attn_parts)
) )
except (RuntimeError, ValueError) as exc: 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}", 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 state.window_args = None
@@ -554,7 +565,7 @@ class ApplyMSWMSAAttention:
return (model,) return (model,)
class ApplyMSWMSAAttentionSimple: class ApplyMSWMSAAttentionSimple(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",) RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",) OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
FUNCTION = "go" FUNCTION = "go"
@@ -597,7 +608,7 @@ class ApplyMSWMSAAttentionSimple:
if preset is None: if preset is None:
errstr = f"Unknown model type {model_type!s}" errstr = f"Unknown model type {model_type!s}"
raise ValueError(errstr) 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}", 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) return ApplyMSWMSAAttention.patch(model=model, **preset.as_dict)
+41 -32
View File
@@ -1,17 +1,16 @@
from __future__ import annotations from __future__ import annotations
import itertools import itertools
import logging
import os import os
import sys
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, NamedTuple from typing import TYPE_CHECKING, Any, NamedTuple
import torch import torch
from comfy.ldm.modules.diffusionmodules import openaimodel from comfy.ldm.modules.diffusionmodules import openaimodel
from . import utils
from .utils import ( from .utils import (
UPSCALE_METHODS, IntegratedNode,
ModelType, ModelType,
TimeMode, TimeMode,
check_time, check_time,
@@ -19,6 +18,7 @@ from .utils import (
fade_scale, fade_scale,
get_sigma, get_sigma,
guess_model_type, guess_model_type,
logger,
parse_blocks, parse_blocks,
scale_samples, scale_samples,
sigma_to_pct, sigma_to_pct,
@@ -29,11 +29,20 @@ if TYPE_CHECKING:
F = torch.nn.functional F = torch.nn.functional
CA_DOWNSCALE_METHODS = ( CA_DOWNSCALE_METHODS = ()
("avg_pool2d", "adaptive_avg_pool2d", *UPSCALE_METHODS)
if "adaptive_avg_pool2d" not in UPSCALE_METHODS
else ("avg_pool2d", *UPSCALE_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): class Preset(NamedTuple):
@@ -138,7 +147,7 @@ class Config:
ca_upscale_mode: str = "bicubic" ca_upscale_mode: str = "bicubic"
ca_downscale_mode: str = "adaptive_avg_pool2d" ca_downscale_mode: str = "adaptive_avg_pool2d"
ca_downscale_factor: float = 2.0 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. # Patches the input block after the skip connection.
ca_input_after_skip_mode: bool = False ca_input_after_skip_mode: bool = False
ca_avg_pool2d_ceil_mode: bool = True ca_avg_pool2d_ceil_mode: bool = True
@@ -154,12 +163,12 @@ class Config:
ca_pre_downscale_multiplier: float = 1.0 ca_pre_downscale_multiplier: float = 1.0
ca_post_downscale_multiplier: float = 1.0 ca_post_downscale_multiplier: float = 1.0
# Allows fading out the scale effect starting from this time. # 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. # Maximum fadeout, as a percentage of the total scale effect.
ca_fadeout_cap: float = 0.0 ca_fadeout_cap: float = 0.0
ca_latent_pixel_increment: int | float = 8 ca_latent_pixel_increment: int | float = 8
verbose: int = 0 verbose: int = 0
curr_sigma: None | float = None curr_sigma: float | None = None
@classmethod @classmethod
def build( def build(
@@ -175,7 +184,7 @@ class Config:
ca_end_time: float, ca_end_time: float,
ca_input_blocks: str | list[int], ca_input_blocks: str | list[int],
ca_output_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, **kwargs: dict,
) -> object: ) -> object:
time_mode: TimeMode = TimeMode(time_mode) time_mode: TimeMode = TimeMode(time_mode)
@@ -254,7 +263,7 @@ class State:
def hd_apply_control( def hd_apply_control(
self, self,
h: torch.Tensor, h: torch.Tensor,
control: None | dict, control: dict | None,
name: str, name: str,
) -> torch.Tensor: ) -> torch.Tensor:
ctrls = control.get(name) if control is not None else None ctrls = control.get(name) if control is not None else None
@@ -264,7 +273,7 @@ class State:
if ctrl is None: if ctrl is None:
return h return h
if ctrl.shape[-2:] != h.shape[-2:]: if ctrl.shape[-2:] != h.shape[-2:]:
logging.info( logger.info(
f"* jankhidiffusion: Scaling controlnet conditioning: {ctrl.shape[-2:]} -> {h.shape[-2:]}", f"* jankhidiffusion: Scaling controlnet conditioning: {ctrl.shape[-2:]} -> {h.shape[-2:]}",
) )
ctrl = F.interpolate(ctrl, size=h.shape[-2:], **self.controlnet_scale_args) ctrl = F.interpolate(ctrl, size=h.shape[-2:], **self.controlnet_scale_args)
@@ -279,7 +288,7 @@ class State:
return return
self.orig_apply_control = openaimodel.apply_control self.orig_apply_control = openaimodel.apply_control
openaimodel.apply_control = self.hd_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. # Try to be compatible with FreeU Advanced.
def try_patch_freeu_advanced(self) -> None: def try_patch_freeu_advanced(self) -> None:
@@ -288,13 +297,13 @@ class State:
# We only try one time. # We only try one time.
self.patched_freeu_advanced = True 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: if not fua_nodes:
return return
self.orig_fua_apply_control = fua_nodes.apply_control self.orig_fua_apply_control = fua_nodes.apply_control
fua_nodes.apply_control = self.hd_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: def apply_patches(self) -> None:
self.try_patch_apply_control() self.try_patch_apply_control()
@@ -303,18 +312,18 @@ class State:
def revert_patches(self) -> None: def revert_patches(self) -> None:
if openaimodel.apply_control == self.hd_apply_control: if openaimodel.apply_control == self.hd_apply_control:
openaimodel.apply_control = self.orig_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: if not self.patched_freeu_advanced:
return return
fua_nodes = sys.modules.get("FreeU_Advanced.nodes") fua_nodes = getattr(utils.MODULES.freeu_advanced, "nodes", None)
if not fua_nodes: if not fua_nodes:
logging.warning( logger.warning(
"** jankhidiffusion: Unexpectedly could not revert FreeU_Advanced patches", "** jankhidiffusion: Unexpectedly could not revert FreeU_Advanced patches",
) )
return return
fua_nodes.apply_control = self.orig_fua_apply_control fua_nodes.apply_control = self.orig_fua_apply_control
self.patched_freeu_advanced = False self.patched_freeu_advanced = False
logging.info("** jankhidiffusion: Reverted FreeU_Advanced patch") logger.info("** jankhidiffusion: Reverted FreeU_Advanced patch")
GLOBAL_STATE: State = State() GLOBAL_STATE: State = State()
@@ -353,7 +362,7 @@ class HDForward:
def forward_upsample( def forward_upsample(
self, self,
x: torch.Tensor, x: torch.Tensor,
output_shape: None | tuple = None, output_shape: tuple | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
config = self.config config = self.config
orig_block = self.orig_block orig_block = self.orig_block
@@ -451,7 +460,7 @@ class HDForward:
) )
class ApplyRAUNet: class ApplyRAUNet(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",) RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",) OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",)
FUNCTION = "patch" FUNCTION = "patch"
@@ -512,7 +521,7 @@ class ApplyRAUNet:
}, },
), ),
"upscale_mode": ( "upscale_mode": (
UPSCALE_METHODS, utils.UPSCALE_METHODS,
{ {
"tooltip": "Method used when upscaling latents in output Upscale blocks.", "tooltip": "Method used when upscaling latents in output Upscale blocks.",
}, },
@@ -554,7 +563,7 @@ class ApplyRAUNet:
}, },
), ),
"ca_upscale_mode": ( "ca_upscale_mode": (
UPSCALE_METHODS, utils.UPSCALE_METHODS,
{ {
"tooltip": "Mode used when upscaling latents in output cross-attention blocks.", "tooltip": "Mode used when upscaling latents in output cross-attention blocks.",
}, },
@@ -577,7 +586,7 @@ class ApplyRAUNet:
}, },
), ),
"two_stage_upscale_mode": ( "two_stage_upscale_mode": (
("disabled", *UPSCALE_METHODS), ("disabled", *utils.UPSCALE_METHODS),
{ {
"default": "disabled", "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.", "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, cls,
*, *,
model: ModelPatcher, model: ModelPatcher,
yaml_parameters: None | str = None, yaml_parameters: str | None = None,
**kwargs: dict[str, Any], **kwargs: dict[str, Any],
) -> tuple[ModelPatcher]: ) -> tuple[ModelPatcher]:
if yaml_parameters: if yaml_parameters:
@@ -630,7 +639,7 @@ class ApplyRAUNet:
"avg_pool2d downscale mode can only be used with integer downscale factors", "avg_pool2d downscale mode can only be used with integer downscale factors",
) )
if config.verbose: 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) have_ca_output_blocks = any(bt == "output" for (bt, _) in config.ca_use_blocks)
model = model.clone() model = model.clone()
@@ -816,7 +825,7 @@ class ApplyRAUNet:
return (model,) return (model,)
class ApplyRAUNetSimple: class ApplyRAUNetSimple(metaclass=IntegratedNode):
RETURN_TYPES = ("MODEL",) RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",) OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",)
FUNCTION = "patch" FUNCTION = "patch"
@@ -852,7 +861,7 @@ class ApplyRAUNetSimple:
"upscale_mode": ( "upscale_mode": (
( (
"default", "default",
*UPSCALE_METHODS, *utils.UPSCALE_METHODS,
), ),
{ {
"tooltip": "Method used when upscaling latents in output Upsample blocks.", "tooltip": "Method used when upscaling latents in output Upsample blocks.",
@@ -861,7 +870,7 @@ class ApplyRAUNetSimple:
"ca_upscale_mode": ( "ca_upscale_mode": (
( (
"default", "default",
*UPSCALE_METHODS, *utils.UPSCALE_METHODS,
), ),
{ {
"tooltip": "Method used when upscaling latents in cross attention blocks.", "tooltip": "Method used when upscaling latents in cross attention blocks.",
@@ -898,7 +907,7 @@ class ApplyRAUNetSimple:
upscale_mode=upscale_mode, upscale_mode=upscale_mode,
ca_upscale_mode=ca_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}", 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) return ApplyRAUNet.patch(model=model, **preset.as_dict)
+141 -26
View File
@@ -1,14 +1,22 @@
from __future__ import annotations from __future__ import annotations
import contextlib
import importlib import importlib
import itertools import itertools
import logging
import math 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 import torch.nn.functional as torchf
from comfy import latent_formats from comfy import latent_formats
from comfy.utils import bislerp from comfy.utils import bislerp
if TYPE_CHECKING:
from collections.abc import Sequence
from types import ModuleType
try: try:
from enum import StrEnum from enum import StrEnum
except ImportError: except ImportError:
@@ -24,6 +32,8 @@ except ImportError:
return str(self.value) return str(self.value)
logger = logging.getLogger(__name__)
UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "nearest", "area") UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "nearest", "area")
@@ -77,7 +87,7 @@ def convert_time(
raise ValueError("invalid time mode") 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): if not isinstance(options, dict):
return None return None
sigmas = options.get(key) sigmas = options.get(key)
@@ -148,7 +158,7 @@ def rescale_size(
raise ValueError(msg) 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") latent_format = model.get_model_object("latent_format")
if isinstance(latent_format, latent_formats.SD15): if isinstance(latent_format, latent_formats.SD15):
return ModelType.SD15 return ModelType.SD15
@@ -179,33 +189,138 @@ def fade_scale(
return max(fade_cap, scaling_pct) return max(fade_cap, scaling_pct)
try: def scale_samples(
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh") samples,
bleh_latentutils = getattr(bleh.py, "latent_utils", None) 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: if bleh_latentutils is None:
raise ImportError # noqa: TRY301 return
bleh_version = getattr(bleh, "BLEH_VERSION", -1) bleh_version = getattr(ext_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
UPSCALE_METHODS = bleh_latentutils.UPSCALE_METHODS UPSCALE_METHODS = bleh_latentutils.UPSCALE_METHODS
except (ImportError, NotImplementedError): if bleh_version >= 0:
scale_samples = bleh_latentutils.scale_samples
return
def scale_samples( def scale_samples_wrapped(*args: list, sigma=None, **kwargs: dict): # noqa: ARG001
samples, return bleh_latentutils.scale_samples(*args, **kwargs)
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)
scale_samples = scale_samples_wrapped
MODULES.register_init_handler(init_integrations)
__all__ = ( __all__ = (
"UPSCALE_METHODS", "UPSCALE_METHODS",
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "comfyui_jankhidiffusion" 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." 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" } license = { file = "LICENSE" }
[project.urls] [project.urls]