Add skip and force_upscale_model parameters

Use ComfyUI union types for wildcard inputs when available
This commit is contained in:
blepping
2024-12-05 03:58:13 -07:00
parent 9c41683e3a
commit 7e173e1611
7 changed files with 100 additions and 52 deletions
+10
View File
@@ -82,6 +82,10 @@ scale_factor: 1.5
Default advanced parameter values:
```yaml
# Mainly useful in iteration overrides, allows skipping an iteration. When defined at
# the toplevel it will skip everything which probably isn't what you want.
skip: false
# Mode used for blending the normal model prediction with the guidance during guidance steps.
# Only has an effect when guidance_factor is less than 1.0
# One of: image, latent, wavelets
@@ -166,6 +170,12 @@ sigma_dishonesty_factor_guidance: null
# iteration overrides.
use_upscale_model: true
# Only has an effect if the upscale model is connected and enabled. This
# will force it to run even if the scale factor is 1 or the size already
# matches the scale. This is to allow use of 1x upscale models that just
# add an effect like film grain.
force_upscale_model: false
# Allows passing extra arguments to the VAE encoder/decoder. Must be null or an object.
# Mainly useful with tiled_diffusion where you could do something like:
# vae_decode_kwargs: { fast: false }
+11 -3
View File
@@ -28,6 +28,7 @@ class Config:
"enable_cache_clearing",
"enable_gc",
"fadeout_factor",
"force_upscale_model",
"guidance_mask_blend_mode",
"guidance_mask_name",
"guidance_factor",
@@ -56,6 +57,7 @@ class Config:
"sharpen_strength",
"sigma_dishonesty_factor_guidance",
"sigma_dishonesty_factor",
"skip",
"skip_callback",
"upscale_model_name",
"use_upscale_model",
@@ -101,6 +103,7 @@ class Config:
enable_gc=True,
enable_cache_clearing=True,
fadeout_factor=0.0,
force_upscale_model=False,
guidance_factor=1.0,
guidance_mode="image",
guidance_restart_s_noise=1.0,
@@ -122,7 +125,7 @@ class Config:
sharpen_reference=True,
sharpen_strength=1.0,
skip_callback=False,
sigma_dishonesty_factor_guidance: None | float = None,
sigma_dishonesty_factor_guidance: float | None = None,
sigma_dishonesty_factor=0.0,
use_upscale_model=True,
vae_decode_kwargs=None,
@@ -141,7 +144,10 @@ class Config:
guidance_mask_name="guidance",
mask_blend_mode="lerp",
guidance_mask_blend_mode="lerp",
skip=False,
):
self.skip = skip
self.vae_name = vae_name
self.upscale_model_name = upscale_model_name
self.highres_sigmas_name = highres_sigmas_name
@@ -211,6 +217,7 @@ class Config:
resample_mode=resample_mode,
rescale_increment=rescale_increment,
upscale_model=upscale_model,
force_upscale_model=force_upscale_model,
)
if schedule_override is not None and not isinstance(schedule_override, dict):
raise TypeError("Bad type for schedule_override: must be null or object")
@@ -285,6 +292,7 @@ class Config:
result["sharpen_gaussian_sigma"] = self.sharpen.gaussian_sigma
result["resample_mode"] = self.upscale.resample_mode
result["rescale_increment"] = self.upscale.rescale_increment
result["force_upscale_model"] = self.upscale.force_upscale_model
return result
def get_iteration_config(self, iteration):
@@ -319,9 +327,9 @@ class ParamGroup:
self,
type_name: str,
*,
name: None | str = "",
name: str | None = "",
param_mode: bool = False,
default: None | Any = None, # noqa: ANN401
default: Any | None = None, # noqa: ANN401
) -> Any: # noqa: ANN401
name = name if name is not None else ""
key = (type_name, name) if not param_mode else (type_name, name, "params")
+2 -2
View File
@@ -32,7 +32,7 @@ class GuidedModel:
):
self.dhso = dh_sampler_object
self.allow_guidance = True
self.force_guidance: None | int = None
self.force_guidance: int | None = None
self.set_guidance_range(guidance_sigmas, guidance_steps)
def set_guidance_range(self, guidance_sigmas, guidance_steps):
@@ -45,7 +45,7 @@ class GuidedModel:
self.guidance_sigmas_list = tuple(guidance_sigmas.detach().cpu().tolist())
self.guidance_steps = guidance_steps
def get_guidance_step(self, sigma: torch.Tensor) -> None | int:
def get_guidance_step(self, sigma: torch.Tensor) -> int | None:
sigma_float = sigma_to_float(sigma)
if not self.allow_guidance:
return None
+47 -31
View File
@@ -9,6 +9,49 @@ from .external import init_integrations
from .sampler import diffusehigh_sampler
from .vae import VAEMode
try:
from comfy_execution import validation as comfy_validation
if not hasattr(comfy_validation, "validate_node_input"):
raise NotImplementedError # noqa: TRY301
HAVE_COMFY_UNION_TYPE = comfy_validation.validate_node_input("B", "A,B")
except (ImportError, NotImplementedError):
HAVE_COMFY_UNION_TYPE = False
except Exception as exc: # noqa: BLE001
HAVE_COMFY_UNION_TYPE = False
print(
f"** jankdiffusehigh: Warning, caught unexpected exception trying to detect ComfyUI union type support. Disabling. Exception: {exc}",
)
PARAM_TYPES = frozenset((
"IMAGE",
"MASK",
"OCS_NOISE",
"SAMPLER",
"SIGMAS",
"SONAR_CUSTOM_NOISE",
"UPSCALE_MODEL",
"VAE",
))
if not HAVE_COMFY_UNION_TYPE:
class Wildcard(str): # noqa: FURB189
__slots__ = ("whitelist",)
@classmethod
def __new__(cls, s, *args: list, whitelist=None, **kwargs: dict):
result = super().__new__(s, *args, **kwargs)
result.whitelist = whitelist
return result
def __ne__(self, other):
return False if self.whitelist is None else other not in self.whitelist
WILDCARD_PARAM = Wildcard("*", whitelist=PARAM_TYPES)
else:
WILDCARD_PARAM = ",".join(PARAM_TYPES)
class DiffuseHighSamplerNode:
DESCRIPTION = "Jank DiffuseHigh sampler node, used for generating directly to resolutions higher than what the model was trained for. Can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input."
@@ -152,8 +195,8 @@ class DiffuseHighSamplerNode:
def go(
cls,
*,
input_params_opt: None | ParamGroup = None,
yaml_parameters: None | str = None,
input_params_opt: ParamGroup | None = None,
yaml_parameters: str | None = None,
**kwargs: dict,
) -> tuple[KSAMPLER]:
init_integrations()
@@ -205,19 +248,6 @@ class DiffuseHighSamplerNode:
)
class Wildcard(str): # noqa: FURB189
__slots__ = ("whitelist",)
@classmethod
def __new__(cls, s, *args: list, whitelist=None, **kwargs: dict):
result = super().__new__(s, *args, **kwargs)
result.whitelist = whitelist
return result
def __ne__(self, other):
return False if self.whitelist is None else other not in self.whitelist
class DiffuseHighParamNode:
RETURN_TYPES = ("DIFFUSEHIGH_PARAMS",)
CATEGORY = "sampling/custom_sampling/JankDiffuseHigh"
@@ -228,20 +258,6 @@ class DiffuseHighParamNode:
FUNCTION = "go"
WC = Wildcard(
"*",
whitelist={
"IMAGE",
"MASK",
"OCS_NOISE",
"SAMPLER",
"SIGMAS",
"SONAR_CUSTOM_NOISE",
"UPSCALE_MODEL",
"VAE",
},
)
PARAM_TYPES = { # noqa: RUF012
"vae": lambda _v: True,
"sampler": lambda v: hasattr(v, "sampler_function"),
@@ -263,9 +279,9 @@ class DiffuseHighParamNode:
},
),
"value": (
cls.WC,
WILDCARD_PARAM,
{
"tooltip": "Connect the type of value expected by the key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the key you will get an error when you run the workflow.",
"tooltip": f"Connect the type of value expected by the key. Allows connecting output from any type of node HOWEVER if it is the wrong type expected by the key you will get an error when you run the workflow.\nThe following input types are supported: {', '.join(PARAM_TYPES)}",
},
),
},
+19 -13
View File
@@ -29,9 +29,9 @@ class DiffuseHighSampler:
initial_x: torch.Tensor,
sigmas: torch.Tensor,
*,
callback: None | callable,
extra_args: None | dict,
disable_pbar: None | bool,
callback: callable | None,
extra_args: dict | None,
disable_pbar: bool | None,
_params,
**kwargs: dict[str, Any],
):
@@ -53,8 +53,8 @@ class DiffuseHighSampler:
_params=_params,
**kwargs,
)
self.highres_sigmas: None | torch.FloatTensor = None
self.reference_image: None | torch.FloatTensor = None
self.highres_sigmas: torch.FloatTensor | None = None
self.reference_image: torch.FloatTensor | None = None
self.guidance_waves = None
self.mask = self.guidance_mask = None
self.seed_offset = 0
@@ -175,10 +175,10 @@ class DiffuseHighSampler:
latent: torch.Tensor,
sigma: float | torch.Tensor,
*,
sigma_next: None | float | torch.Tensor = None,
sigma_next: float | torch.Tensor | None = None,
factor=1.0,
allow_max_denoise=True,
noise_sampler: None | callable = None,
noise_sampler: callable | None = None,
) -> torch.Tensor:
self.seed_offset += 1
sigma_next = fallback(sigma_next, sigma)
@@ -215,7 +215,7 @@ class DiffuseHighSampler:
self,
x: torch.Tensor,
*,
sigmas: None | torch.Tensor = None,
sigmas: torch.Tensor | None = None,
) -> torch.Tensor:
sigmas = self.highres_sigmas if sigmas is None else sigmas
sigmas_len = len(sigmas)
@@ -340,8 +340,8 @@ class DiffuseHighSampler:
x: torch.Tensor,
sigmas: torch.Tensor,
*,
model: None | object = None,
sampler: None | object = None,
model: object | None = None,
sampler: object | None = None,
disable_pbar: bool = False,
):
sampler = fallback(sampler, self.sampler)
@@ -477,6 +477,8 @@ class DiffuseHighSampler:
desc="DiffuseHigh iteration",
):
self.update_iteration_config(iteration)
if self.skip:
continue
del x_new
if self.highres_sigmas[-1] != 0:
raise ValueError(
@@ -528,6 +530,10 @@ class DiffuseHighSampler:
x_new,
disable_pbar=self.disable_pbar,
)
if x_new is None:
raise ValueError(
"All iterations skipped, cannot return a result from sampler!",
)
return x_new
@@ -537,9 +543,9 @@ def diffusehigh_sampler(
sigmas: torch.Tensor,
*,
diffusehigh_options: dict[str, Any],
disable: None | bool = None,
extra_args: None | dict[str, Any] = None,
callback: None | callable = None,
disable: bool | None = None,
extra_args: dict[str, Any] | None = None,
callback: callable | None = None,
) -> torch.Tensor:
return DiffuseHighSampler(
model,
+4 -1
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from enum import Enum, auto
from typing import Sequence
from typing import TYPE_CHECKING
import numpy as np
import torch
@@ -10,6 +10,9 @@ from PIL import Image as PILImage
from .external import EXTERNAL
if TYPE_CHECKING:
from collections.abc import Sequence
F = torch.nn.functional
BLENDING_MODES = {
+7 -2
View File
@@ -15,13 +15,15 @@ class Upscale:
resample_mode="bicubic",
rescale_increment=64,
upscale_model=None,
force_upscale_model=False,
):
self.resample_mode = resample_mode
self.rescale_increment = scale_dim(max(8, rescale_increment), increment=8)
self.upscale_model = upscale_model
self.force_upscale_model = upscale_model is not None and force_upscale_model
def __call__(self, imgbatch, scale_factor, *, pbar=None, use_upscale_model=True):
if scale_factor == 1.0:
if scale_factor == 1.0 and not self.force_upscale_model:
return imgbatch
_batch, height, width, _channels = imgbatch.shape
target_height = scale_dim(
@@ -35,7 +37,10 @@ class Upscale:
increment=self.rescale_increment,
)
# tqdm.write(f">> UPSCALE: {width}x{height} -> {target_width}x{target_height}")
if (target_height, target_width) == (height, width):
if (target_height, target_width) == (
height,
width,
) and not self.force_upscale_model:
return imgbatch
if use_upscale_model and self.upscale_model is not None:
if pbar is not None: