Add skip and force_upscale_model parameters
Use ComfyUI union types for wildcard inputs when available
This commit is contained in:
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user