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.
|
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
@@ -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
@@ -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
@@ -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
@@ -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]
|
||||||
|
|||||||
Reference in New Issue
Block a user