Author SHA1 Message Date
blepping 78eb9a8447 Update changelog 2024-12-24 21:42:23 -07:00
blepping a54e89efa5 Remove unused import 2024-12-24 21:39:51 -07:00
blepping f6449b0ab0 Integration refactor part 2 2024-12-22 10:37:28 -07:00
blepping 0ad10230cf Different approach to integrating external modules
Other internal cleanups
2024-12-20 07:01:39 -07:00
blepping 4e66c60163 Merge pull request #26 from blepping/attention_improvements
October update
2024-10-14 04:34:14 -06:00
blepping 71435fc0ec Update changelog 2024-10-14 04:29:27 -06:00
blepping 30ffe2d939 Make block in MSW-MSA attention warnings more user friendly 2024-10-12 08:21:56 -06:00
blepping debeccf722 Fix compatibility with pre-Python 3.11.
Use enums a bit more responsibly.

Make the verbose config field an integer verbosity level.
2024-10-12 05:50:25 -06:00
blepping e62bd8ea41 Allow applying MSW-MSA attention to attn2 (not a good idea though).
Move MSW-MSA attention scaling modes into advanced YAML options.
2024-10-10 17:10:44 -06:00
blepping 88dd9340aa Add advanced YAML parameters input to normal nodes.
Allow fading out CA downscale.

Use a pixel increment when scaling to ensure more compatible sizes.

Allow multiplying tensors on input/out for both RAUNet and MSW-MSA attention.

Add an auto model type for simple nodes and try to guess the model.

Allow CA scaling to use different factors for width and height.

Default CA downscaling method is now adaptive_avg_pool2d (and also add that scaling type).

CA input block scaling can now apply after the skip connection like Kohya Deep Shrink.

