Compare commits
4
Commits
main
...
fix_integration
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
78eb9a8447 | ||
|
|
a54e89efa5 | ||
|
|
f6449b0ab0 | ||
|
|
0ad10230cf |
@@ -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
@@ -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
@@ -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
@@ -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