15 Commits
Author SHA1 Message Date
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
7 changed files with 1226 additions and 372 deletions
+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
+21
View File
@@ -2,6 +2,27 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 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.
+397 -111
View File
@@ -1,16 +1,65 @@
from __future__ import annotations
import itertools
import logging
from typing import TYPE_CHECKING, NamedTuple
import math
from time import time
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
from .utils import *
from .utils import (
UPSCALE_METHODS,
ModelType,
StrEnum,
TimeMode,
block_to_num,
check_time,
convert_time,
get_sigma,
guess_model_type,
parse_blocks,
rescale_size,
scale_samples,
)
F = torch.nn.functional
if TYPE_CHECKING:
import comfy
SCALE_METHODS = ("disabled", "skip", *UPSCALE_METHODS)
REVERSE_SCALE_METHODS = UPSCALE_METHODS
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
width: int
@@ -27,6 +76,132 @@ class ShiftSize(WindowSize):
pass
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
):
logging.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:
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
@@ -56,16 +231,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 +270,17 @@ class ApplyMSWMSAAttention:
},
),
},
"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,
},
),
},
}
# reference: https://github.com/microsoft/Swin-Transformer
@@ -105,53 +288,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 +393,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,88 +415,142 @@ 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: None | str = 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:
logging.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:
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)
)
except (RuntimeError, ValueError) 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}",
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}",
)
window_args = None
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,)
@@ -275,9 +566,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 +583,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)
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)
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],
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")
+542 -241
View File
@@ -1,57 +1,247 @@
from __future__ import annotations
import itertools
import logging
import os
import sys
from functools import partial
from typing import TYPE_CHECKING
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, NamedTuple
import torch
from comfy.ldm.modules.diffusionmodules import openaimodel
from .utils import (
UPSCALE_METHODS,
ModelType,
TimeMode,
check_time,
convert_time,
fade_scale,
get_sigma,
guess_model_type,
parse_blocks,
scale_samples,
sigma_to_pct,
)
if TYPE_CHECKING:
from typing import Callable
from comfy.model_patcher import ModelPatcher
F = torch.nn.functional
CA_DOWNSCALE_METHODS = (
("avg_pool2d", "adaptive_avg_pool2d", *UPSCALE_METHODS)
if "adaptive_avg_pool2d" not in UPSCALE_METHODS
else ("avg_pool2d", *UPSCALE_METHODS)
)
class HDConfig:
def __init__(
self,
start_sigma: float,
end_sigma: float,
use_blocks: dict,
upscale_mode: str,
two_stage_upscale_mode: str,
):
self.curr_sigma: None | float = None
self.start_sigma = start_sigma
self.end_sigma = end_sigma
self.use_blocks = use_blocks
self.upscale_mode = upscale_mode
self.two_stage_upscale_mode = two_stage_upscale_mode
def check(self, topts: dict) -> bool:
if not isinstance(topts, dict) or topts.get("block") not in self.use_blocks:
class Preset(NamedTuple):
input_blocks: str = ""
output_blocks: str = ""
time_mode: TimeMode = TimeMode.PERCENT
start_time: float = 1.0
end_time: float = 1.0
upscale_mode: str = "bicubic"
ca_start_time: float = 1.0
ca_end_time: float = 1.0
ca_downscale_factor: float = 2.0
ca_input_blocks: str = ""
ca_output_blocks: str = ""
ca_upscale_mode: str = "bicubic"
ca_downscale_mode: str = "avg_pool2d"
ca_input_after_skip_mode: bool = False
two_stage_upscale_mode: str = "disabled"
def _pretty_blocks(self, *, ca: bool = False) -> str:
if ca:
blocks = (
self.ca_input_blocks,
self.ca_output_blocks,
)
else:
blocks = (self.input_blocks, set(), self.output_blocks)
return " / ".join(b or "none" for b in blocks)
@property
def pretty_blocks(self) -> str:
return self._pretty_blocks(ca=False)
@property
def ca_pretty_blocks(self) -> str:
return self._pretty_blocks(ca=True)
@property
def as_dict(self):
return {k: getattr(self, k) for k in self._fields}
def edited(self, **kwargs: dict) -> NamedTuple:
kwargs = self.as_dict | kwargs
return self.__class__(**kwargs)
SD15_PRESET = Preset(
input_blocks="3",
output_blocks="8",
ca_input_blocks="1",
ca_output_blocks="11",
)
SDXL_PRESET = Preset(
input_blocks="3",
output_blocks="5",
ca_input_blocks="4",
ca_output_blocks="5",
)
SIMPLE_PRESETS = {
"SD15_low": SD15_PRESET.edited(
start_time=0.0,
end_time=0.4,
),
"SD15_high": SD15_PRESET.edited(
start_time=0.0,
end_time=0.5,
ca_start_time=0.0,
ca_end_time=0.35,
),
"SD15_ultra": SD15_PRESET.edited(
start_time=0.0,
end_time=0.6,
ca_start_time=0.0,
ca_end_time=0.45,
),
"SDXL_low": SDXL_PRESET.edited(), # ???
"SDXL_high": SDXL_PRESET.edited(
ca_start_time=0.0,
ca_end_time=0.5,
),
"SDXL_ultra": SDXL_PRESET.edited(
start_time=0.0,
end_time=0.45,
ca_start_time=0.0,
ca_end_time=0.6,
),
}
@dataclass
class Config:
start_sigma: float
end_sigma: float
ca_start_sigma: float
ca_end_sigma: float
use_blocks: set
ca_use_blocks: set
upscale_mode: str = "bicubic"
two_stage_upscale_mode: str = "disabled"
ca_upscale_mode: str = "bicubic"
ca_downscale_mode: str = "adaptive_avg_pool2d"
ca_downscale_factor: float = 2.0
ca_downscale_factor_w: None | float = None
# Patches the input block after the skip connection.
ca_input_after_skip_mode: bool = False
ca_avg_pool2d_ceil_mode: bool = True
# Hack for ComfyUI-bleh latent effects. # noqa: FIX004
ca_output_sigma_hack: bool = True
# Scaling on the tensors going in/out of scaling.
pre_upscale_multiplier: float = 1.0
post_upscale_multiplier: float = 1.0
pre_downscale_multiplier: float = 1.0
post_downscale_multiplier: float = 1.0
ca_pre_upscale_multiplier: float = 1.0
ca_post_upscale_multiplier: float = 1.0
ca_pre_downscale_multiplier: float = 1.0
ca_post_downscale_multiplier: float = 1.0
# Allows fading out the scale effect starting from this time.
ca_fadeout_start_sigma: None | float = None
# Maximum fadeout, as a percentage of the total scale effect.
ca_fadeout_cap: float = 0.0
ca_latent_pixel_increment: int | float = 8
verbose: int = 0
curr_sigma: None | float = None
@classmethod
def build(
cls,
ms: object,
*,
input_blocks: str | list[int],
output_blocks: str | list[int],
time_mode: str | TimeMode,
start_time: float,
end_time: float,
ca_start_time: float,
ca_end_time: float,
ca_input_blocks: str | list[int],
ca_output_blocks: str | list[int],
ca_fadeout_start_time: None | float = None,
**kwargs: dict,
) -> object:
time_mode: TimeMode = TimeMode(time_mode)
start_sigma, end_sigma = convert_time(ms, time_mode, start_time, end_time)
ca_start_sigma, ca_end_sigma = convert_time(
ms,
time_mode,
ca_start_time,
ca_end_time,
)
if ca_fadeout_start_time is not None:
ca_fadeout_start_sigma = convert_time(
ms,
time_mode,
ca_fadeout_start_time,
ca_fadeout_start_time,
)[0]
else:
ca_fadeout_start_sigma = None
input_blocks, output_blocks = itertools.starmap(
parse_blocks,
(
("input", input_blocks),
("output", output_blocks),
),
)
ca_input_blocks, ca_output_blocks = itertools.starmap(
parse_blocks,
(
("input", ca_input_blocks),
("output", ca_output_blocks),
),
)
return cls(
start_sigma=start_sigma,
end_sigma=end_sigma,
ca_start_sigma=ca_start_sigma,
ca_end_sigma=ca_end_sigma,
ca_fadeout_start_sigma=ca_fadeout_start_sigma,
use_blocks=input_blocks | output_blocks,
ca_use_blocks=ca_input_blocks | ca_output_blocks,
**kwargs,
)
def check(self, topts: dict, *, ca=False) -> bool:
start_sigma, end_sigma, use_blocks = (
(self.ca_start_sigma, self.ca_end_sigma, self.ca_use_blocks)
if ca
else (self.start_sigma, self.end_sigma, self.use_blocks)
)
if not isinstance(topts, dict) or topts.get("block") not in use_blocks:
return False
return check_time(topts, self.start_sigma, self.end_sigma)
return check_time(topts, start_sigma, end_sigma)
@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
GLOBAL_STATE: HDState
class HDState:
class State:
def __init__(self):
self.no_controlnet_workaround = (
"JANKHIDIFFUSION_NO_CONTROLNET_WORKAROUND" in os.environ
@@ -61,9 +251,8 @@ class HDState:
self.orig_apply_control = openaimodel.apply_control
self.orig_fua_apply_control = None
@classmethod
def hd_apply_control(
cls,
self,
h: torch.Tensor,
control: None | dict,
name: str,
@@ -78,7 +267,7 @@ class HDState:
logging.info(
f"* jankhidiffusion: Scaling controlnet conditioning: {ctrl.shape[-2:]} -> {h.shape[-2:]}",
)
ctrl = F.interpolate(ctrl, size=h.shape[-2:], **cls.controlnet_scale_args)
ctrl = F.interpolate(ctrl, size=h.shape[-2:], **self.controlnet_scale_args)
h += ctrl
return h
@@ -128,90 +317,138 @@ class HDState:
logging.info("** jankhidiffusion: Reverted FreeU_Advanced patch")
GLOBAL_STATE = HDState()
GLOBAL_STATE: State = State()
def forward_upsample( # noqa: PLR0917
block_index: int,
model: object,
orig_forward: Callable,
hdconfig: HDConfig,
x: torch.Tensor,
output_shape: None | tuple = None,
) -> torch.Tensor:
if (
model.dims == 3
or not model.use_conv
or not hdconfig.check({
"sigmas": hdconfig.curr_sigma,
"block": ("output", block_index),
})
):
return orig_forward(x, output_shape=output_shape)
shape = (
output_shape[2:4]
if output_shape is not None
else (x.shape[2] * 4, x.shape[3] * 4)
class HDForward:
FORWARD_DOWNSAMPLE_COPY_OP_KEYS = (
"comfy_cast_weights",
"weight_function",
"bias_function",
"weight",
"bias",
)
if hdconfig.two_stage_upscale_mode != "disabled":
def __init__(
self,
orig_block: object,
config: Config,
block_index: int,
is_up: bool,
):
self.orig_block = orig_block
orig_forward = orig_block.forward
# This is weird but apparently when we patch the model, the previous object patches
# may still exist, so we have to make sure we get the _real_ original forward function.
while isinstance(orig_forward, HDForward):
orig_forward = orig_forward.orig_forward
self.orig_forward = orig_forward
self.config = config
self.block_index = block_index
self.forward = self.forward_upsample if is_up else self.forward_downsample
def __call__(self, *args: list, **kwargs: dict) -> torch.Tensor:
return self.forward(*args, **kwargs)
def forward_upsample(
self,
x: torch.Tensor,
output_shape: None | tuple = None,
) -> torch.Tensor:
config = self.config
orig_block = self.orig_block
block_index = self.block_index
if (
orig_block.dims == 3
or not orig_block.use_conv
or not config.check({
"sigmas": config.curr_sigma,
"block": ("output", block_index),
})
):
return self.orig_forward(x, output_shape=output_shape)
shape = (
output_shape[2:4]
if output_shape is not None
else (x.shape[2] * 4, x.shape[3] * 4)
)
x = config.maybe_multiply(x, config.pre_upscale_multiplier)
if config.two_stage_upscale_mode != "disabled":
x = scale_samples(
x,
shape[1] // 2,
shape[0] // 2,
mode=config.two_stage_upscale_mode,
sigma=config.curr_sigma,
)
x = scale_samples(
x,
shape[1] // 2,
shape[0] // 2,
mode=hdconfig.two_stage_upscale_mode,
sigma=hdconfig.curr_sigma,
shape[1],
shape[0],
mode=config.upscale_mode,
sigma=config.curr_sigma,
)
return config.maybe_multiply(
orig_block.conv(x),
config.post_upscale_multiplier,
post=True,
)
x = scale_samples(
x,
shape[1],
shape[0],
mode=hdconfig.upscale_mode,
sigma=hdconfig.curr_sigma,
)
return model.conv(x)
def forward_downsample(
self,
x: torch.Tensor,
) -> torch.Tensor:
config = self.config
orig_block = self.orig_block
block_index = self.block_index
if (
orig_block.dims == 3
or not orig_block.use_conv
or not config.check({
"sigmas": config.curr_sigma,
"block": ("input", block_index),
})
):
return self.orig_forward(x)
FORWARD_DOWNSAMPLE_COPY_OP_KEYS = (
"comfy_cast_weights",
"weight_function",
"bias_function",
"weight",
"bias",
)
tempop = openaimodel.ops.conv_nd(
orig_block.dims,
orig_block.channels,
orig_block.out_channels,
3, # kernel size
stride=(4, 4),
padding=(2, 2),
dilation=(2, 2),
dtype=x.dtype,
device=x.device,
)
if (
orig_block.op.__class__.__base__ is not None
and orig_block.op.__class__.__base__.__name__ == "GGMLLayer"
):
# Workaround for GGML quantized Downsample blocks.
if not hasattr(orig_block.op, "get_weights"):
errstr = f"Cannot handle downsample block {block_index} which appears to be GGUF quantized but has no get_weights method!"
raise RuntimeError(errstr)
tempop.comfy_cast_weights = True
tempop.weight, tempop.bias = (
torch.nn.Parameter(p).to(device=x.device)
for p in orig_block.op.get_weights(x.dtype)
)
return tempop(x)
def forward_downsample(
block_index: int,
model: object,
orig_forward: Callable,
hdconfig: HDConfig,
x: torch.Tensor,
) -> torch.Tensor:
if (
model.dims == 3
or not model.use_conv
or not hdconfig.check({
"sigmas": hdconfig.curr_sigma,
"block": ("input", block_index),
})
):
return orig_forward(x)
tempop = openaimodel.ops.conv_nd(
model.dims,
model.channels,
model.out_channels,
3, # kernel size
stride=(4, 4),
padding=(2, 2),
dilation=(2, 2),
dtype=x.dtype,
device=x.device,
)
for k in FORWARD_DOWNSAMPLE_COPY_OP_KEYS:
setattr(tempop, k, getattr(model.op, k))
return tempop(x)
for k in self.FORWARD_DOWNSAMPLE_COPY_OP_KEYS:
setattr(tempop, k, getattr(orig_block.op, k))
x = config.maybe_multiply(x, config.pre_downscale_multiplier)
if config.pre_downscale_multiplier != 1.0:
x = x * config.pre_downscale_multiplier
return config.maybe_multiply(
tempop(x),
config.post_downscale_multiplier,
post=True,
)
class ApplyRAUNet:
@@ -235,19 +472,20 @@ class ApplyRAUNet:
"STRING",
{
"default": "3",
"tooltip": "Comma-separated list of input Downsample blocks. The default of 3 will work with SD1.x and SDXL.",
"tooltip": "Comma-separated list of input Downsample blocks. Default is for SD 1.5. The corresponding valid block from output_blocks must be set along with input.\nValid blocks for SD1.5: 3, 6, 9\nValid blocks for SDXL: 3, 6. Original Hidiffusion implementation uses 6 for SDXL.",
},
),
"output_blocks": (
"STRING",
{
"default": "8",
"tooltip": "Comma-separated list of output Upsample blocks. The default is for SD1.x, for SDXL use 5.",
"tooltip": "Comma-separated list of output Upsample blocks. Default is for SD 1.5. The corresponding valid block from input_blocks must be set along with output.\nValid blocks for SD1.5: 8, 5, 2\nValid blocks for SDXL: 5, 2. Original Hidiffusion implementation uses 2 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.",
},
),
@@ -305,14 +543,14 @@ class ApplyRAUNet:
"STRING",
{
"default": "4",
"tooltip": "Comma separated list of input cross-attention blocks. Default is for SD1.x, for SDXL you can try using 2 (or just disable it).",
"tooltip": "Comma separated list of input cross-attention blocks. Default is for SD1.x, for SDXL you can try using 5 (or just disable it).",
},
),
"ca_output_blocks": (
"STRING",
{
"default": "8",
"tooltip": "Comma-separated list of output cross-attention blocks. Default is for SD1.x, for SDXL you can try using 7 (or just disable it).",
"tooltip": "Comma-separated list of output cross-attention blocks. Default is for SD1.x, for SDXL you can try using 4 (or just disable it).",
},
),
"ca_upscale_mode": (
@@ -322,10 +560,10 @@ class ApplyRAUNet:
},
),
"ca_downscale_mode": (
("avg_pool2d", *UPSCALE_METHODS),
CA_DOWNSCALE_METHODS,
{
"default": "avg_pool2d",
"tooltip": "Mode used when downscaling latents in output cross-attention blocks (use avg_pool2d for normal Hidiffusion behavior).",
"default": "adaptive_avg_pool2d",
"tooltip": "Mode used when downscaling latents in output cross-attention blocks (use avg_pool2d for normal Hidiffusion behavior). adaptive_avg_pool2d should be the same and also supports fractional scales.",
},
),
"ca_downscale_factor": (
@@ -346,83 +584,170 @@ class ApplyRAUNet:
},
),
},
"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,
},
),
},
}
@classmethod
def patch(
def patch( # noqa: PLR0914
cls,
*,
model: ModelPatcher,
input_blocks: str,
output_blocks: str,
time_mode: str,
start_time: float,
end_time: float,
upscale_mode: str,
ca_start_time: float,
ca_end_time: float,
ca_input_blocks: str,
ca_output_blocks: str,
ca_upscale_mode: str,
ca_downscale_mode: str = "avg_pool2d",
ca_downscale_factor: float = 2.0,
two_stage_upscale_mode: str = "disabled",
yaml_parameters: None | str = None,
**kwargs: dict[str, Any],
) -> tuple[ModelPatcher]:
if ca_downscale_mode == "avg_pool2d" and not ca_downscale_factor.is_integer():
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(
"RAUNet: yaml_parameters must either be null or an object",
)
else:
kwargs |= extra_params
ms = model.get_model_object("model_sampling")
config = Config.build(ms, **kwargs)
if config.ca_downscale_mode == "avg_pool2d" and (
not config.ca_downscale_factor.is_integer()
or not (
config.ca_downscale_factor_w is None
or config.ca_downscale_factor_w.is_integer()
)
):
raise ValueError(
"avg_pool2d downscale mode can only be used with integer downscale factors",
)
use_blocks = parse_blocks("output", output_blocks)
use_blocks |= parse_blocks("input", input_blocks)
ca_use_blocks = parse_blocks("output", ca_output_blocks)
have_ca_output_blocks = len(ca_use_blocks) > 0
ca_use_blocks |= parse_blocks("input", ca_input_blocks)
if config.verbose:
logging.info(f"** jankhidiffusion: RAUNet: Using config: {config}")
have_ca_output_blocks = any(bt == "output" for (bt, _) in config.ca_use_blocks)
model = model.clone()
model.unpatch_model(device_to=model.model.device)
ms = model.get_model_object("model_sampling")
ca_start_sigma, ca_end_sigma = convert_time(
ms,
time_mode,
ca_start_time,
ca_end_time,
downscale_factor = config.ca_downscale_factor
downscale_factor_w = (
downscale_factor
if config.ca_downscale_factor_w is None
else config.ca_downscale_factor_w
)
hdconfig = HDConfig(
*convert_time(
if config.ca_fadeout_start_sigma is not None:
ca_start_pct = sigma_to_pct(
ms,
time_mode,
start_time,
end_time,
),
use_blocks,
upscale_mode,
two_stage_upscale_mode,
)
torch.tensor(config.ca_start_sigma, dtype=torch.float32),
)
ca_end_pct = sigma_to_pct(
ms,
torch.tensor(config.ca_end_sigma, dtype=torch.float32),
)
ca_fadeout_start_pct = sigma_to_pct(
ms,
torch.tensor(config.ca_fadeout_start_sigma, dtype=torch.float32),
)
else:
del ms
ca_fadeout_start_pct = None
ca_pixel_increment = max(1, config.ca_latent_pixel_increment)
def input_block_patch(h: torch.Tensor, extra_options: dict) -> torch.Tensor:
block_type, block_index = extra_options.get("block", ("unknown", -1))
_block_type, block_index = extra_options.get("block", ("unknown", -1))
if block_index == 0:
hdconfig.curr_sigma = get_sigma(extra_options)
if (block_type, block_index) not in ca_use_blocks or not check_time(
hdconfig.curr_sigma,
ca_start_sigma,
ca_end_sigma,
):
config.curr_sigma = get_sigma(extra_options)
if not config.check(extra_options, ca=True):
return h
if ca_downscale_mode == "avg_pool2d":
return F.avg_pool2d(
h,
kernel_size=(int(ca_downscale_factor), int(ca_downscale_factor)),
curr_downscale_factor, curr_downscale_factor_w = (
downscale_factor,
downscale_factor_w,
)
if ca_fadeout_start_pct is not None:
pct = sigma_to_pct(ms, extra_options["sigmas"].max())
scale_scale = fade_scale(
pct,
ca_start_pct,
ca_end_pct,
ca_fadeout_start_pct,
config.ca_fadeout_cap,
)
return scale_samples(
h,
max(1, int(h.shape[-1] // ca_downscale_factor)),
max(1, int(h.shape[-2] // ca_downscale_factor)),
mode=ca_downscale_mode,
sigma=hdconfig.curr_sigma,
if scale_scale <= 0.0:
return h
if scale_scale < 1.0:
curr_downscale_factor = curr_downscale_factor - (
curr_downscale_factor - 1.0
) * (1.0 - scale_scale)
curr_downscale_factor_w = curr_downscale_factor_w - (
curr_downscale_factor_w - 1.0
) * (1.0 - scale_scale)
# print(
# f"\n>>> scale_scale={scale_scale:0.4f}, down=({curr_downscale_factor:0.4f}, {curr_downscale_factor_w:0.4f})",
# )
height, width = h.shape[-2:]
target_h = int(
max(
ca_pixel_increment,
((height / ca_pixel_increment) // curr_downscale_factor)
* ca_pixel_increment,
),
)
target_w = int(
max(
ca_pixel_increment,
((width / ca_pixel_increment) // curr_downscale_factor_w)
* ca_pixel_increment,
),
)
# When downscaling, make sure not to overshoot the original size.
# When upscaling, don't undershoot the original size.
target_h = (
min(height, target_h)
if curr_downscale_factor >= 1
else max(height, target_h)
)
target_w = (
min(width, target_w)
if curr_downscale_factor_w >= 1
else max(width, target_w)
)
if (target_h, target_w) == h.shape[-2:]:
return h
h = config.maybe_multiply(h, config.ca_pre_downscale_multiplier)
if config.ca_downscale_mode == "avg_pool2d":
return config.maybe_multiply(
F.avg_pool2d(
h,
kernel_size=(
max(1, int(height // target_h)),
max(1, int(width // target_w)),
),
ceil_mode=config.ca_avg_pool2d_ceil_mode,
),
config.ca_post_downscale_multiplier,
post=True,
)
# print(f"\n>> h,w={(height, width)}, targets=({target_h}, {target_w})")
if config.ca_downscale_mode == "adaptive_avg_pool2d":
result = F.adaptive_avg_pool2d(h, (target_h, target_w))
else:
result = scale_samples(
h,
target_w,
target_h,
mode=config.ca_downscale_mode,
sigma=config.curr_sigma,
)
return config.maybe_multiply(
result,
config.ca_post_downscale_multiplier,
post=True,
)
def output_block_patch(
@@ -430,36 +755,48 @@ class ApplyRAUNet:
hsp: torch.Tensor,
extra_options: dict,
) -> torch.Tensor:
if extra_options.get("block") not in ca_use_blocks or not check_time(
hdconfig.curr_sigma,
ca_start_sigma,
ca_end_sigma,
if (
not config.check(extra_options, ca=True)
or h.shape[-2:] == hsp.shape[-2:]
):
return h, hsp
sigma = hdconfig.curr_sigma
sigma = config.curr_sigma
block = extra_options.get("block", ("", 0))[1]
if sigma is not None and (block < 3 or block > 6):
if (
sigma is not None
and config.ca_output_sigma_hack
and (block < 3 or block > 6)
):
sigma /= 16
return scale_samples(
h,
hsp.shape[-1],
hsp.shape[-2],
mode=ca_upscale_mode,
sigma=sigma,
h = config.maybe_multiply(h, config.ca_pre_upscale_multiplier)
return config.maybe_multiply(
scale_samples(
h,
hsp.shape[-1],
hsp.shape[-2],
mode=config.ca_upscale_mode,
sigma=sigma,
),
config.ca_post_upscale_multiplier,
post=True,
), hsp
model.set_model_input_block_patch(input_block_patch)
if config.ca_input_after_skip_mode:
model.set_model_input_block_patch_after_skip(input_block_patch)
else:
model.set_model_input_block_patch(input_block_patch)
if have_ca_output_blocks:
model.set_model_output_block_patch(output_block_patch)
for block_type, block_index in use_blocks:
for block_type, block_index in config.use_blocks:
main_block = model.get_model_object(
f"diffusion_model.{block_type}_blocks.{block_index}",
)
block_fun, expected_class = (
(forward_downsample, openaimodel.Downsample)
expected_class = (
openaimodel.Downsample
if block_type == "input"
else (forward_upsample, openaimodel.Upsample)
else openaimodel.Upsample
)
block_name = f"diffusion_model.{block_type}_blocks.{block_index}.{len(main_block) - 1}"
block = model.get_model_object(block_name)
@@ -471,7 +808,7 @@ class ApplyRAUNet:
raise ValueError(error_message) # noqa: TRY004
model.add_object_patch(
f"{block_name}.forward",
partial(block_fun, block_index, block, block.forward, hdconfig),
HDForward(block, config, block_index, block_type != "input"),
)
GLOBAL_STATE.apply_patches()
@@ -497,9 +834,9 @@ class ApplyRAUNetSimple:
},
),
"model_type": (
("SD15", "SDXL"),
("auto", "SD15", "SDXL"),
{
"tooltip": "Model type being patched. Choose SD15 for SD 1.4 or SD 2.x.",
"tooltip": "Model type being patched. Generally safe to leave on auto. Choose SD15 for SD 1.4 or SD 2.x.",
},
),
"res_mode": (
@@ -509,7 +846,7 @@ class ApplyRAUNetSimple:
"ultra (over 2048)",
),
{
"tooltip": "Resolution mode hint, does not have to correspond to the actual size.",
"tooltip": "Resolution mode hint, does not have to correspond to the actual size. Note: Choosing `low` with SDXL simply disables RAUNet as SDXL can natively generate at 1024x1024.",
},
),
"upscale_mode": (
@@ -538,69 +875,33 @@ class ApplyRAUNetSimple:
cls,
*,
model: ModelPatcher,
model_type: str,
model_type: str | ModelType,
res_mode: str,
upscale_mode: str,
ca_upscale_mode: str,
) -> tuple[ModelPatcher]:
if model_type == "auto":
model_type = guess_model_type(model)
if model_type not in ModelType:
raise RuntimeError("Unable to guess model type")
if upscale_mode == "default":
upscale_mode = "bicubic"
if ca_upscale_mode == "default":
ca_upscale_mode = "bicubic"
res = res_mode.split(" ", 1)[0]
if model_type == "SD15":
blocks = ("3", "8")
ca_blocks = ("1", "11")
time_range = (0.0, 0.6)
if res == "low":
time_range = (0.0, 0.4)
ca_time_range = (1.0, 0.0)
ca_blocks = ("", "")
elif res == "high":
time_range = (0.0, 0.5)
ca_time_range = (0.0, 0.35)
elif res == "ultra":
time_range = (0.0, 0.6)
ca_time_range = (0.0, 0.45)
else:
raise ValueError("Unknown res_mode")
elif model_type == "SDXL":
blocks = ("3", "5")
ca_blocks = ("4", "5")
if res == "low":
time_range = (1.0, 0.0)
ca_time_range = (1.0, 0.0)
ca_blocks = ("", "")
elif res == "high":
time_range = (0.0, 0.5)
ca_time_range = (1.0, 0.0)
elif res == "ultra":
time_range = (0.0, 0.6)
ca_time_range = (0.0, 0.45)
else:
raise ValueError("Unknown res_mode")
else:
raise ValueError("Unknown model type")
prettyblocks = " / ".join(b or "none" for b in blocks)
prettycablocks = " / ".join(b or "none" for b in ca_blocks)
logging.info(
f"** ApplyRAUNetSimple: Using preset {model_type} {res}: upscale {upscale_mode}, in/out blocks [{prettyblocks}], start/end percent {time_range[0]:.2}/{time_range[1]:.2} | CA upscale {ca_upscale_mode}, CA in/out blocks [{prettycablocks}], CA start/end percent {ca_time_range[0]:.2}/{ca_time_range[1]:.2}",
)
return ApplyRAUNet.patch(
model=model,
input_blocks=blocks[0],
output_blocks=blocks[1],
time_mode="percent",
start_time=time_range[0],
end_time=time_range[1],
preset_key = f"{model_type!s}_{res}"
preset = SIMPLE_PRESETS.get(preset_key)
if preset is None:
errstr = f"Unsupported model_type/res_mode combination {preset_key}"
raise ValueError(errstr)
preset = preset.edited(
upscale_mode=upscale_mode,
ca_start_time=ca_time_range[0],
ca_end_time=ca_time_range[1],
ca_input_blocks=ca_blocks[0],
ca_output_blocks=ca_blocks[1],
ca_upscale_mode=ca_upscale_mode,
)
logging.info(
f"** ApplyRAUNetSimple: Using preset {model_type!s} {res}: upscale {upscale_mode}, in/out blocks [{preset.pretty_blocks}], start/end percent {preset.start_time:.2}/{preset.end_time:.2} | CA upscale {preset.ca_upscale_mode}, CA in/out blocks [{preset.ca_pretty_blocks}], CA start/end percent {preset.ca_start_time:.2}/{preset.ca_end_time:.2}",
)
return ApplyRAUNet.patch(model=model, **preset.as_dict)
__all__ = ("ApplyRAUNet", "ApplyRAUNetSimple")
+144 -19
View File
@@ -1,40 +1,79 @@
from __future__ import annotations
import importlib
import itertools
import math
from typing import Sequence
import torch.nn.functional as torchf
from comfy import latent_formats
from comfy.utils import bislerp
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)
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")
@@ -56,6 +95,90 @@ def check_time(time_arg: dict | float, start_sigma: float, end_sigma: float) ->
return sigma <= start_sigma and sigma >= end_sigma
__block_to_num_map = {"input": 0, "middle": 1, "output": 2}
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)
# 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) -> None | ModelType:
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),
):
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)
try:
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
bleh_latentutils = getattr(bleh.py, "latent_utils", None)
@@ -89,6 +212,8 @@ __all__ = (
"check_time",
"convert_time",
"get_sigma",
"guess_model_type",
"parse_blocks",
"rescale_size",
"scale_samples",
)
+1 -1
View File
@@ -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.1"
version = "0.8.4"
license = { file = "LICENSE" }
[project.urls]
+1
View File
@@ -31,6 +31,7 @@ ignore = [
"PLR0913",
"PLR0915",
"PLR2004",
"PLR6104",
"T201",
"TD001",
"TD002",