Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8de02bc769 | ||
|
|
54d60e0d18 | ||
|
|
4e66c60163 | ||
|
|
71435fc0ec | ||
|
|
30ffe2d939 | ||
|
|
debeccf722 | ||
|
|
e62bd8ea41 | ||
|
|
88dd9340aa | ||
|
|
f723a994a1 | ||
|
|
eca9689c03 | ||
|
|
c79cc33ff9 | ||
|
|
f17b2efe23 | ||
|
|
86afeb70f9 | ||
|
|
2b3fadfaf2 | ||
|
|
2b9c5b1c2e | ||
|
|
71f2ef42fd | ||
|
|
d4d3fd0ff3 | ||
|
|
64090c80b7 | ||
|
|
4f48873f98 | ||
|
|
922a400f6f |
@@ -1,11 +1,7 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
release: { types: ["published"] }
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
|
||||
@@ -187,6 +187,67 @@ of 32, 64 or 128 (may need to experiment). Known to work with ELLA, FreeU (V2),
|
||||
Input blocks downscale and output blocks upscale so the biggest effect on performance will be applying this
|
||||
to input blocks with a low block number and output blocks with a high block number.
|
||||
|
||||
<details>
|
||||
|
||||
<summary>YAML parameters</summary>
|
||||
|
||||
This input can be converted to a multi-line text widget. Allows setting advanced/rare parameters. You can also override the node parameters here. JSON is valid YAML so you can use that if you prefer.
|
||||
|
||||
Default parameter values:
|
||||
|
||||
```yaml
|
||||
# In addition to the extra advanced options, you can override any fields from
|
||||
# the node here. For example:
|
||||
# time_mode: percent
|
||||
|
||||
# Scale mode used as a fallback only when image sizes are not multiples of 64. May decrease image quality.
|
||||
# May also be set to "disabled" to disable the workaround or "skip" to skip MSW-MSA attention on incompatible sizes.
|
||||
scale_mode: nearest-exact
|
||||
|
||||
# Scale mode used to reverse the scale_mode scaling.
|
||||
reverse_scale_mode: nearest-exact
|
||||
|
||||
# One of global, block, both, ignore
|
||||
last_shift_mode: global
|
||||
|
||||
# One of decrement, increment, retry
|
||||
last_shift_strategy: decrement
|
||||
|
||||
# Can be enabled to disable the log warning about incompatible image sizes.
|
||||
silent: false
|
||||
|
||||
# Allow scaling the window before/after the window or window reverse operation.
|
||||
pre_window_multiplier: 1.0
|
||||
post_window_multiplier: 1.0
|
||||
pre_window_reverse_multiplier: 1.0
|
||||
post_window_reverse_multiplier: 1.0
|
||||
|
||||
# Positive/negative distance to search for candidate rescales when dealing with incompatible
|
||||
# resolutions. Can possibly be used to brute force attn2 application (you can set it to something
|
||||
# absurd like 32).
|
||||
rescale_search_tolerance: 1
|
||||
|
||||
# Not recommended. Forces applying the attention patch to attn2.
|
||||
force_apply_attn2: false
|
||||
|
||||
# Logging verbosity level. 1 - Dumps config at startup. 2 - Warnings are also no longer throttled.
|
||||
verbose: 0
|
||||
```
|
||||
|
||||
|
||||
* `scale_mode`: Scale mode used as a fallback only when image sizes are not multiples of 64. May decrease image quality. Use `disabled` to bypass the fallback (may result in error) or `skip` to skip using MSW-MSA attention when the image size is incompatible. Any of the available scaling modes may be used here.
|
||||
* `reverse_scale_mode`: Scale mode used to reverse the scaling done by `scale_mode`. No effect when `scale_mode` is not being applied.
|
||||
* `last_shift_mode`: `global` - tracking is independent of blocks. `block` - remembers the last shift by block. `both` - avoids using the last shift both by block and globally. `ignore` - just uses whatever shift was randomly picked.
|
||||
* `last_shift_strategy`: Only has an effect when `last_shift_mode` is not `ignore`. There are four possible shift types. `decrement` - decrements the shift type. `increment` - increments the shift type. `retry` - keeps generating random shifts until it hits one not on the ignore list (changes seeds most significantly).
|
||||
* `pre_window_multipler` (etc): You can multiply the tensor before/after the window or window reverse operation. There's generally no difference between doing it before or after unless you're using weird upscale modes from [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh). I don't know why/when this would be useful, but it's there if you want to mess with it!
|
||||
* `force_apply_attn2`: Forces applying to attn2 rather than attn1. **Warning**: MSW-MSA attention was not made for `attn2` and the sizes are guaranteed to be incompatible and require scaling. Using it also doesn't seem to improve performance, there isn't much reason to enable this unless you're a weirdo like me and just like trying strange things.
|
||||
|
||||
The last shift options are for trying to avoid choosing the same shift size consecutively. This may or may not actually be helpful.
|
||||
|
||||
**Note**: Normal error checking generally doesn't apply to parameters set/overriden here. You are allowed to shoot yourself in the foot and will likely just get an exception if you enter the wrong type/an absurd value.
|
||||
|
||||
</details>
|
||||
|
||||
### `ApplyRAUNet`
|
||||
|
||||
**Use case**: Helps avoid artifacts when generating at resolutions significantly higher than what the model
|
||||
@@ -258,6 +319,65 @@ probably work best if you don't want to manually set segments.
|
||||
other scaling effects that target the same blocks (i.e. Deep Shrink). By itself, I think it should be fine with
|
||||
HyperTile and Deep Cache though I haven't actually tested that. May not work properly with ControlNet.
|
||||
|
||||
<details>
|
||||
|
||||
<summary>YAML parameters</summary>
|
||||
|
||||
This input can be converted to a multi-line text widget. Allows setting advanced/rare parameters. You can also override the node parameters here. JSON is valid YAML so you can use that if you prefer.
|
||||
|
||||
Default parameter values:
|
||||
|
||||
```yaml
|
||||
# In addition to the extra advanced options, you can override any fields from
|
||||
# the node here. For example:
|
||||
# time_mode: percent
|
||||
|
||||
# Patches input blocks after the skip connection when enabled (similar to Kohya deep shrink).
|
||||
ca_input_after_skip_mode: false
|
||||
|
||||
# Either null or set to a time (with the same time mode as the other times).
|
||||
# Starts fading out the CA scaling effect, starting from the specified time.
|
||||
ca_fadeout_start_time: null
|
||||
|
||||
# Maximum fadeout, specified as a percentage of the total scaling effect.
|
||||
ca_fadeout_cap: 0.0
|
||||
|
||||
# null or float. Allows setting the width scale separately. When null the same
|
||||
# factor will be used for height and width.
|
||||
ca_downscale_factor_w: null
|
||||
|
||||
# When applying CA scaling, ensures the rescaled latent is divisible by the specified incremenrt.
|
||||
ca_latent_pixel_increment: 8
|
||||
|
||||
# When using the avg_pool2d method, enable ceil mode.
|
||||
# See: https://pytorch.org/docs/stable/generated/torch.nn.functional.avg_pool2d.html
|
||||
ca_avg_pool2d_ceil_mode: true
|
||||
|
||||
# Allows applying a multiplier to the tensor: can be set separately for before/after upscale, downscale
|
||||
# and whether it's CA or not.
|
||||
pre_upscale_multiplier: 1.0
|
||||
post_upscale_multiplier: 1.0
|
||||
pre_downscale_multiplier: 1.0
|
||||
post_downscale_multiplier: 1.0
|
||||
ca_pre_upscale_multiplier: 1.0
|
||||
ca_post_upscale_multiplier: 1.0
|
||||
ca_pre_downscale_multiplier: 1.0
|
||||
ca_post_downscale_multiplier: 1.0
|
||||
|
||||
# Logging verbosity level. 1 - Dumps configuration on startup.
|
||||
verbose: 0
|
||||
```
|
||||
|
||||
* `ca_input_after_skip_mode`: When applying CA scaling, the effect will occur after the skip connection. This is the default for Kohya Deep Shrink and may produce less noisy results. **Note**: This changes the corresponding output block you need to set if not targeting a downscale block (i.e. ones you can target with the main RAUNet effect). It seems like you generally just subtract one. Example: Using SD15 and targeting input 4, you'd normally use output 8 - use output 7 instead.
|
||||
* `ca_latent_pixel_increment`: Ensures the scaled sizes are a multiple of the latent pixel increment. The default of 8 should ensure the scaled size is compatible with MSW-MSA attention without scaling workarounds. *Note*: Has no effect when downscaling with `avg_pool2d`.
|
||||
* `ca_fadeout_start_time`: Will start fading out the CA downscale factor starting from the specified time (which uses the same time mode as other configured times). The fadeout occurs such that the downscale factor will reach `1.0` (no downscaling) at `ca_end_time`. This can (sometimes) help decrease artifacts compared to simply ending the scale effect abruptly.
|
||||
* `ca_fadeout_cap`: Only has an effect when fadeout is in effect (see above). This is expressed as a percentage of the scaling effect, so, for example, you could set it to `0.5` to fade out the first 50% of the downscale effect and after that the downscale would stay at 50% (of the total downscale effect) until `ca_end_time` is reached.
|
||||
* `pre_upscale_multipler` (etc): You can multiply the tensor before/after it's upscaled or downscaled. There's generally no difference between doing it before or after unless you're using weird upscale modes from [ComfyUI-bleh](https://github.com/blepping/ComfyUI-bleh). Should you multiply it? Maybe not! It's a setting to possibly mess with and (not very scientifically) it seems like applying a mild positive multiplier can help.
|
||||
|
||||
**Note**: Normal error checking generally doesn't apply to parameters set/overriden here. You are allowed to shoot yourself in the foot and will likely just get an exception if you enter the wrong type/an absurd value.
|
||||
|
||||
</details>
|
||||
|
||||
## Credits
|
||||
|
||||
Code based on the HiDiffusion original implementation: https://github.com/megvii-research/HiDiffusion
|
||||
|
||||
@@ -2,6 +2,38 @@
|
||||
|
||||
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.
|
||||
|
||||
_Note_: This update may slightly change seeds.
|
||||
|
||||
* MSW-MSA attention can now work with all images sizes. When the size is incompatible it will scale the latent which may affect quality. Contributed by @pamparamm. Thanks!
|
||||
* Scaling now tries to make the output size a multiple of 8 so it's compatible with MSW-MSA attention. May change seeds, set `ca_latent_pixel_increment: 1` in YAML parameters for the old behavior. *Note*: Does not apply if you use `avg_pool2d` for downscaling.
|
||||
* CA downscaling now uses `adaptive_avg_pool2d` as the default method which supports fractional downscale sizes. As far as I know, it's the same as `avg_pool2d` with integer sizes but it's possible this will change seeds.
|
||||
* Simple nodes now support an "auto" model type parameter that will try to guess the model from the latent type.
|
||||
* Added a `yaml_parameters` input to the advanced nodes which allows specifying advanced/uncommon parameters. See main README for possible settings.
|
||||
* You can now use a different scale factor for width and height in RAUNet CA scaling. See `ca_downscale_factor_w` in YAML parameters.
|
||||
* You can now fade out the CA scaling effect in RAUNet node. See `ca_fadeout_start_time` and `ca_fadeout_cap` in YAML parameters.
|
||||
* Simple nodes default parameters for SDXL models adjusted to match the official HiDiffusion ones more closely.
|
||||
|
||||
Check the expandable "YAML Parameters" sections in the main README for more information about advanced parameters added in this update.
|
||||
|
||||
## 20240827
|
||||
|
||||
* Fixed (hopefully) an issue with RAUNet model patching that could cause semi-non-deterministic output. Unfortunately the fix also may change seeds.
|
||||
|
||||
## 20240813
|
||||
|
||||
_Note_: Advanced RAUNet node parameters changed, will break workflows.
|
||||
|
||||
+412
-116
@@ -1,15 +1,75 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, NamedTuple
|
||||
import itertools
|
||||
import math
|
||||
from time import time
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
from .utils import *
|
||||
from . import utils
|
||||
from .utils import (
|
||||
IntegratedNode,
|
||||
ModelType,
|
||||
StrEnum,
|
||||
TimeMode,
|
||||
block_to_num,
|
||||
check_time,
|
||||
convert_time,
|
||||
get_sigma,
|
||||
guess_model_type,
|
||||
logger,
|
||||
parse_blocks,
|
||||
rescale_size,
|
||||
scale_samples,
|
||||
)
|
||||
|
||||
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)
|
||||
|
||||
DEFAULT_WARN_INTERVAL = 60
|
||||
|
||||
|
||||
class Preset(NamedTuple):
|
||||
input_blocks: str = ""
|
||||
middle_blocks: str = ""
|
||||
output_blocks: str = ""
|
||||
time_mode: TimeMode = TimeMode.PERCENT
|
||||
start_time: float = 0.2
|
||||
end_time: float = 1.0
|
||||
scale_mode: str = "nearest-exact"
|
||||
reverse_scale_mode: str = "nearest-exact"
|
||||
|
||||
@property
|
||||
def as_dict(self):
|
||||
return {k: getattr(self, k) for k in self._fields}
|
||||
|
||||
@property
|
||||
def pretty_blocks(self):
|
||||
blocks = (self.input_blocks, self.middle_blocks, self.output_blocks)
|
||||
return " / ".join(b or "none" for b in blocks)
|
||||
|
||||
|
||||
SIMPLE_PRESETS = {
|
||||
ModelType.SD15: Preset(input_blocks="1,2", output_blocks="11,10,9"),
|
||||
ModelType.SDXL: Preset(input_blocks="4,5", output_blocks="3,4,5"),
|
||||
}
|
||||
|
||||
|
||||
class WindowSize(NamedTuple):
|
||||
height: int
|
||||
@@ -27,7 +87,133 @@ class ShiftSize(WindowSize):
|
||||
pass
|
||||
|
||||
|
||||
class ApplyMSWMSAAttention:
|
||||
class LastShiftMode(StrEnum):
|
||||
GLOBAL = "global"
|
||||
BLOCK = "block"
|
||||
BOTH = "both"
|
||||
IGNORE = "ignore"
|
||||
|
||||
|
||||
class LastShiftStrategy(StrEnum):
|
||||
INCREMENT = "increment"
|
||||
DECREMENT = "decrement"
|
||||
RETRY = "retry"
|
||||
|
||||
|
||||
class Config(NamedTuple):
|
||||
start_sigma: float
|
||||
end_sigma: float
|
||||
use_blocks: set
|
||||
scale_mode: str = "nearest-exact"
|
||||
reverse_scale_mode: str = "nearest-exact"
|
||||
# Allows disabling the log warning for incompatible sizes.
|
||||
silent: bool = False
|
||||
# Mode for trying to avoid using the same window size consecutively.
|
||||
last_shift_mode: LastShiftMode = LastShiftMode.GLOBAL
|
||||
# Strategy to use when avoiding a duplicate window size.
|
||||
last_shift_strategy: LastShiftStrategy = LastShiftStrategy.INCREMENT
|
||||
# Allows multiplying the tensor going into/out of the window or window reverse effect.
|
||||
pre_window_multiplier: float = 1.0
|
||||
post_window_multiplier: float = 1.0
|
||||
pre_window_reverse_multiplier: float = 1.0
|
||||
post_window_reverse_multiplier: float = 1.0
|
||||
force_apply_attn2: bool = False
|
||||
rescale_search_tolerance: int = 1
|
||||
verbose: int = 0
|
||||
|
||||
@classmethod
|
||||
def build(
|
||||
cls,
|
||||
*,
|
||||
ms: object,
|
||||
input_blocks: str | list[int],
|
||||
middle_blocks: str | list[int],
|
||||
output_blocks: str | list[int],
|
||||
time_mode: str | TimeMode,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
**kwargs: dict,
|
||||
) -> object:
|
||||
time_mode: TimeMode = TimeMode(time_mode)
|
||||
start_sigma, end_sigma = convert_time(ms, time_mode, start_time, end_time)
|
||||
input_blocks, middle_blocks, output_blocks = itertools.starmap(
|
||||
parse_blocks,
|
||||
(
|
||||
("input", input_blocks),
|
||||
("middle", middle_blocks),
|
||||
("output", output_blocks),
|
||||
),
|
||||
)
|
||||
return cls.__new__(
|
||||
cls,
|
||||
start_sigma=start_sigma,
|
||||
end_sigma=end_sigma,
|
||||
use_blocks=input_blocks | middle_blocks | output_blocks,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def maybe_multiply(
|
||||
t: torch.Tensor,
|
||||
multiplier: float = 1.0,
|
||||
post: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if multiplier == 1.0:
|
||||
return t
|
||||
return t.mul_(multiplier) if post else t * multiplier
|
||||
|
||||
|
||||
class State:
|
||||
__slots__ = (
|
||||
"config",
|
||||
"last_block",
|
||||
"last_shift",
|
||||
"last_shifts",
|
||||
"last_sigma",
|
||||
"last_warned",
|
||||
"window_args",
|
||||
)
|
||||
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.last_warned = None
|
||||
self.reset()
|
||||
|
||||
def reset(self):
|
||||
self.window_args = None
|
||||
self.last_sigma = None
|
||||
self.last_block = None
|
||||
self.last_shift = None
|
||||
self.last_shifts = {}
|
||||
|
||||
@property
|
||||
def pretty_last_block(self) -> str:
|
||||
if self.last_block is None:
|
||||
return "unknown"
|
||||
bt, bnum = self.last_block
|
||||
attstr = "" if not self.config.force_apply_attn2 else "attn2."
|
||||
btstr = ("in", "mid", "out")[bt]
|
||||
return f"{attstr}{btstr}.{bnum}"
|
||||
|
||||
def maybe_warning(self, s):
|
||||
if self.config.silent:
|
||||
return
|
||||
now = time()
|
||||
if (
|
||||
self.config.verbose >= 2
|
||||
or self.last_warned is None
|
||||
or now - self.last_warned >= DEFAULT_WARN_INTERVAL
|
||||
):
|
||||
logger.warning(
|
||||
f"** jankhidiffusion: MSW-MSA attention({self.pretty_last_block}): {s}",
|
||||
)
|
||||
self.last_warned = now
|
||||
|
||||
def __repr__(self):
|
||||
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(metaclass=IntegratedNode):
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
|
||||
FUNCTION = "patch"
|
||||
@@ -56,16 +242,13 @@ class ApplyMSWMSAAttention:
|
||||
"STRING",
|
||||
{
|
||||
"default": "9,10,11",
|
||||
"tooltip": "Comma-separated list of output blocks to patch. Default is for SD1.x, you can try 5,4 for SDXL",
|
||||
"tooltip": "Comma-separated list of output blocks to patch. Default is for SD1.x, you can try 3,4,5 for SDXL",
|
||||
},
|
||||
),
|
||||
"time_mode": (
|
||||
(
|
||||
"percent",
|
||||
"timestep",
|
||||
"sigma",
|
||||
),
|
||||
tuple(str(val) for val in TimeMode),
|
||||
{
|
||||
"default": "percent",
|
||||
"tooltip": "Time mode controls how to interpret the values in start_time and end_time.",
|
||||
},
|
||||
),
|
||||
@@ -98,6 +281,16 @@ class ApplyMSWMSAAttention:
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"yaml_parameters": (
|
||||
"STRING",
|
||||
{
|
||||
"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,
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
# reference: https://github.com/microsoft/Swin-Transformer
|
||||
@@ -105,53 +298,104 @@ class ApplyMSWMSAAttention:
|
||||
@staticmethod
|
||||
def window_partition(
|
||||
x: torch.Tensor,
|
||||
window_size: WindowSize,
|
||||
shift_size: ShiftSize,
|
||||
height: int,
|
||||
width: int,
|
||||
state: State,
|
||||
window_index: int,
|
||||
) -> torch.Tensor:
|
||||
config = state.config
|
||||
scale_mode = config.scale_mode
|
||||
x = config.maybe_multiply(x, config.pre_window_multiplier)
|
||||
window_size, shift_size, height, width = state.window_args[window_index]
|
||||
do_rescale = (height % 2 + width % 2) != 0
|
||||
if do_rescale:
|
||||
if scale_mode == "skip":
|
||||
state.maybe_warning(
|
||||
"Incompatible latent size - skipping MSW-MSA attention.",
|
||||
)
|
||||
return x
|
||||
if scale_mode == "disabled":
|
||||
state.maybe_warning(
|
||||
"Incompatible latent size - trying to proceed anyway. This may result in an error.",
|
||||
)
|
||||
do_rescale = False
|
||||
else:
|
||||
state.maybe_warning(
|
||||
"Incompatible latent size - applying scaling workaround. Note: This may reduce quality - use resolutions that are multiples of 64 when possible.",
|
||||
)
|
||||
batch, _features, channels = x.shape
|
||||
wheight, wwidth = window_size
|
||||
x = x.view(batch, height, width, channels)
|
||||
if do_rescale:
|
||||
x = (
|
||||
scale_samples(
|
||||
x.permute(0, 3, 1, 2).contiguous(),
|
||||
wwidth * 2,
|
||||
wheight * 2,
|
||||
mode=scale_mode,
|
||||
sigma=state.last_sigma,
|
||||
)
|
||||
.permute(0, 2, 3, 1)
|
||||
.contiguous()
|
||||
)
|
||||
if shift_size.sum > 0:
|
||||
x = torch.roll(x, shifts=-shift_size, dims=(1, 2))
|
||||
x = x.view(
|
||||
batch,
|
||||
height // wheight,
|
||||
wheight,
|
||||
width // wwidth,
|
||||
wwidth,
|
||||
channels,
|
||||
)
|
||||
x = x.view(batch, 2, wheight, 2, wwidth, channels)
|
||||
windows = (
|
||||
x.permute(0, 1, 3, 2, 4, 5)
|
||||
.contiguous()
|
||||
.view(-1, window_size.height, window_size.width, channels)
|
||||
)
|
||||
return windows.view(-1, window_size.sum, channels)
|
||||
return config.maybe_multiply(
|
||||
windows.view(-1, window_size.sum, channels),
|
||||
config.post_window_multiplier,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def window_reverse(
|
||||
windows: torch.Tensor,
|
||||
window_size: WindowSize,
|
||||
shift_size: WindowSize,
|
||||
height: int,
|
||||
width: int,
|
||||
state: State,
|
||||
window_index: int = 0,
|
||||
) -> torch.Tensor:
|
||||
config = state.config
|
||||
windows = config.maybe_multiply(windows, config.pre_window_reverse_multiplier)
|
||||
window_size, shift_size, height, width = state.window_args[window_index]
|
||||
do_rescale = (height % 2 + width % 2) != 0
|
||||
if do_rescale:
|
||||
if config.scale_mode == "skip":
|
||||
return windows
|
||||
if config.scale_mode == "disabled":
|
||||
do_rescale = False
|
||||
batch, _features, channels = windows.shape
|
||||
wheight, wwidth = window_size
|
||||
windows = windows.view(-1, wheight, wwidth, channels)
|
||||
batch = int(
|
||||
windows.shape[0] / (height * width / wheight / wwidth),
|
||||
batch = int(windows.shape[0] / 4)
|
||||
x = windows.view(batch, 2, 2, wheight, wwidth, -1)
|
||||
x = (
|
||||
x.permute(0, 1, 3, 2, 4, 5)
|
||||
.contiguous()
|
||||
.view(batch, wheight * 2, wwidth * 2, -1)
|
||||
)
|
||||
x = windows.view(batch, height // wheight, width // wwidth, wheight, wwidth, -1)
|
||||
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(batch, height, width, -1)
|
||||
if shift_size.sum > 0:
|
||||
x = torch.roll(x, shifts=shift_size, dims=(1, 2))
|
||||
return x.view(batch, height * width, channels)
|
||||
if do_rescale:
|
||||
x = (
|
||||
scale_samples(
|
||||
x.permute(0, 3, 1, 2).contiguous(),
|
||||
width,
|
||||
height,
|
||||
mode=config.reverse_scale_mode,
|
||||
sigma=state.last_sigma,
|
||||
)
|
||||
.permute(0, 2, 3, 1)
|
||||
.contiguous()
|
||||
)
|
||||
return config.maybe_multiply(
|
||||
x.view(batch, height * width, channels),
|
||||
config.post_window_reverse_multiplier,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_window_args(
|
||||
config: Config,
|
||||
n: torch.Tensor,
|
||||
orig_shape: tuple,
|
||||
shift: int,
|
||||
@@ -159,14 +403,17 @@ class ApplyMSWMSAAttention:
|
||||
_batch, features, _channels = n.shape
|
||||
orig_height, orig_width = orig_shape[-2:]
|
||||
|
||||
downsample_ratio = int(
|
||||
((orig_height * orig_width) // features) ** 0.5,
|
||||
width, height = rescale_size(
|
||||
orig_width,
|
||||
orig_height,
|
||||
features,
|
||||
tolerance=config.rescale_search_tolerance,
|
||||
)
|
||||
height, width = (
|
||||
orig_height // downsample_ratio,
|
||||
orig_width // downsample_ratio,
|
||||
)
|
||||
wheight, wwidth = height // 2, width // 2
|
||||
# if (height, width) != (orig_height, orig_width):
|
||||
# print(
|
||||
# f"\nRESC: features={features}, orig={(orig_height, orig_width)}, new={(height, width)}",
|
||||
# )
|
||||
wheight, wwidth = math.ceil(height / 2), math.ceil(width / 2)
|
||||
|
||||
if shift == 0:
|
||||
shift_size = ShiftSize(0, 0)
|
||||
@@ -178,92 +425,146 @@ class ApplyMSWMSAAttention:
|
||||
shift_size = ShiftSize(wheight // 4 * 3, wwidth // 4 * 3)
|
||||
return (WindowSize(wheight, wwidth), shift_size, height, width)
|
||||
|
||||
@staticmethod
|
||||
def get_shift(
|
||||
curr_block: tuple,
|
||||
state: State,
|
||||
*,
|
||||
shift_count=4,
|
||||
) -> int:
|
||||
mode = state.config.last_shift_mode
|
||||
strat = state.config.last_shift_strategy
|
||||
shift = int(torch.rand(1, device="cpu").item() * shift_count)
|
||||
block_last_shift = state.last_shifts.get(curr_block)
|
||||
last_shift = state.last_shift
|
||||
if mode == LastShiftMode.BOTH:
|
||||
avoid = {block_last_shift, last_shift}
|
||||
elif mode == LastShiftMode.BLOCK:
|
||||
avoid = {block_last_shift}
|
||||
elif mode == LastShiftMode.GLOBAL:
|
||||
avoid = {last_shift}
|
||||
else:
|
||||
avoid = {}
|
||||
if shift in avoid:
|
||||
if strat == LastShiftStrategy.DECREMENT:
|
||||
while shift in avoid:
|
||||
shift -= 1
|
||||
if shift < 0:
|
||||
shift = shift_count - 1
|
||||
elif strat == LastShiftStrategy.RETRY:
|
||||
while shift in avoid:
|
||||
shift = int(torch.rand(1, device="cpu").item() * shift_count)
|
||||
else:
|
||||
# Increment
|
||||
while shift in avoid:
|
||||
shift = (shift + 1) % shift_count
|
||||
return shift
|
||||
|
||||
@classmethod
|
||||
def patch(
|
||||
cls,
|
||||
*,
|
||||
model: comfy.model_patcher.ModelPatcher,
|
||||
input_blocks: str,
|
||||
middle_blocks: str,
|
||||
output_blocks: str,
|
||||
time_mode: str,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
yaml_parameters: str | None = None,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> tuple[comfy.model_patcher.ModelPatcher]:
|
||||
use_blocks = parse_blocks("input", input_blocks)
|
||||
use_blocks |= parse_blocks("middle", middle_blocks)
|
||||
use_blocks |= parse_blocks("output", output_blocks)
|
||||
if yaml_parameters:
|
||||
import yaml # noqa: PLC0415
|
||||
|
||||
extra_params = yaml.safe_load(yaml_parameters)
|
||||
if extra_params is None:
|
||||
pass
|
||||
elif not isinstance(extra_params, dict):
|
||||
raise ValueError(
|
||||
"MSWMSAAttention: yaml_parameters must either be null or an object",
|
||||
)
|
||||
else:
|
||||
kwargs |= extra_params
|
||||
config = Config.build(
|
||||
ms=model.get_model_object("model_sampling"),
|
||||
**kwargs,
|
||||
)
|
||||
if not config.use_blocks:
|
||||
return (model,)
|
||||
if config.verbose:
|
||||
logger.info(
|
||||
f"** jankhidiffusion: MSW-MSA Attention: Using config: {config}",
|
||||
)
|
||||
|
||||
model = model.clone()
|
||||
if not use_blocks:
|
||||
return (model,)
|
||||
state = State(config)
|
||||
|
||||
window_args = last_block = last_shift = None
|
||||
|
||||
ms = model.get_model_object("model_sampling")
|
||||
|
||||
start_sigma, end_sigma = convert_time(ms, time_mode, start_time, end_time)
|
||||
|
||||
def attn1_patch(
|
||||
def attn_patch(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
extra_options: dict,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
nonlocal window_args, last_shift, last_block
|
||||
window_args = None
|
||||
last_block = extra_options.get("block")
|
||||
if last_block not in use_blocks or not check_time(
|
||||
extra_options,
|
||||
start_sigma,
|
||||
end_sigma,
|
||||
state.window_args = None
|
||||
sigma = get_sigma(extra_options)
|
||||
block = extra_options.get("block", ("missing", 0))
|
||||
curr_block = block_to_num(*block)
|
||||
if state.last_sigma is not None and sigma > state.last_sigma:
|
||||
# logging.warning(
|
||||
# f"Doing reset: block={block}, sigma={sigma}, state={state}",
|
||||
# )
|
||||
state.reset()
|
||||
state.last_block = curr_block
|
||||
state.last_sigma = sigma
|
||||
if block not in config.use_blocks or not check_time(
|
||||
sigma,
|
||||
config.start_sigma,
|
||||
config.end_sigma,
|
||||
):
|
||||
return q, k, v
|
||||
orig_shape = extra_options["original_shape"]
|
||||
# MSW-MSA
|
||||
shift = int(torch.rand(1, device="cpu").item() * 4)
|
||||
if shift == last_shift:
|
||||
shift = (shift + 1) % 4
|
||||
last_shift = shift
|
||||
window_args = tuple(
|
||||
cls.get_window_args(x, orig_shape, shift) if x is not None else None
|
||||
for x in (q, k, v)
|
||||
)
|
||||
shift = cls.get_shift(curr_block, state)
|
||||
state.last_shifts[curr_block] = state.last_shift = shift
|
||||
try:
|
||||
if q is not None and q is k and q is v:
|
||||
return (
|
||||
cls.window_partition(
|
||||
q,
|
||||
*window_args[0],
|
||||
),
|
||||
) * 3
|
||||
return tuple(
|
||||
cls.window_partition(x, *window_args[idx])
|
||||
# get_window_args() can fail with ValueError in rescale_size() for some weird resolutions/aspect ratios
|
||||
# so we catch it here and skip MSW-MSA attention in that case.
|
||||
state.window_args = tuple(
|
||||
cls.get_window_args(config, x, orig_shape, shift)
|
||||
if x is not None
|
||||
else None
|
||||
for idx, x in enumerate((q, k, v))
|
||||
for x in (q, k, v)
|
||||
)
|
||||
except RuntimeError as exc:
|
||||
logging.warning(
|
||||
f"** jankhidiffusion: MSW-MSA attention not applied: Incompatible model patches or bad resolution. Try using resolutions that are multiples of 32 or 64. Original exception: {exc}",
|
||||
attn_parts = (q,) if q is not None and q is k and q is v else (q, k, v)
|
||||
result = tuple(
|
||||
cls.window_partition(tensor, state, idx)
|
||||
if tensor is not None
|
||||
else None
|
||||
for idx, tensor in enumerate(attn_parts)
|
||||
)
|
||||
window_args = None
|
||||
except (RuntimeError, ValueError) as exc:
|
||||
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
|
||||
return q, k, v
|
||||
return result * 3 if len(result) == 1 else result
|
||||
|
||||
def attn1_output_patch(n: torch.Tensor, extra_options: dict) -> torch.Tensor:
|
||||
nonlocal window_args
|
||||
if window_args is None or last_block != extra_options.get("block"):
|
||||
window_args = None
|
||||
def attn_output_patch(n: torch.Tensor, extra_options: dict) -> torch.Tensor:
|
||||
if state.window_args is None or state.last_block != block_to_num(
|
||||
*extra_options.get("block", ("missing", 0)),
|
||||
):
|
||||
state.window_args = None
|
||||
return n
|
||||
args, window_args = window_args[0], None
|
||||
return cls.window_reverse(n, *args)
|
||||
result = cls.window_reverse(n, state)
|
||||
state.window_args = None
|
||||
return result
|
||||
|
||||
model.set_model_attn1_patch(attn1_patch)
|
||||
model.set_model_attn1_output_patch(attn1_output_patch)
|
||||
if not config.force_apply_attn2:
|
||||
model.set_model_attn1_patch(attn_patch)
|
||||
model.set_model_attn1_output_patch(attn_output_patch)
|
||||
else:
|
||||
model.set_model_attn2_patch(attn_patch)
|
||||
model.set_model_attn2_output_patch(attn_output_patch)
|
||||
return (model,)
|
||||
|
||||
|
||||
class ApplyMSWMSAAttentionSimple:
|
||||
class ApplyMSWMSAAttentionSimple(metaclass=IntegratedNode):
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
|
||||
FUNCTION = "go"
|
||||
@@ -275,9 +576,9 @@ class ApplyMSWMSAAttentionSimple:
|
||||
return {
|
||||
"required": {
|
||||
"model_type": (
|
||||
("SD15", "SDXL"),
|
||||
("auto", "SD15", "SDXL"),
|
||||
{
|
||||
"tooltip": "Model type being patched. Choose SD15 for SD 1.4, SD 2.x.",
|
||||
"tooltip": "Model type being patched. Generally safe to leave on auto. Choose SD15 for SD 1.4, SD 2.x.",
|
||||
},
|
||||
),
|
||||
"model": (
|
||||
@@ -292,29 +593,24 @@ class ApplyMSWMSAAttentionSimple:
|
||||
@classmethod
|
||||
def go(
|
||||
cls,
|
||||
model_type: str,
|
||||
model_type: str | ModelType,
|
||||
model: comfy.model_patcher.ModelPatcher,
|
||||
) -> tuple[comfy.model_patcher.ModelPatcher]:
|
||||
time_range = (0.2, 1.0)
|
||||
if model_type == "SD15":
|
||||
blocks = ("1,2", "", "11,10,9")
|
||||
elif model_type == "SDXL":
|
||||
blocks = ("4,5", "", "5,4")
|
||||
if model_type == "auto":
|
||||
guessed_model_type = guess_model_type(model)
|
||||
if guessed_model_type not in SIMPLE_PRESETS:
|
||||
raise RuntimeError("Unable to guess model type")
|
||||
model_type = guessed_model_type
|
||||
else:
|
||||
raise ValueError("Unknown model type")
|
||||
prettyblocks = " / ".join(b or "none" for b in blocks)
|
||||
logging.info(
|
||||
f"** ApplyMSWMSAAttentionSimple: Using preset {model_type}: in/mid/out blocks [{prettyblocks}], start/end percent {time_range[0]:.2}/{time_range[1]:.2}",
|
||||
)
|
||||
return ApplyMSWMSAAttention.patch(
|
||||
model=model,
|
||||
input_blocks=blocks[0],
|
||||
middle_blocks=blocks[1],
|
||||
output_blocks=blocks[2],
|
||||
time_mode="percent",
|
||||
start_time=time_range[0],
|
||||
end_time=time_range[1],
|
||||
model_type = ModelType(model_type)
|
||||
preset = SIMPLE_PRESETS.get(model_type)
|
||||
if preset is None:
|
||||
errstr = f"Unknown model type {model_type!s}"
|
||||
raise ValueError(errstr)
|
||||
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)
|
||||
|
||||
|
||||
__all__ = ("ApplyMSWMSAAttention", "ApplyMSWMSAAttentionSimple")
|
||||
|
||||
+585
-267
File diff suppressed because it is too large
Load Diff
+282
-42
@@ -1,44 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import importlib
|
||||
import itertools
|
||||
import logging
|
||||
import math
|
||||
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:
|
||||
# Compatibility workaround for pre-3.11 Python versions.
|
||||
from enum import Enum
|
||||
|
||||
class StrEnum(str, Enum):
|
||||
@staticmethod
|
||||
def _generate_next_value_(name: str, *_unused: list) -> str:
|
||||
return name.lower()
|
||||
|
||||
def __str__(self) -> str:
|
||||
return str(self.value)
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
UPSCALE_METHODS = ("bicubic", "bislerp", "bilinear", "nearest-exact", "nearest", "area")
|
||||
|
||||
|
||||
def parse_blocks(name: str, s: str) -> set:
|
||||
vals = (rawval.strip() for rawval in s.split(","))
|
||||
class TimeMode(StrEnum):
|
||||
PERCENT = "percent"
|
||||
TIMESTEP = "timestep"
|
||||
SIGMA = "sigma"
|
||||
|
||||
|
||||
class ModelType(StrEnum):
|
||||
SD15 = "SD15"
|
||||
SDXL = "SDXL"
|
||||
|
||||
|
||||
def parse_blocks(name: str, val: str | Sequence[int]) -> set[tuple[str, int]]:
|
||||
if isinstance(val, (tuple, list)):
|
||||
# Handle a sequence passed in via YAML parameters.
|
||||
if not all(isinstance(item, int) and item >= 0 for item in val):
|
||||
raise ValueError(
|
||||
"Bad blocks definition, must be comma separated string or sequence of positive int",
|
||||
)
|
||||
return {(name, item) for item in val}
|
||||
vals = (rawval.strip() for rawval in val.split(","))
|
||||
return {(name, int(val.strip())) for val in vals if val}
|
||||
|
||||
|
||||
def convert_time(
|
||||
ms: object,
|
||||
time_mode: str,
|
||||
time_mode: TimeMode,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
) -> tuple:
|
||||
if time_mode == "sigma":
|
||||
) -> tuple[float, float]:
|
||||
if time_mode == TimeMode.SIGMA:
|
||||
return (start_time, end_time)
|
||||
if time_mode in {"percent", "timestep"}:
|
||||
if time_mode == "timestep":
|
||||
start_time = 1.0 - (start_time / 999.0)
|
||||
end_time = 1.0 - (end_time / 999.0)
|
||||
else:
|
||||
if start_time > 1.0 or start_time < 0.0:
|
||||
raise ValueError(
|
||||
"invalid value for start percent",
|
||||
)
|
||||
if end_time > 1.0 or end_time < 0.0:
|
||||
raise ValueError(
|
||||
"invalid value for end percent",
|
||||
)
|
||||
return (ms.percent_to_sigma(start_time), ms.percent_to_sigma(end_time))
|
||||
if time_mode == TimeMode.TIMESTEP:
|
||||
start_time = 1.0 - (start_time / 999.0)
|
||||
end_time = 1.0 - (end_time / 999.0)
|
||||
else:
|
||||
if start_time > 1.0 or start_time < 0.0:
|
||||
raise ValueError(
|
||||
"invalid value for start percent",
|
||||
)
|
||||
if end_time > 1.0 or end_time < 0.0:
|
||||
raise ValueError(
|
||||
"invalid value for end percent",
|
||||
)
|
||||
return (
|
||||
round(ms.percent_to_sigma(start_time), 4),
|
||||
round(ms.percent_to_sigma(end_time), 4),
|
||||
)
|
||||
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)
|
||||
@@ -56,39 +105,230 @@ def check_time(time_arg: dict | float, start_sigma: float, end_sigma: float) ->
|
||||
return sigma <= start_sigma and sigma >= end_sigma
|
||||
|
||||
|
||||
try:
|
||||
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
|
||||
bleh_latentutils = getattr(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:
|
||||
__block_to_num_map = {"input": 0, "middle": 1, "output": 2}
|
||||
|
||||
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
|
||||
except (ImportError, NotImplementedError):
|
||||
def block_to_num(block_type: str, block_id: int) -> tuple[int, int]:
|
||||
type_id = __block_to_num_map.get(block_type)
|
||||
if type_id is None:
|
||||
errstr = f"Got unexpected block type {block_type}!"
|
||||
raise ValueError(errstr)
|
||||
return (type_id, block_id)
|
||||
|
||||
def scale_samples(
|
||||
samples,
|
||||
width,
|
||||
height,
|
||||
mode="bicubic",
|
||||
sigma=None, # noqa: ARG001
|
||||
|
||||
# Naive and totally inaccurate way to factorize target_res into rescaled integer width/height
|
||||
def rescale_size(
|
||||
width: int,
|
||||
height: int,
|
||||
target_res: int,
|
||||
*,
|
||||
tolerance=1,
|
||||
) -> tuple[int, int]:
|
||||
tolerance = min(target_res, tolerance)
|
||||
|
||||
def get_neighbors(num: float):
|
||||
if num < 1:
|
||||
return None
|
||||
numi = int(num)
|
||||
return tuple(
|
||||
numi + adj
|
||||
for adj in sorted(
|
||||
range(
|
||||
-min(numi - 1, tolerance),
|
||||
tolerance + 1 + math.ceil(num - numi),
|
||||
),
|
||||
key=abs,
|
||||
)
|
||||
)
|
||||
|
||||
scale = math.sqrt(height * width / target_res)
|
||||
height_scaled, width_scaled = height / scale, width / scale
|
||||
height_rounded = get_neighbors(height_scaled)
|
||||
width_rounded = get_neighbors(width_scaled)
|
||||
for h, w in itertools.zip_longest(height_rounded, width_rounded):
|
||||
h_adj = target_res / w if w is not None else 0.1
|
||||
if h_adj % 1 == 0:
|
||||
return (w, int(h_adj))
|
||||
if h is None:
|
||||
continue
|
||||
w_adj = target_res / h
|
||||
if w_adj % 1 == 0:
|
||||
return (int(w_adj), h)
|
||||
msg = f"Can't rescale {width} and {height} to fit {target_res}"
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
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
|
||||
if isinstance(
|
||||
latent_format,
|
||||
(latent_formats.SDXL, latent_formats.SDXL_Playground_2_5),
|
||||
):
|
||||
if mode == "bislerp":
|
||||
return bislerp(samples, width, height)
|
||||
return torchf.interpolate(samples, size=(height, width), mode=mode)
|
||||
return ModelType.SDXL
|
||||
return None
|
||||
|
||||
|
||||
def sigma_to_pct(ms, sigma):
|
||||
return (1.0 - (ms.timestep(sigma).detach().cpu() / 999.0)).clamp(0.0, 1.0).item()
|
||||
|
||||
|
||||
def fade_scale(
|
||||
pct,
|
||||
start_pct=0.0,
|
||||
end_pct=1.0,
|
||||
fade_start=1.0,
|
||||
fade_cap=0.0,
|
||||
):
|
||||
if not (start_pct <= pct <= end_pct) or start_pct > end_pct:
|
||||
return 0.0
|
||||
if pct < fade_start:
|
||||
return 1.0
|
||||
scaling_pct = 1.0 - ((pct - fade_start) / (end_pct - fade_start))
|
||||
return max(fade_cap, scaling_pct)
|
||||
|
||||
|
||||
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:
|
||||
return
|
||||
bleh_version = getattr(ext_bleh, "BLEH_VERSION", -1)
|
||||
UPSCALE_METHODS = bleh_latentutils.UPSCALE_METHODS
|
||||
if bleh_version >= 0:
|
||||
scale_samples = bleh_latentutils.scale_samples
|
||||
return
|
||||
|
||||
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",
|
||||
"check_time",
|
||||
"convert_time",
|
||||
"get_sigma",
|
||||
"guess_model_type",
|
||||
"parse_blocks",
|
||||
"rescale_size",
|
||||
"scale_samples",
|
||||
)
|
||||
|
||||
+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.0"
|
||||
version = "0.8.5"
|
||||
license = { file = "LICENSE" }
|
||||
|
||||
[project.urls]
|
||||
|
||||
Reference in New Issue
Block a user