diff --git a/README.md b/README.md index 0166b25..9edac5b 100644 --- a/README.md +++ b/README.md @@ -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. +
+ +YAML parameters + +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. + +
+ ### `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. +
+ +YAML parameters + +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. + +
+ ## Credits Code based on the HiDiffusion original implementation: https://github.com/megvii-research/HiDiffusion diff --git a/changelog.md b/changelog.md index 18a777a..9c6c127 100644 --- a/changelog.md +++ b/changelog.md @@ -2,6 +2,23 @@ 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. diff --git a/py/msw_msa_attention.py b/py/msw_msa_attention.py index 3683f12..3eea234 100644 --- a/py/msw_msa_attention.py +++ b/py/msw_msa_attention.py @@ -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"" + + 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") diff --git a/py/raunet.py b/py/raunet.py index 2ace7b0..29e330f 100644 --- a/py/raunet.py +++ b/py/raunet.py @@ -1,20 +1,27 @@ from __future__ import annotations +import itertools import logging import os import sys -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: @@ -22,36 +29,219 @@ if TYPE_CHECKING: F = torch.nn.functional +CA_DOWNSCALE_METHODS = ( + ("avg_pool2d", "adaptive_avg_pool2d", *UPSCALE_METHODS) + if "adaptive_avg_pool2d" not in UPSCALE_METHODS + else ("avg_pool2d", *UPSCALE_METHODS) +) -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) - def __str__(self): - return f"" + @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 @@ -127,7 +317,7 @@ class HDState: logging.info("** jankhidiffusion: Reverted FreeU_Advanced patch") -GLOBAL_STATE = HDState() +GLOBAL_STATE: State = State() class HDForward: @@ -142,7 +332,7 @@ class HDForward: def __init__( self, orig_block: object, - hdconfig: HDConfig, + config: Config, block_index: int, is_up: bool, ): @@ -153,7 +343,7 @@ class HDForward: while isinstance(orig_forward, HDForward): orig_forward = orig_forward.orig_forward self.orig_forward = orig_forward - self.hdconfig = hdconfig + self.config = config self.block_index = block_index self.forward = self.forward_upsample if is_up else self.forward_downsample @@ -165,14 +355,14 @@ class HDForward: x: torch.Tensor, output_shape: None | tuple = None, ) -> torch.Tensor: - hdconfig = self.hdconfig + 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 hdconfig.check({ - "sigmas": hdconfig.curr_sigma, + or not config.check({ + "sigmas": config.curr_sigma, "block": ("output", block_index), }) ): @@ -183,35 +373,40 @@ class HDForward: if output_shape is not None else (x.shape[2] * 4, x.shape[3] * 4) ) - if hdconfig.two_stage_upscale_mode != "disabled": + 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=hdconfig.two_stage_upscale_mode, - sigma=hdconfig.curr_sigma, + mode=config.two_stage_upscale_mode, + sigma=config.curr_sigma, ) x = scale_samples( x, shape[1], shape[0], - mode=hdconfig.upscale_mode, - sigma=hdconfig.curr_sigma, + mode=config.upscale_mode, + sigma=config.curr_sigma, + ) + return config.maybe_multiply( + orig_block.conv(x), + config.post_upscale_multiplier, + post=True, ) - return orig_block.conv(x) def forward_downsample( self, x: torch.Tensor, ) -> torch.Tensor: - hdconfig = self.hdconfig + 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 hdconfig.check({ - "sigmas": hdconfig.curr_sigma, + or not config.check({ + "sigmas": config.curr_sigma, "block": ("input", block_index), }) ): @@ -246,7 +441,14 @@ class HDForward: for k in self.FORWARD_DOWNSAMPLE_COPY_OP_KEYS: setattr(tempop, k, getattr(orig_block.op, k)) - return tempop(x) + 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: @@ -270,19 +472,20 @@ class ApplyRAUNet: "STRING", { "default": "3", - "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", + "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. 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", + "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.", }, ), @@ -340,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": ( @@ -357,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": ( @@ -381,82 +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() - 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( @@ -464,29 +755,41 @@ 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}", ) @@ -505,7 +808,7 @@ class ApplyRAUNet: raise ValueError(error_message) # noqa: TRY004 model.add_object_patch( f"{block_name}.forward", - HDForward(block, hdconfig, block_index, block_type != "input"), + HDForward(block, config, block_index, block_type != "input"), ) GLOBAL_STATE.apply_patches() @@ -531,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": ( @@ -543,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": ( @@ -572,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") diff --git a/py/utils.py b/py/utils.py index 40c02e0..2aaebc2 100644 --- a/py/utils.py +++ b/py/utils.py @@ -1,43 +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 ( - round(ms.percent_to_sigma(start_time), 4), - round(ms.percent_to_sigma(end_time), 4), - ) + 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") @@ -59,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) @@ -92,6 +212,8 @@ __all__ = ( "check_time", "convert_time", "get_sigma", + "guess_model_type", "parse_blocks", + "rescale_size", "scale_samples", ) diff --git a/pyproject.toml b/pyproject.toml index 7db1ace..161696e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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.3" +version = "0.8.4" license = { file = "LICENSE" } [project.urls] diff --git a/ruff.toml b/ruff.toml index 9f5e8af..7cd2146 100644 --- a/ruff.toml +++ b/ruff.toml @@ -31,6 +31,7 @@ ignore = [ "PLR0913", "PLR0915", "PLR2004", + "PLR6104", "T201", "TD001", "TD002",