Various refactoring and cleanups.
2024-10-10 09:05:45 -06:00
blepping f723a994a1 Merge pull request #25 from pamparamm/main
Allow all resolutions for MSW-MSA. Change default values for SDXL
2024-09-30 07:50:14 -06:00
Pam eca9689c03 Add upscale mode selector for MSW MSA 2024-09-30 18:03:53 +05:00
Pam c79cc33ff9 Rewrite downsample_ratio calculation 2024-09-29 11:08:17 +05:00
Pam f17b2efe23 update msw msa
change SDXL default values
2024-09-27 22:19:13 +05:00
blepping 86afeb70f9 Bump version 2024-08-27 22:12:46 -06:00
blepping 2b3fadfaf2 Try to fix issue with RAUNet model patching changing seeds
Better tooltips
2024-08-27 22:07:02 -06:00
blepping 2b9c5b1c2e ComfyUI-GGUF compatibility for RAUNet 2024-08-25 16:46:34 -06:00
blepping 71f2ef42fd Bump version 2024-08-17 22:13:03 -06:00
blepping d4d3fd0ff3 Fix issue with controlnet workaround 2024-08-17 01:34:02 -06:00
blepping 64090c80b7 Version bump 2024-08-16 13:06:57 -06:00
blepping 4f48873f98 Make up/downsample block targeting more resiliant in RAUNet 2024-08-15 19:06:56 -06:00
blepping 922a400f6f Change publish workflow to trigger on release 2024-08-15 03:12:19 -06:00
blepping 7548ad6d07 Merge pull request #19 from blepping/comfyorg_publish
Set up Comfy Registry publishing
2024-08-15 02:57:11 -06:00
blepping 028a831031 Set up Comfy Registry publishing 2024-08-15 02:54:51 -06:00
blepping 4e8ef65a7e Slight consistency tweak for tooltip text. 2024-08-15 02:39:49 -06:00
blepping 3d33f3f7e7 Merge pull request #11 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-08-14 11:07:16 -06:00
blepping be31421715 Merge pull request #12 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-08-14 11:06:41 -06:00
blepping c1f7e80f78 Fix output blocks tooltip in ApplyRAUNet node 2024-08-14 10:56:44 -06:00
blepping 00e41dc6b1 Add tooltips metadata to nodes 2024-08-14 10:54:24 -06:00
blepping 4925c89a31 Merge pull request #18 from blepping/refactor
* Refactor RAUNet code to avoid monkeypatching Upsample/Downsample blocks (by pamparamm)
* Move two_stage_upscale toggle into two_stage_upscale_mode (by pamparamm)
* Refactor RAUNet code to avoid monkeypatching forward_timestep_embed
* Allow setting a downscale factor and mode for CA downsampling in advanced RAUNet node
* Make it so MSW-MSA attention failing due to size mismatches is a warning rather than hard error
2024-08-13 03:57:49 -06:00
haohaocreates 156ae752e1 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-22 19:18:35 -04:00
haohaocreates 89e2ce4c44 chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-22 19:18:31 -04:00
8 changed files with 1586 additions and 429 deletions
+16
View File
@@ -0,0 +1,16 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
release: { types: ["published"] }
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
personal_access_token: ${{ secrets.COMFYORG_REGISTRY_API_KEY }}
+120
View File
@@ -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
+25
View File
@@ -2,6 +2,31 @@
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.
_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.
+458 -119
View File
@@ -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,24 +87,170 @@ 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"
CATEGORY = "model_patches/unet"
DESCRIPTION = "This node applies an attention patch which _may_ slightly improve quality especially when generating at high resolutions. It is a large performance increase on SD1.x, may improve performance on SDXL. This is the advanced version of the node with more parameters, use ApplyMSWMSAAttentionSimple if this seems too complex. NOTE: Only supports SD1.x, SD2.x and SDXL."
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"input_blocks": ("STRING", {"default": "1,2"}),
"middle_blocks": ("STRING", {"default": ""}),
"output_blocks": ("STRING", {"default": "9,10,11"}),
"input_blocks": (
"STRING",
{
"default": "1,2",
"tooltip": "Comma-separated list of input blocks to patch. Default is for SD1.x, you can try 4,5 for SDXL",
},
),
"middle_blocks": (
"STRING",
{
"default": "",
"tooltip": "Comma-separated list of middle blocks to patch. Generally not recommended.",
},
),
"output_blocks": (
"STRING",
{
"default": "9,10,11",
"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.",
},
),
"start_time": (
"FLOAT",
@@ -54,6 +260,7 @@ class ApplyMSWMSAAttention:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time the MSW-MSA attention effect starts applying - value is inclusive.",
},
),
"end_time": (
@@ -64,9 +271,26 @@ class ApplyMSWMSAAttention:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time the MSW-MSA attention effect ends - value is inclusive.",
},
),
"model": (
"MODEL",
{
"tooltip": "Model to patch with the MSW-MSA attention effect.",
},
),
},
"optional": {
"yaml_parameters": (
"STRING",
{
"tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. This input can be converted into a multiline text widget. See main README for possible options. Note: When specifying paramaters this way, there is very little error checking.",
"dynamicPrompts": False,
"multiline": True,
"defaultInput": True,
},
),
"model": ("MODEL",),
},
}
@@ -75,53 +299,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,
@@ -129,14 +404,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)
@@ -148,131 +426,192 @@ 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"
CATEGORY = "model_patches/unet"
DESCRIPTION = "This node applies an attention patch which _may_ slightly improve quality especially when generating at high resolutions. It is a large performance increase on SD1.x, may improve performance on SDXL. This is the simplified version of the node with less parameters. Use ApplyMSWMSAAttention if you require more control. NOTE: Only supports SD1.x, SD2.x and SDXL."
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"model_type": (("SD15", "SDXL"),),
"model": ("MODEL",),
"model_type": (
("auto", "SD15", "SDXL"),
{
"tooltip": "Model type being patched. Generally safe to leave on auto. Choose SD15 for SD 1.4, SD 2.x.",
},
),
"model": (
"MODEL",
{
"tooltip": "Model to patch with the MSW-MSA attention effect.",
},
),
},
}
@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")
+670 -268
View File
File diff suppressed because it is too large Load Diff
+282 -42
View File
@@ -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",
)
+14
View File
@@ -0,0 +1,14 @@
[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.5"
license = { file = "LICENSE" }
[project.urls]
Repository = "https://github.com/blepping/comfyui_jankhidiffusion"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "blepping"
DisplayName = "comfyui_jankhidiffusion"
Icon = ""
+1
View File
@@ -31,6 +31,7 @@ ignore = [
"PLR0913",
"PLR0915",
"PLR2004",
"PLR6104",
"T201",
"TD001",
"TD002",