Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8de02bc769 | ||
|
|
54d60e0d18 |
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user