18 Commits
Author SHA1 Message Date
blepping 86afeb70f9 Bump version 2024-08-27 22:12:46 -06:00
blepping 2b3fadfaf2 Try to fix issue with RAUNet model patching changing seeds
Better tooltips
2024-08-27 22:07:02 -06:00
blepping 2b9c5b1c2e ComfyUI-GGUF compatibility for RAUNet 2024-08-25 16:46:34 -06:00
blepping 71f2ef42fd Bump version 2024-08-17 22:13:03 -06:00
blepping d4d3fd0ff3 Fix issue with controlnet workaround 2024-08-17 01:34:02 -06:00
blepping 64090c80b7 Version bump 2024-08-16 13:06:57 -06:00
blepping 4f48873f98 Make up/downsample block targeting more resiliant in RAUNet 2024-08-15 19:06:56 -06:00
blepping 922a400f6f Change publish workflow to trigger on release 2024-08-15 03:12:19 -06:00
blepping 7548ad6d07 Merge pull request #19 from blepping/comfyorg_publish
Set up Comfy Registry publishing
2024-08-15 02:57:11 -06:00
blepping 028a831031 Set up Comfy Registry publishing 2024-08-15 02:54:51 -06:00
blepping 4e8ef65a7e Slight consistency tweak for tooltip text. 2024-08-15 02:39:49 -06:00
blepping 3d33f3f7e7 Merge pull request #11 from haohaocreates/publish
Add Github Action for Publishing to Comfy Registry
2024-08-14 11:07:16 -06:00
blepping be31421715 Merge pull request #12 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-08-14 11:06:41 -06:00
blepping c1f7e80f78 Fix output blocks tooltip in ApplyRAUNet node 2024-08-14 10:56:44 -06:00
blepping 00e41dc6b1 Add tooltips metadata to nodes 2024-08-14 10:54:24 -06:00
blepping 4925c89a31 Merge pull request #18 from blepping/refactor
* Refactor RAUNet code to avoid monkeypatching Upsample/Downsample blocks (by pamparamm)
* Move two_stage_upscale toggle into two_stage_upscale_mode (by pamparamm)
* Refactor RAUNet code to avoid monkeypatching forward_timestep_embed
* Allow setting a downscale factor and mode for CA downsampling in advanced RAUNet node
* Make it so MSW-MSA attention failing due to size mismatches is a warning rather than hard error
2024-08-13 03:57:49 -06:00
haohaocreates 156ae752e1 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-22 19:18:35 -04:00
haohaocreates 89e2ce4c44 chore(publish): Add Github Action for Publishing to Comfy Registry 2024-05-22 19:18:31 -04:00
6 changed files with 311 additions and 106 deletions
+16
View File
@@ -0,0 +1,16 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
release: { types: ["published"] }
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
personal_access_token: ${{ secrets.COMFYORG_REGISTRY_API_KEY }}
+4
View File
@@ -2,6 +2,10 @@
Note, only relatively significant changes to user-visible functionality will be included here. Most recent changes at the top.
## 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.
+48 -6
View File
@@ -29,22 +29,45 @@ class ShiftSize(WindowSize):
class ApplyMSWMSAAttention:
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
FUNCTION = "patch"
CATEGORY = "model_patches/unet"
DESCRIPTION = "This node applies an attention patch which _may_ slightly improve quality especially when generating at high resolutions. It is a large performance increase on SD1.x, may improve performance on SDXL. This is the advanced version of the node with more parameters, use ApplyMSWMSAAttentionSimple if this seems too complex. NOTE: Only supports SD1.x, SD2.x and SDXL."
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"input_blocks": ("STRING", {"default": "1,2"}),
"middle_blocks": ("STRING", {"default": ""}),
"output_blocks": ("STRING", {"default": "9,10,11"}),
"input_blocks": (
"STRING",
{
"default": "1,2",
"tooltip": "Comma-separated list of input blocks to patch. Default is for SD1.x, you can try 4,5 for SDXL",
},
),
"middle_blocks": (
"STRING",
{
"default": "",
"tooltip": "Comma-separated list of middle blocks to patch. Generally not recommended.",
},
),
"output_blocks": (
"STRING",
{
"default": "9,10,11",
"tooltip": "Comma-separated list of output blocks to patch. Default is for SD1.x, you can try 5,4 for SDXL",
},
),
"time_mode": (
(
"percent",
"timestep",
"sigma",
),
{
"tooltip": "Time mode controls how to interpret the values in start_time and end_time.",
},
),
"start_time": (
"FLOAT",
@@ -54,6 +77,7 @@ class ApplyMSWMSAAttention:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time the MSW-MSA attention effect starts applying - value is inclusive.",
},
),
"end_time": (
@@ -64,9 +88,15 @@ class ApplyMSWMSAAttention:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time the MSW-MSA attention effect ends - value is inclusive.",
},
),
"model": (
"MODEL",
{
"tooltip": "Model to patch with the MSW-MSA attention effect.",
},
),
"model": ("MODEL",),
},
}
@@ -235,15 +265,27 @@ class ApplyMSWMSAAttention:
class ApplyMSWMSAAttentionSimple:
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the MSW-MSA attention effect.",)
FUNCTION = "go"
CATEGORY = "model_patches/unet"
DESCRIPTION = "This node applies an attention patch which _may_ slightly improve quality especially when generating at high resolutions. It is a large performance increase on SD1.x, may improve performance on SDXL. This is the simplified version of the node with less parameters. Use ApplyMSWMSAAttention if you require more control. NOTE: Only supports SD1.x, SD2.x and SDXL."
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"model_type": (("SD15", "SDXL"),),
"model": ("MODEL",),
"model_type": (
("SD15", "SDXL"),
{
"tooltip": "Model type being patched. Choose SD15 for SD 1.4, SD 2.x.",
},
),
"model": (
"MODEL",
{
"tooltip": "Model to patch with the MSW-MSA attention effect.",
},
),
},
}
+225 -99
View File
@@ -3,7 +3,6 @@ from __future__ import annotations
import logging
import os
import sys
from functools import partial
from typing import TYPE_CHECKING
import torch
@@ -19,8 +18,6 @@ from .utils import (
)
if TYPE_CHECKING:
from typing import Callable
from comfy.model_patcher import ModelPatcher
F = torch.nn.functional
@@ -47,6 +44,9 @@ class HDConfig:
return False
return check_time(topts, self.start_sigma, self.end_sigma)
def __str__(self):
return f"<HDConfig: curr_sigma={self.curr_sigma}, start_sigma={self.start_sigma}, end_sigma={self.end_sigma}, use_blocks={self.use_blocks!r}, upscale_mode={self.upscale_mode}, two_stage_upscale_mode={self.two_stage_upscale_mode}>"
GLOBAL_STATE: HDState
@@ -54,16 +54,15 @@ GLOBAL_STATE: HDState
class HDState:
def __init__(self):
self.no_controlnet_workaround = (
os.environ.get("JANKHIDIFFUSION_NO_CONTROLNET_WORKAROUND") is not None
"JANKHIDIFFUSION_NO_CONTROLNET_WORKAROUND" in os.environ
)
self.controlnet_scale_args = {"mode": "bilinear", "align_corners": False}
self.patched_freeu_advanced = False
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 +77,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
@@ -131,102 +130,162 @@ class HDState:
GLOBAL_STATE = HDState()
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,
hdconfig: HDConfig,
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.hdconfig = hdconfig
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:
hdconfig = self.hdconfig
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,
"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)
)
if hdconfig.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,
)
x = scale_samples(
x,
shape[1] // 2,
shape[0] // 2,
mode=hdconfig.two_stage_upscale_mode,
shape[1],
shape[0],
mode=hdconfig.upscale_mode,
sigma=hdconfig.curr_sigma,
)
x = scale_samples(
x,
shape[1],
shape[0],
mode=hdconfig.upscale_mode,
sigma=hdconfig.curr_sigma,
)
return model.conv(x)
return orig_block.conv(x)
def forward_downsample(
self,
x: torch.Tensor,
) -> torch.Tensor:
hdconfig = self.hdconfig
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,
"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))
return tempop(x)
class ApplyRAUNet:
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",)
FUNCTION = "patch"
CATEGORY = "model_patches/unet"
DESCRIPTION = "This node is used to enable generation at higher resolutions than a model was trained for with less artifacts or other negative effects. This is the advanced version with more tuneable parameters, use ApplyRAUNetSimple if this seems too complex. NOTE: Only supports SD1.x, SD2.x and SDXL."
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"model": ("MODEL",),
"input_blocks": ("STRING", {"default": "3"}),
"output_blocks": ("STRING", {"default": "8"}),
"time_mode": (("percent", "timestep", "sigma"),),
"model": (
"MODEL",
{
"tooltip": "Model to be patched with the RAUNet effect.",
},
),
"input_blocks": (
"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",
},
),
"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",
},
),
"time_mode": (
("percent", "timestep", "sigma"),
{
"tooltip": "Time mode controls how to interpret the values in start_time and end_time.",
},
),
"start_time": (
"FLOAT",
{
@@ -235,6 +294,7 @@ class ApplyRAUNet:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time normal RAUNet effects start applying - value is inclusive.",
},
),
"end_time": (
@@ -245,9 +305,15 @@ class ApplyRAUNet:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time normal RAUNet effects end - value is inclusive.",
},
),
"upscale_mode": (
UPSCALE_METHODS,
{
"tooltip": "Method used when upscaling latents in output Upscale blocks.",
},
),
"upscale_mode": (UPSCALE_METHODS,),
"ca_start_time": (
"FLOAT",
{
@@ -256,6 +322,7 @@ class ApplyRAUNet:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time normal cross-attention effects start applying - value is inclusive..",
},
),
"ca_end_time": (
@@ -266,22 +333,52 @@ class ApplyRAUNet:
"max": 999.0,
"round": False,
"step": 0.01,
"tooltip": "Time normal cross-attention effects end - value is inclusive.",
},
),
"ca_input_blocks": (
"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).",
},
),
"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).",
},
),
"ca_upscale_mode": (
UPSCALE_METHODS,
{
"tooltip": "Mode used when upscaling latents in output cross-attention blocks.",
},
),
"ca_input_blocks": ("STRING", {"default": "4"}),
"ca_output_blocks": ("STRING", {"default": "8"}),
"ca_upscale_mode": (UPSCALE_METHODS,),
"ca_downscale_mode": (
("avg_pool2d", *UPSCALE_METHODS),
{"default": "avg_pool2d"},
{
"default": "avg_pool2d",
"tooltip": "Mode used when downscaling latents in output cross-attention blocks (use avg_pool2d for normal Hidiffusion behavior).",
},
),
"ca_downscale_factor": (
"FLOAT",
{"default": 2.0, "min": 0.01, "step": 0.1, "round": False},
{
"default": 2.0,
"min": 0.01,
"step": 0.1,
"round": False,
"tooltip": "Factor to downscale with in cross-attention, 2.0 means downscale to half size. Must be an integer when using ca_downscale_mode avg_pool2d.",
},
),
"two_stage_upscale_mode": (
("disabled", *UPSCALE_METHODS),
{"default": "disabled"},
{
"default": "disabled",
"tooltip": "When upscaling in output Upscale blocks (non-NA), do half the upscale with this mode and half with the normal upscale mode. May produce a different effect, isn't necessarily better.",
},
),
},
}
@@ -318,7 +415,6 @@ class ApplyRAUNet:
ca_use_blocks |= parse_blocks("input", ca_input_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(
@@ -391,16 +487,25 @@ class ApplyRAUNet:
model.set_model_output_block_patch(output_block_patch)
for block_type, block_index in use_blocks:
subidx, block_fun = (
(0, forward_downsample)
if block_type == "input"
else (2, forward_upsample)
main_block = model.get_model_object(
f"diffusion_model.{block_type}_blocks.{block_index}",
)
block_name = f"diffusion_model.{block_type}_blocks.{block_index}.{subidx}"
expected_class = (
openaimodel.Downsample
if block_type == "input"
else openaimodel.Upsample
)
block_name = f"diffusion_model.{block_type}_blocks.{block_index}.{len(main_block) - 1}"
block = model.get_model_object(block_name)
if not isinstance(block, expected_class):
block_type_name = getattr(type(block), "__name__", "unknown")
error_message = (
f"User error: {block_type} {block_index} requires targeting an {expected_class.__name__} block but got block of type {block_type_name} instead.",
)
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, hdconfig, block_index, block_type != "input"),
)
GLOBAL_STATE.apply_patches()
@@ -410,33 +515,54 @@ class ApplyRAUNet:
class ApplyRAUNetSimple:
RETURN_TYPES = ("MODEL",)
OUTPUT_TOOLTIPS = ("Model patched with the RAUNet effect.",)
FUNCTION = "patch"
CATEGORY = "model_patches/unet"
DESCRIPTION = "This node is used to enable generation at higher resolutions than a model was trained for with less artifacts or other negative effects. This is the simplified version with less parameters, use ApplyRAUNet if you require more control. NOTE: Only supports SD1.x, SD2.x and SDXL."
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"model": ("MODEL",),
"model_type": (("SD15", "SDXL"),),
"model": (
"MODEL",
{
"tooltip": "Model to be patched with the RAUNet effect.",
},
),
"model_type": (
("SD15", "SDXL"),
{
"tooltip": "Model type being patched. Choose SD15 for SD 1.4 or SD 2.x.",
},
),
"res_mode": (
(
"high (1536-2048)",
"low (1024 or lower)",
"ultra (over 2048)",
),
{
"tooltip": "Resolution mode hint, does not have to correspond to the actual size.",
},
),
"upscale_mode": (
(
"default",
*UPSCALE_METHODS,
),
{
"tooltip": "Method used when upscaling latents in output Upsample blocks.",
},
),
"ca_upscale_mode": (
(
"default",
*UPSCALE_METHODS,
),
{
"tooltip": "Method used when upscaling latents in cross attention blocks.",
},
),
},
}
+4 -1
View File
@@ -34,7 +34,10 @@ def convert_time(
raise ValueError(
"invalid value for end percent",
)
return (ms.percent_to_sigma(start_time), ms.percent_to_sigma(end_time))
return (
round(ms.percent_to_sigma(start_time), 4),
round(ms.percent_to_sigma(end_time), 4),
)
raise ValueError("invalid time mode")
+14
View File
@@ -0,0 +1,14 @@
[project]
name = "comfyui_jankhidiffusion"
description = "Janky implementation of HiDiffusion for ComfyUI. Enables generating at resolutions higher than what the model was trained for. Only supports SD 1.x (maybe 2.x) and SDXL."
version = "0.8.3"
license = { file = "LICENSE" }
[project.urls]
Repository = "https://github.com/blepping/comfyui_jankhidiffusion"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "blepping"
DisplayName = "comfyui_jankhidiffusion"
Icon = ""