From 7e173e161113cfbb6317cdeef77c460172ec1f4e Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 5 Dec 2024 03:58:13 -0700 Subject: [PATCH] Add skip and force_upscale_model parameters Use ComfyUI union types for wildcard inputs when available --- README.md | 10 ++++++ py/config.py | 14 ++++++-- py/guided_model.py | 4 +-- py/nodes.py | 78 +++++++++++++++++++++++++----------------- py/sampler.py | 32 ++++++++++------- py/tensor_image_ops.py | 5 ++- py/upscale.py | 9 +++-- 7 files changed, 100 insertions(+), 52 deletions(-) diff --git a/README.md b/README.md index 041e218..3b3db6d 100644 --- a/README.md +++ b/README.md @@ -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 } diff --git a/py/config.py b/py/config.py index db91dc1..76e4d43 100644 --- a/py/config.py +++ b/py/config.py @@ -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") diff --git a/py/guided_model.py b/py/guided_model.py index cccd532..fc1f5b5 100644 --- a/py/guided_model.py +++ b/py/guided_model.py @@ -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 diff --git a/py/nodes.py b/py/nodes.py index 434750d..d4f3d6f 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -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)}", }, ), }, diff --git a/py/sampler.py b/py/sampler.py index a0f08c7..1a9ea16 100644 --- a/py/sampler.py +++ b/py/sampler.py @@ -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, diff --git a/py/tensor_image_ops.py b/py/tensor_image_ops.py index f85a10c..5ae39a8 100644 --- a/py/tensor_image_ops.py +++ b/py/tensor_image_ops.py @@ -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 = { diff --git a/py/upscale.py b/py/upscale.py index bcd98ff..31d4c17 100644 --- a/py/upscale.py +++ b/py/upscale.py @@ -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: