Second pass
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
comfyui_jankdiffusehigh
|
||||
|
||||
Copyright https://gitub.com/blepping
|
||||
|
||||
This project was referenced from the original implementation at https://github.com/yhyun225/DiffuseHigh
|
||||
+214
@@ -0,0 +1,214 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from comfy.samplers import ksampler
|
||||
from pytorch_wavelets import DTCWTForward, DTCWTInverse, DWTForward, DWTInverse
|
||||
|
||||
from .tensor_image_ops import (
|
||||
BLENDING_MODES,
|
||||
Sharpen,
|
||||
)
|
||||
from .upscale import Upscale
|
||||
from .utils import fallback
|
||||
from .vae import VAEHelper
|
||||
|
||||
|
||||
class Config:
|
||||
_overridable_fields = { # noqa: RUF012
|
||||
"blend_by_mode",
|
||||
"blend_mode",
|
||||
"denoised_wavelet_multiplier",
|
||||
"dtcwt_biort",
|
||||
"dtcwt_mode",
|
||||
"dtcwt_qshift",
|
||||
"dwt_flip_filters",
|
||||
"dwt_level",
|
||||
"dwt_mode",
|
||||
"dwt_wave",
|
||||
"fadeout_factor",
|
||||
"guidance_factor",
|
||||
"guidance_mode",
|
||||
"guidance_restart_s_noise",
|
||||
"guidance_restart",
|
||||
"guidance_steps",
|
||||
"iteration_override",
|
||||
"reference_wavelet_multiplier",
|
||||
"renoise_factor",
|
||||
"resample_mode",
|
||||
"rescale_increment",
|
||||
"scale_factor",
|
||||
"sharpen_gaussian_kernel_size",
|
||||
"sharpen_gaussian_sigma",
|
||||
"sharpen_mode",
|
||||
"sharpen_reference",
|
||||
"sharpen_strength",
|
||||
"sigma_offset",
|
||||
"vae_decode_kwargs",
|
||||
"vae_encode_kwargs",
|
||||
"vae_mode",
|
||||
}
|
||||
|
||||
_dict_exclude_keys = { # noqa: RUF012
|
||||
"as_dict",
|
||||
"blend_function",
|
||||
"dwt",
|
||||
"get_iteration_config",
|
||||
"idwt",
|
||||
"iteration_override",
|
||||
"sharpen",
|
||||
"upscale",
|
||||
"vae",
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device,
|
||||
dtype,
|
||||
latent_format,
|
||||
*,
|
||||
blend_mode="lerp",
|
||||
blend_by_mode="image",
|
||||
denoised_wavelet_multiplier=1.0,
|
||||
dtcwt_biort="near_sym_a",
|
||||
dtcwt_mode=False,
|
||||
dtcwt_qshift="qshift_a",
|
||||
dwt_flip_filters=False,
|
||||
dwt_level=1,
|
||||
dwt_mode="symmetric",
|
||||
dwt_wave="db4",
|
||||
fadeout_factor=0.0,
|
||||
guidance_factor=1.0,
|
||||
guidance_mode="image",
|
||||
guidance_restart_s_noise=1.0,
|
||||
guidance_restart=0,
|
||||
guidance_sampler=None,
|
||||
guidance_steps=5,
|
||||
iteration_override=None,
|
||||
iterations=1,
|
||||
reference_sampler=None,
|
||||
reference_wavelet_multiplier=1.0,
|
||||
renoise_factor=1.0,
|
||||
resample_mode="bicubic",
|
||||
rescale_increment=64,
|
||||
sampler=None,
|
||||
scale_factor=2.0,
|
||||
sharpen_gaussian_kernel_size=3,
|
||||
sharpen_gaussian_sigma=(0.1, 2.0),
|
||||
sharpen_mode="gaussian",
|
||||
sharpen_reference=True,
|
||||
sharpen_strength=1.0,
|
||||
sigma_offset=0,
|
||||
upscale_model=None,
|
||||
vae_decode_kwargs=None,
|
||||
vae_encode_kwargs=None,
|
||||
vae_mode="normal",
|
||||
vae=None,
|
||||
):
|
||||
sampler = fallback(
|
||||
sampler,
|
||||
lambda: ksampler("euler"),
|
||||
default_is_fun=True,
|
||||
)
|
||||
self.sigma_offset = sigma_offset
|
||||
self.fadeout_factor = fadeout_factor
|
||||
self.scale_factor = scale_factor
|
||||
self.guidance_factor = guidance_factor
|
||||
self.renoise_factor = renoise_factor
|
||||
self.iterations = iterations
|
||||
self.guidance_steps = guidance_steps
|
||||
self.guidance_mode = guidance_mode
|
||||
self.guidance_restart = guidance_restart
|
||||
self.guidance_restart_s_noise = guidance_restart_s_noise
|
||||
self.sampler = sampler
|
||||
self.guidance_sampler = fallback(guidance_sampler, sampler)
|
||||
self.reference_sampler = fallback(reference_sampler, sampler)
|
||||
self.vae = VAEHelper(
|
||||
vae_mode,
|
||||
latent_format,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
vae=vae,
|
||||
encode_kwargs=fallback(vae_encode_kwargs, {}),
|
||||
decode_kwargs=fallback(vae_decode_kwargs, {}),
|
||||
)
|
||||
self.sharpen = Sharpen(
|
||||
mode=sharpen_mode,
|
||||
strength=sharpen_strength if sharpen_reference else 0,
|
||||
gaussian_kernel_size=sharpen_gaussian_kernel_size,
|
||||
gaussian_sigma=sharpen_gaussian_sigma,
|
||||
)
|
||||
self.upscale = Upscale(
|
||||
resample_mode=resample_mode,
|
||||
rescale_increment=rescale_increment,
|
||||
upscale_model=upscale_model,
|
||||
)
|
||||
self.dwt_mode = dwt_mode
|
||||
self.dwt_level = dwt_level
|
||||
self.dwt_wave = dwt_wave
|
||||
self.dtcwt_mode = dtcwt_mode
|
||||
self.dtcwt_biort = dtcwt_biort
|
||||
self.dtcwt_qshift = dtcwt_qshift
|
||||
if dtcwt_mode:
|
||||
self.dwt = DTCWTForward(
|
||||
J=dwt_level,
|
||||
mode=dwt_mode,
|
||||
biort=dtcwt_biort,
|
||||
qshift=dtcwt_qshift,
|
||||
).to(device)
|
||||
self.idwt = DTCWTInverse(
|
||||
mode=dwt_mode,
|
||||
biort=dtcwt_biort,
|
||||
qshift=dtcwt_qshift,
|
||||
).to(device)
|
||||
else:
|
||||
self.dwt = DWTForward(J=dwt_level, wave=dwt_wave, mode=dwt_mode).to(device)
|
||||
self.idwt = DWTInverse(wave=dwt_wave, mode=dwt_mode).to(device)
|
||||
self.dwt_flip_filters = dwt_flip_filters
|
||||
self.reference_wavelet_multiplier = reference_wavelet_multiplier
|
||||
self.denoised_wavelet_multiplier = denoised_wavelet_multiplier
|
||||
self.blend_mode = blend_mode
|
||||
if blend_by_mode not in {"image", "latent", "wavelet"}:
|
||||
raise ValueError("Bad blend_by_mode: must be one of image, latent, wavelet")
|
||||
self.blend_by_mode = blend_by_mode
|
||||
self.blend_function = BLENDING_MODES[blend_mode]
|
||||
self.iteration_override = {}
|
||||
if iteration_override is None or iteration_override == {}:
|
||||
return
|
||||
if not isinstance(iteration_override, dict):
|
||||
raise TypeError("Iteration override must be an object")
|
||||
# if isinstance(next(iter(iteration_override.values())), self.__class__):
|
||||
# self.iteration_Override = iteration_override
|
||||
# return
|
||||
selfdict = self.as_dict()
|
||||
overrides = self.iteration_override
|
||||
for k, v in iteration_override.items():
|
||||
if not isinstance(k, (int, str)) or not isinstance(v, dict):
|
||||
raise TypeError(
|
||||
"Bad type for override item: key must be integer or string, value must be an object",
|
||||
)
|
||||
okwargs = selfdict | {
|
||||
ok: ov for ok, ov in v.items() if ok in self._overridable_fields
|
||||
}
|
||||
overrides[k] = self.__class__(device, dtype, latent_format, **okwargs)
|
||||
|
||||
def as_dict(self) -> dict:
|
||||
result = {
|
||||
k: getattr(self, k)
|
||||
for k in dir(self)
|
||||
if not k.startswith("_") and k not in self._dict_exclude_keys
|
||||
}
|
||||
result["vae_mode"] = self.vae.mode.name.lower()
|
||||
result["vae"] = self.vae.vae
|
||||
result["vae_encode_kwargs"] = self.vae.encode_kwargs
|
||||
result["vae_decode_kwargs"] = self.vae.decode_kwargs
|
||||
result["sharpen_reference"] = self.sharpen.strength != 0
|
||||
result["sharpen_strength"] = self.sharpen.strength
|
||||
result["sharpen_gaussian_kernel_size"] = self.sharpen.gaussian_kernel_size
|
||||
result["sharpen_gaussian_sigma"] = self.sharpen.gaussian_sigma
|
||||
result["resample_mode"] = self.upscale.resample_mode
|
||||
result["rescale_increment"] = self.upscale.rescale_increment
|
||||
result["upscale_model"] = self.upscale.upscale_model
|
||||
return result
|
||||
|
||||
def get_iteration_config(self, iteration):
|
||||
override = self.iteration_override.get(iteration)
|
||||
return override.get_iteration_config(iteration) if override else self
|
||||
@@ -7,3 +7,10 @@ with contextlib.suppress(ImportError):
|
||||
EXTERNAL["tiled_diffusion"] = importlib.import_module(
|
||||
"custom_nodes.ComfyUI-TiledDiffusion",
|
||||
)
|
||||
|
||||
with contextlib.suppress(ImportError, NotImplementedError):
|
||||
bleh = importlib.import_module("custom_nodes.ComfyUI-bleh")
|
||||
bleh_version = getattr(bleh, "BLEH_VERSION", -1)
|
||||
if bleh_version < 1:
|
||||
raise NotImplementedError
|
||||
EXTERNAL["bleh"] = bleh.py
|
||||
|
||||
+95
-18
@@ -4,9 +4,15 @@ import yaml
|
||||
from comfy.samplers import KSAMPLER
|
||||
|
||||
from .sampler import diffusehigh_sampler
|
||||
from .vae import VAEMode
|
||||
|
||||
|
||||
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."
|
||||
OUTPUT_TOOLTIPS = (
|
||||
"SAMPLER that can be connected to a SamplerCustom or other sampler node that supports a SAMPLER input.",
|
||||
)
|
||||
CATEGORY = "sampling/custom_sampling/JankDiffuseHigh"
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
FUNCTION = "go"
|
||||
|
||||
@@ -14,13 +20,29 @@ class DiffuseHighSamplerNode:
|
||||
def INPUT_TYPES(cls) -> dict:
|
||||
return {
|
||||
"required": {
|
||||
"highres_sigmas": ("SIGMAS",),
|
||||
"guidance_steps": ("INT", {"default": 5, "min": 0}),
|
||||
"highres_sigmas": (
|
||||
"SIGMAS",
|
||||
{
|
||||
"tooltip": "Sigmas used for steps after upscaling. Generally should be around 0.3-0.5 denoise. NOTE: I do not recommend plugging in raw 1.0 denoise sigmas here.",
|
||||
},
|
||||
),
|
||||
"guidance_steps": (
|
||||
"INT",
|
||||
{
|
||||
"default": 5,
|
||||
"min": 0,
|
||||
"tooltip": "Number of guidance steps after an upscale.",
|
||||
},
|
||||
),
|
||||
"guidance_mode": (
|
||||
(
|
||||
"image",
|
||||
"latent",
|
||||
),
|
||||
{
|
||||
"default": "image",
|
||||
"tooltip": "The original implementation uses image guidance. This requires a VAE encode/decode per guidance step. Alternatively, you can try using guidance via the latent instead which is much faster.",
|
||||
},
|
||||
),
|
||||
"guidance_factor": (
|
||||
"FLOAT",
|
||||
@@ -28,28 +50,83 @@ class DiffuseHighSamplerNode:
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"tooltip": "Mix factor used on guidance steps. 1.0 means use 100% DiffuseHigh guidance for those steps (like the original implementation).",
|
||||
},
|
||||
),
|
||||
"fadeout_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.0,
|
||||
"tooltip": "Can be enabled to fade out guidance_factor. For example, if guidance_factor is 1 and guidance_steps is 4 then fadeout_factor would use these guidance_factors for the guidance steps: 1.00, 0.75, 0.50, 0.25",
|
||||
},
|
||||
),
|
||||
"scale_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 2.0,
|
||||
"tooltip": "Upscale factor per iteration.",
|
||||
},
|
||||
),
|
||||
"renoise_factor": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"tooltip": "Strength of noise added at the start of each iteration. The default of 1.0 (100%) is the normal amount, but you can increase this slightly to add more detail.",
|
||||
},
|
||||
),
|
||||
"iterations": (
|
||||
"INT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": 0,
|
||||
"tooltip": "Number of upscale iterations to run. Be careful, this can add up fast - if you start at 512x512 with a 2.0 scale factor then 3 iterations will get you to 4096x4096.",
|
||||
},
|
||||
),
|
||||
"fadeout_factor": ("FLOAT", {"default": 0.0}),
|
||||
"scale_factor": ("FLOAT", {"default": 2.0}),
|
||||
"renoise_factor": ("FLOAT", {"default": 1.0}),
|
||||
"iterations": ("INT", {"default": 1, "min": 0}),
|
||||
"sampler": ("SAMPLER",),
|
||||
"vae_mode": (
|
||||
(
|
||||
"taesd",
|
||||
"normal",
|
||||
"tiled",
|
||||
"tiled_diffusion",
|
||||
),
|
||||
tuple(vm.name.lower() for vm in VAEMode),
|
||||
{
|
||||
"default": "normal",
|
||||
"tooltip": "Mode used for encoding/decoding images. TAESD is fast/low VRAM but may reduce quality (you will also need the TAESD encoders installed). Normal will just use the normal VAE node, tiled with use the tiled VAE node. Alternatively, if you have ComfyUI-TiledDiffusion installed you can use tiled_diffusion here.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"reference_image_opt": ("IMAGE",),
|
||||
"guidance_sampler_opt": ("SAMPLER",),
|
||||
"reference_sampler_opt": ("SAMPLER",),
|
||||
"vae_opt": ("VAE",),
|
||||
"upscale_model_opt": ("UPSCALE_MODEL",),
|
||||
"sampler": (
|
||||
"SAMPLER",
|
||||
{
|
||||
"tooltip": "Default sampler used for steps. If not specified the sampler will default to non-ancestral Euler.",
|
||||
},
|
||||
),
|
||||
"reference_image_opt": (
|
||||
"IMAGE",
|
||||
{
|
||||
"tooltip": "Optional: Image used for the initial pass. If not connected, a low-res initial reference will be generated using the schedule from the normal sigmas.",
|
||||
},
|
||||
),
|
||||
"guidance_sampler_opt": (
|
||||
"SAMPLER",
|
||||
{
|
||||
"tooltip": "Optional: Sampler used for guidance steps. If not specified, will fallback to the base sampler. Note: The sampler is called on individual steps, samplers that keep history will not work well here.",
|
||||
},
|
||||
),
|
||||
"reference_sampler_opt": (
|
||||
"SAMPLER",
|
||||
{
|
||||
"tooltip": "Optional: Sampler used to generate the initial low-resolution reference. Only used if reference_image_opt is not connected.",
|
||||
},
|
||||
),
|
||||
"vae_opt": (
|
||||
"VAE",
|
||||
{
|
||||
"tooltip": "Optional when vae_mode is set to `taesd`, otherwise this is the VAE that will be used for encoding/decoding images.",
|
||||
},
|
||||
),
|
||||
"upscale_model_opt": (
|
||||
"UPSCALE_MODEL",
|
||||
{
|
||||
"tooltip": "Optional: Model used for upscaling. When not attached, simple image scaling will be used. Regardless, the image will be scaled to match the size expected based on scale_factor. For example, if you use scale_factor 2 and a 4x upscale model, the image will get scaled down after the upscale model runs.",
|
||||
},
|
||||
),
|
||||
"yaml_parameters": (
|
||||
"STRING",
|
||||
{
|
||||
|
||||
+116
-189
@@ -1,29 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import PIL.Image as PILImage
|
||||
import torch
|
||||
import torchvision
|
||||
from comfy_extras.nodes_upscale_model import ImageUpscaleWithModel
|
||||
from pytorch_wavelets import DTCWTForward, DTCWTInverse, DWTForward, DWTInverse
|
||||
from tqdm import tqdm
|
||||
from tqdm.auto import trange
|
||||
|
||||
from .utils import (
|
||||
ensure_model,
|
||||
pilimgbatch_to_torch,
|
||||
torch_to_pilimgbatch,
|
||||
from .config import Config
|
||||
from .tensor_image_ops import (
|
||||
blend_wavelets,
|
||||
scale_wavelets,
|
||||
)
|
||||
from .vae import VAEHelper
|
||||
|
||||
|
||||
def gaussian_blur_image_sharpening(image, kernel_size=3, sigma=(0.1, 2.0), alpha=1):
|
||||
gaussian_blur = torchvision.transforms.GaussianBlur(
|
||||
kernel_size=kernel_size,
|
||||
sigma=sigma,
|
||||
)
|
||||
image_blurred = gaussian_blur(image)
|
||||
return (alpha + 1) * image - alpha * image_blurred
|
||||
from .utils import ensure_model, fallback
|
||||
|
||||
|
||||
class DiffuseHighSampler:
|
||||
@@ -37,102 +23,37 @@ class DiffuseHighSampler:
|
||||
extra_args,
|
||||
disable_pbar,
|
||||
highres_sigmas,
|
||||
sampler,
|
||||
guidance_steps=5,
|
||||
guidance_mode="image",
|
||||
guidance_factor=1.0,
|
||||
guidance_restart=0,
|
||||
guidance_restart_s_noise=1.0,
|
||||
fadeout_factor=0.0,
|
||||
scale_factor=2.0,
|
||||
renoise_factor=1.0,
|
||||
iterations=1,
|
||||
vae_mode="normal",
|
||||
dwt_level=1,
|
||||
dwt_wave="db4",
|
||||
dwt_mode="symmetric",
|
||||
dwt_flip_filters=False,
|
||||
dtcwt_mode=False,
|
||||
dtcwt_biort="near_sym_a",
|
||||
dtcwt_qshift="qshift_a",
|
||||
reference_wavelet_multiplier=1.0,
|
||||
denoised_wavelet_multiplier=1.0,
|
||||
sharpen_reference=True,
|
||||
sharpen_kernel_size=3,
|
||||
sharpen_sigma=(0.1, 2.0),
|
||||
sharpen_alpha=1.0,
|
||||
resample_mode="bicubic",
|
||||
rescale_increment=64,
|
||||
guidance_sampler_opt=None,
|
||||
reference_sampler_opt=None,
|
||||
reference_image_opt=None,
|
||||
vae_opt=None,
|
||||
upscale_model_opt=None,
|
||||
**kwargs: dict,
|
||||
):
|
||||
self.s_in = initial_x.new_ones((initial_x.shape[0],))
|
||||
self.initial_x = initial_x
|
||||
self.callback = callback
|
||||
self.disable_pbar = disable_pbar
|
||||
self.sigmas = sigmas
|
||||
self.extra_args = extra_args if extra_args is not None else {}
|
||||
self.extra_args = fallback(extra_args, {})
|
||||
self.model = model
|
||||
self.latent_format = model.inner_model.inner_model.latent_format
|
||||
self.fadeout_factor = fadeout_factor
|
||||
self.scale_factor = scale_factor
|
||||
self.guidance_factor = guidance_factor
|
||||
self.renoise_factor = renoise_factor
|
||||
self.iterations = iterations
|
||||
self.highres_sigmas = highres_sigmas.clone().to(sigmas)
|
||||
self.guidance_steps = guidance_steps
|
||||
self.guidance_mode = guidance_mode
|
||||
self.guidance_restart = guidance_restart
|
||||
self.guidance_restart_s_noise = guidance_restart_s_noise
|
||||
self.sampler = sampler
|
||||
self.guidance_sampler = guidance_sampler_opt or sampler
|
||||
self.reference_sampler = reference_sampler_opt or sampler
|
||||
self.vae = VAEHelper(
|
||||
vae_mode,
|
||||
self.config = self.base_config = Config(
|
||||
initial_x.device,
|
||||
initial_x.dtype,
|
||||
self.latent_format,
|
||||
device=initial_x.device,
|
||||
dtype=initial_x.dtype,
|
||||
guidance_sampler=guidance_sampler_opt,
|
||||
reference_sampler=reference_sampler_opt,
|
||||
vae=vae_opt,
|
||||
upscale_model=upscale_model_opt,
|
||||
**kwargs,
|
||||
)
|
||||
self.highres_sigmas = highres_sigmas.detach().clone().to(sigmas)
|
||||
self.reference_image = reference_image_opt
|
||||
self.sharpen_reference = sharpen_reference
|
||||
self.sharpen_kernel_size = sharpen_kernel_size
|
||||
self.sharpen_sigma = sharpen_sigma
|
||||
self.sharpen_alpha = sharpen_alpha
|
||||
self.resample_mode = getattr(PILImage, resample_mode.upper())
|
||||
self.rescale_increment = self.scale_dim(
|
||||
max(8, rescale_increment),
|
||||
1,
|
||||
increment=8,
|
||||
)
|
||||
if dtcwt_mode:
|
||||
self.dwt = DTCWTForward(
|
||||
J=dwt_level,
|
||||
mode=dwt_mode,
|
||||
biort=dtcwt_biort,
|
||||
qshift=dtcwt_qshift,
|
||||
).to(
|
||||
initial_x.device,
|
||||
)
|
||||
self.idwt = DTCWTInverse(
|
||||
mode=dwt_mode,
|
||||
biort=dtcwt_biort,
|
||||
qshift=dtcwt_qshift,
|
||||
).to(initial_x.device)
|
||||
else:
|
||||
self.dwt = DWTForward(J=dwt_level, wave=dwt_wave, mode=dwt_mode).to(
|
||||
initial_x.device,
|
||||
)
|
||||
self.idwt = DWTInverse(wave=dwt_wave, mode=dwt_mode).to(initial_x.device)
|
||||
self.dwt_flip_filters = dwt_flip_filters
|
||||
self.reference_wavelet_multiplier = reference_wavelet_multiplier
|
||||
self.denoised_wavelet_multiplier = denoised_wavelet_multiplier
|
||||
self.guidance_waves = None
|
||||
self.guidance_latent = None
|
||||
self.upscale_model = upscale_model_opt
|
||||
|
||||
def __getattr__(self, key):
|
||||
return getattr(self.config, key)
|
||||
|
||||
def call_model(self, x, sigma):
|
||||
return self.model(x, sigma * self.s_in, **self.extra_args)
|
||||
@@ -153,43 +74,66 @@ class DiffuseHighSampler:
|
||||
return denoised
|
||||
mix_scale = (
|
||||
self.guidance_factor
|
||||
- ((self.guidance_factor / (self.guidance_steps + 1)) * idx)
|
||||
* self.fadeout_factor
|
||||
- ((self.guidance_factor / self.guidance_steps) * idx) * self.fadeout_factor
|
||||
)
|
||||
if mix_scale == 0:
|
||||
return denoised
|
||||
print("GUIDANCE APPLY", idx)
|
||||
if self.guidance_mode not in {"image", "latent"}:
|
||||
raise ValueError("ohno")
|
||||
if self.guidance_mode == "image":
|
||||
dn_img = self.vae.decode(denoised).to(denoised).movedim(-1, 1)
|
||||
print("DN_IMG", dn_img.shape)
|
||||
dn_img = (
|
||||
self.vae.decode(denoised, disable_pbar=self.disable_pbar)
|
||||
.to(denoised)
|
||||
.movedim(-1, 1)
|
||||
)
|
||||
denoised_waves = self.dwt(dn_img)
|
||||
del dn_img
|
||||
elif self.guidance_mode == "latent":
|
||||
denoised_waves = self.dwt(denoised)
|
||||
denoised_waves_orig = denoised_waves
|
||||
if self.denoised_wavelet_multiplier != 1:
|
||||
denoised_waves = (
|
||||
denoised_waves[0] * self.denoised_wavelet_multiplier,
|
||||
tuple(t * self.denoised_wavelet_multiplier for t in denoised_waves[1]),
|
||||
)
|
||||
denoised_waves = scale_wavelets(self.denoised_wavelet_multiplier)
|
||||
coeffs = (
|
||||
(self.guidance_waves[0], denoised_waves[1])
|
||||
if not self.dwt_flip_filters
|
||||
else (denoised_waves[0], self.guidance_waves[1])
|
||||
)
|
||||
if self.blend_by_mode == "wavelet" or (
|
||||
self.blend_by_mode == "image" and self.guidance_mode != "image"
|
||||
):
|
||||
coeffs = blend_wavelets(
|
||||
denoised_waves_orig,
|
||||
coeffs,
|
||||
mix_scale,
|
||||
self.blend_function,
|
||||
)
|
||||
result = self.idwt(coeffs)
|
||||
if self.guidance_mode == "image":
|
||||
result = self.vae.encode(result.cpu(), fix_dims=True)
|
||||
|
||||
print("GUIDE OUT", denoised.shape, result.shape, mix_scale)
|
||||
return torch.lerp(denoised, result.to(denoised), mix_scale)
|
||||
if self.blend_by_mode == "image":
|
||||
result = self.blend_function(
|
||||
dn_img,
|
||||
result.to(dn_img),
|
||||
dn_img.new_full((1,), mix_scale),
|
||||
).clamp_(0, 1)
|
||||
result = self.vae.encode(
|
||||
result.cpu(),
|
||||
fix_dims=True,
|
||||
disable_pbar=self.disable_pbar,
|
||||
)
|
||||
# tqdm.write(str(("GUIDE OUT", denoised.shape, result.shape, mix_scale)))
|
||||
if self.blend_by_mode != "latent":
|
||||
return result.to(denoised)
|
||||
return self.blend_function(
|
||||
denoised,
|
||||
result.to(denoised),
|
||||
denoised.new_full((1,), mix_scale),
|
||||
)
|
||||
|
||||
def run_steps(self, *, x=None, sigmas=None):
|
||||
x = self.initial_x if x is None else x
|
||||
sigmas = self.sigmas if sigmas is None else sigmas
|
||||
guidance_sigmas = sigmas[: self.guidance_steps + 1]
|
||||
normal_sigmas = sigmas[self.guidance_steps :]
|
||||
soffset = self.sigma_offset
|
||||
guidance_sigmas = sigmas[soffset : soffset + self.guidance_steps + 1]
|
||||
normal_sigmas = sigmas[soffset + self.guidance_steps :]
|
||||
step_idx = 0
|
||||
model = self.model
|
||||
|
||||
@@ -206,7 +150,12 @@ class DiffuseHighSampler:
|
||||
if hasattr(model, k):
|
||||
setattr(model_wrapper, k, getattr(model, k))
|
||||
|
||||
for repidx in range(self.guidance_restart + 1):
|
||||
for repidx in trange(
|
||||
self.guidance_restart + 1,
|
||||
initial=1,
|
||||
disable=self.guidance_restart < 1 or self.disable_pbar,
|
||||
desc="guidance steps iteration",
|
||||
):
|
||||
if repidx > 0:
|
||||
noise_factor = (
|
||||
guidance_sigmas[0] ** 2 - guidance_sigmas[-1] ** 2
|
||||
@@ -214,109 +163,84 @@ class DiffuseHighSampler:
|
||||
x = x + torch.randn_like(x) * (
|
||||
noise_factor * self.guidance_restart_s_noise
|
||||
)
|
||||
for idx in range(len(guidance_sigmas) - 1):
|
||||
guidance_steps = len(guidance_sigmas) - 1
|
||||
for idx in trange(
|
||||
guidance_steps,
|
||||
initial=1,
|
||||
disable=self.disable_pbar,
|
||||
desc="guidance step",
|
||||
):
|
||||
step_idx = idx
|
||||
x = self.run_sampler(
|
||||
x,
|
||||
guidance_sigmas[idx : idx + 2],
|
||||
model=model_wrapper,
|
||||
sampler=self.guidance_sampler,
|
||||
disable_pbar=True,
|
||||
)
|
||||
if len(normal_sigmas) > 1:
|
||||
ensure_model(model)
|
||||
x = self.run_sampler(x, normal_sigmas)
|
||||
with tqdm(disable=self.disable_pbar, total=1, desc="normal steps") as pbar:
|
||||
x = self.run_sampler(x, normal_sigmas)
|
||||
pbar.update()
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def scale_dim(n, factor, *, increment=64) -> int:
|
||||
return math.ceil((n * factor) / increment) * increment
|
||||
|
||||
def upscale(self, imgbatch):
|
||||
_batch, height, width, _channels = imgbatch.shape
|
||||
target_height = self.scale_dim(
|
||||
height,
|
||||
self.scale_factor,
|
||||
increment=self.rescale_increment,
|
||||
)
|
||||
target_width = self.scale_dim(
|
||||
width,
|
||||
self.scale_factor,
|
||||
increment=self.rescale_increment,
|
||||
)
|
||||
print(f">> UPSCALE: {width}x{height} -> {target_width}x{target_height}")
|
||||
if (target_height, target_width) == (height, width):
|
||||
return imgbatch
|
||||
if self.upscale_model is not None:
|
||||
print("** Upscaling with model")
|
||||
imgbatch = ImageUpscaleWithModel().upscale(self.upscale_model, imgbatch)[0]
|
||||
if imgbatch.shape[1:3] == (target_height, target_width):
|
||||
return imgbatch
|
||||
print(
|
||||
f"** PIL upscale {imgbatch.shape[2]}x{imgbatch.shape[1]} -> {target_width}x{target_height}",
|
||||
)
|
||||
ref_imgbatch = torch_to_pilimgbatch(self.reference_image)
|
||||
return pilimgbatch_to_torch(
|
||||
tuple(
|
||||
i.resize((target_width, target_height), resample=self.resample_mode)
|
||||
for i in ref_imgbatch
|
||||
),
|
||||
)
|
||||
|
||||
def run_sampler(self, x, sigmas, *, model=None, sampler=None):
|
||||
model = model or self.model
|
||||
sampler = sampler or self.sampler
|
||||
def run_sampler(self, x, sigmas, *, model=None, sampler=None, disable_pbar=False):
|
||||
sampler = fallback(sampler, self.sampler)
|
||||
return sampler.sampler_function(
|
||||
model,
|
||||
fallback(model, self.model),
|
||||
x,
|
||||
sigmas,
|
||||
callback=self.callback,
|
||||
extra_args=self.extra_args.copy(),
|
||||
disable=self.disable_pbar,
|
||||
disable=disable_pbar or self.disable_pbar,
|
||||
**sampler.extra_options,
|
||||
)
|
||||
|
||||
def __call__(self):
|
||||
self.config = self.base_config.get_iteration_config("reference")
|
||||
if self.reference_image is None:
|
||||
x_lr = self.run_sampler(
|
||||
self.initial_x,
|
||||
self.sigmas,
|
||||
sampler=self.reference_sampler,
|
||||
)
|
||||
with tqdm(disable=self.disable_pbar, desc="reference steps"):
|
||||
x_lr = self.run_sampler(
|
||||
self.initial_x,
|
||||
self.sigmas,
|
||||
sampler=self.reference_sampler,
|
||||
)
|
||||
if self.iterations < 1:
|
||||
return x_lr
|
||||
self.reference_image = self.vae.decode(x_lr)
|
||||
self.reference_image = self.vae.decode(x_lr, disable_pbar=self.disable_pbar)
|
||||
elif self.iterations < 1:
|
||||
return self.vae.encode(self.reference_image)
|
||||
for iteration in trange(self.iterations, disable=self.disable_pbar):
|
||||
print(
|
||||
f"\nIT({iteration}): shp={self.reference_image.shape}, min={self.reference_image.min()}, max={self.reference_image.max()}",
|
||||
)
|
||||
img_hr = self.upscale(self.reference_image)
|
||||
print("IMG_HR", img_hr.shape)
|
||||
if self.sharpen_reference:
|
||||
img_hr = gaussian_blur_image_sharpening(
|
||||
img_hr.movedim(-1, 1),
|
||||
kernel_size=self.sharpen_kernel_size,
|
||||
sigma=self.sharpen_sigma,
|
||||
alpha=self.sharpen_alpha,
|
||||
).movedim(1, -1)
|
||||
self.reference_image = img_hr
|
||||
x_new = self.vae.encode(self.reference_image).to(self.initial_x)
|
||||
self.guidance_latent = x_new.clone()
|
||||
return self.vae.encode(self.reference_image, disable_pbar=self.disable_pbar)
|
||||
self.config = self.base_config
|
||||
for iteration in trange(
|
||||
self.iterations,
|
||||
disable=self.disable_pbar,
|
||||
initial=1,
|
||||
desc="DiffuseHigh iteration",
|
||||
):
|
||||
self.config = self.base_config.get_iteration_config(iteration)
|
||||
with tqdm(disable=self.disable_pbar, total=1, desc="upscale") as pbar:
|
||||
img_hr = self.upscale(
|
||||
self.reference_image,
|
||||
self.scale_factor,
|
||||
pbar=pbar,
|
||||
)
|
||||
pbar.update()
|
||||
self.reference_image = self.sharpen(img_hr, fix_dims=True)
|
||||
x_new = self.vae.encode(
|
||||
self.reference_image,
|
||||
disable_pbar=self.disable_pbar,
|
||||
).to(self.initial_x)
|
||||
if self.guidance_mode == "image":
|
||||
print("REF IMG", self.reference_image.shape)
|
||||
self.guidance_waves = self.dwt(
|
||||
self.reference_image.clone().movedim(-1, 1).to(self.initial_x),
|
||||
)
|
||||
elif self.guidance_mode == "latent":
|
||||
self.guidance_waves = self.dwt(self.guidance_latent)
|
||||
self.guidance_waves = self.dwt(x_new.clone())
|
||||
if self.reference_wavelet_multiplier != 1:
|
||||
self.guidance_waves = (
|
||||
self.guidance_waves[0] * self.reference_wavelet_multiplier,
|
||||
tuple(
|
||||
t * self.reference_wavelet_multiplier
|
||||
for t in self.guidance_waves[1]
|
||||
),
|
||||
self.guidance_waves = scale_wavelets(
|
||||
self.guidance_waves,
|
||||
self.reference_wavelet_multiplier,
|
||||
)
|
||||
else:
|
||||
raise ValueError("ohno")
|
||||
@@ -330,7 +254,10 @@ class DiffuseHighSampler:
|
||||
result = self.run_steps(x=x_new, sigmas=self.highres_sigmas)
|
||||
if iteration == self.iterations - 1:
|
||||
break
|
||||
self.reference_image = self.vae.decode(result)
|
||||
self.reference_image = self.vae.decode(
|
||||
result,
|
||||
disable_pbar=self.disable_pbar,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
from enum import Enum, auto
|
||||
|
||||
import torch
|
||||
import torchvision
|
||||
|
||||
from .external import EXTERNAL
|
||||
|
||||
F = torch.nn.functional
|
||||
|
||||
EXT_BLEH = EXTERNAL.get("bleh")
|
||||
|
||||
if EXT_BLEH is not None:
|
||||
BLENDING_MODES = EXT_BLEH.latent_utils.BLENDING_MODES
|
||||
else:
|
||||
BLENDING_MODES = {
|
||||
"lerp": torch.lerp,
|
||||
}
|
||||
|
||||
|
||||
class SharpenMode(Enum):
|
||||
GAUSSIAN = auto()
|
||||
CONTRAST_ADAPTIVE = auto()
|
||||
|
||||
|
||||
def scale_wavelets(waves, factor=1.0):
|
||||
if factor == 1:
|
||||
return waves
|
||||
return (waves[0] * factor, tuple(t * factor for t in waves[1]))
|
||||
|
||||
|
||||
def blend_wavelets(a, b, factor, blend_function):
|
||||
if not isinstance(factor, torch.Tensor):
|
||||
factor = a[0].new_full((1,), factor)
|
||||
return (
|
||||
blend_function(a[0], b[0], factor),
|
||||
tuple(blend_function(ta, tb, factor) for ta, tb in zip(a[1], b[1])),
|
||||
)
|
||||
|
||||
|
||||
class Sharpen:
|
||||
def __init__(
|
||||
self,
|
||||
mode="gaussian",
|
||||
strength=1.0,
|
||||
gaussian_kernel_size=3,
|
||||
gaussian_sigma=(0.1, 2.0),
|
||||
):
|
||||
self.mode = getattr(SharpenMode, mode.upper(), None)
|
||||
if self.mode is None:
|
||||
raise ValueError("Bad sharpen mode")
|
||||
self.strength = strength
|
||||
self.gaussian_kernel_size = gaussian_kernel_size
|
||||
self.gaussian_sigma = gaussian_sigma
|
||||
|
||||
def __call__(self, t, *, fix_dims=False):
|
||||
if self.strength == 0:
|
||||
return t
|
||||
if fix_dims:
|
||||
t = t.movedim(-1, 1)
|
||||
if self.mode == SharpenMode.GAUSSIAN:
|
||||
result = gaussian_blur_image_sharpening(
|
||||
t,
|
||||
kernel_size=self.gaussian_kernel_size,
|
||||
sigma=self.gaussian_sigma,
|
||||
alpha=self.strength,
|
||||
)
|
||||
elif self.mode == SharpenMode.CONTRAST_ADAPTIVE:
|
||||
result = contrast_adaptive_sharpening(t, amount=self.strength)
|
||||
if fix_dims:
|
||||
result = result.movedim(1, -1)
|
||||
return result
|
||||
|
||||
|
||||
def gaussian_blur_image_sharpening(image, kernel_size=3, sigma=(0.1, 2.0), alpha=1):
|
||||
gaussian_blur = torchvision.transforms.GaussianBlur(
|
||||
kernel_size=kernel_size,
|
||||
sigma=sigma,
|
||||
)
|
||||
image_blurred = gaussian_blur(image)
|
||||
return (alpha + 1) * image - alpha * image_blurred
|
||||
|
||||
|
||||
# The following is modified to work with latent images of ~0 mean from https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/tree/main.
|
||||
def contrast_adaptive_sharpening(x, amount=0.8, *, epsilon=1e-06): # noqa: D417, PLR0914
|
||||
"""Performs contrast adaptive sharpening on the batch of images x.
|
||||
|
||||
The algorithm is directly implemented from FidelityFX's source code,
|
||||
that can be found here
|
||||
https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
x : Tensor
|
||||
Image or stack of images, of shape [batch, channels, ny, nx].
|
||||
Batch and channel dimensions can be ommited.
|
||||
amount : int [0, 1]
|
||||
Amount of sharpening to do, 0 being minimum and 1 maximum
|
||||
|
||||
Returns
|
||||
-------
|
||||
Tensor
|
||||
Processed stack of images.
|
||||
|
||||
""" # noqa: D401
|
||||
|
||||
def on_abs_stacked(tensor_list, f, *args: list, **kwargs: dict):
|
||||
return f(torch.abs(torch.stack(tensor_list)), *args, **kwargs)[0]
|
||||
|
||||
x_padded = F.pad(x, pad=(1, 1, 1, 1))
|
||||
x_padded = torch.complex(x_padded, torch.zeros_like(x_padded))
|
||||
# each side gets padded with 1 pixel
|
||||
# padding = same by default
|
||||
|
||||
# Extracting the 3x3 neighborhood around each pixel
|
||||
# a b c
|
||||
# d e f
|
||||
# g h i
|
||||
|
||||
a = x_padded[..., :-2, :-2]
|
||||
b = x_padded[..., :-2, 1:-1]
|
||||
c = x_padded[..., :-2, 2:]
|
||||
d = x_padded[..., 1:-1, :-2]
|
||||
e = x_padded[..., 1:-1, 1:-1]
|
||||
f = x_padded[..., 1:-1, 2:]
|
||||
g = x_padded[..., 2:, :-2]
|
||||
h = x_padded[..., 2:, 1:-1]
|
||||
i = x_padded[..., 2:, 2:]
|
||||
|
||||
# Computing contrast
|
||||
cross = (b, d, e, f, h)
|
||||
mn = on_abs_stacked(cross, torch.min, axis=0)
|
||||
mx = on_abs_stacked(cross, torch.max, axis=0)
|
||||
|
||||
diag = (a, c, g, i)
|
||||
mn2 = on_abs_stacked(diag, torch.min, axis=0)
|
||||
mx2 = on_abs_stacked(diag, torch.max, axis=0)
|
||||
|
||||
mx = mx + mx2
|
||||
mn = mn + mn2
|
||||
|
||||
# Computing local weight
|
||||
inv_mx = torch.reciprocal(mx + epsilon) # 1/mx
|
||||
|
||||
amp = inv_mx * mn
|
||||
|
||||
# scaling
|
||||
amp = torch.sqrt(amp)
|
||||
|
||||
w = -amp * (amount * (1 / 5 - 1 / 8) + 1 / 8)
|
||||
# w scales from 0 when amp=0 to K for amp=1
|
||||
# K scales from -1/5 when amount=1 to -1/8 for amount=0
|
||||
|
||||
# The local conv filter is
|
||||
# 0 w 0
|
||||
# w 1 w
|
||||
# 0 w 0
|
||||
div = torch.reciprocal(1 + 4 * w)
|
||||
output = ((b + d + f + h) * w + e) * div
|
||||
|
||||
return output.real.clamp(x.min(), x.max())
|
||||
@@ -0,0 +1,64 @@
|
||||
from comfy_extras.nodes_upscale_model import ImageUpscaleWithModel
|
||||
from PIL import Image
|
||||
|
||||
from .utils import (
|
||||
pilimgbatch_to_torch,
|
||||
scale_dim,
|
||||
torch_to_pilimgbatch,
|
||||
)
|
||||
|
||||
|
||||
class Upscale:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
resample_mode="bicubic",
|
||||
rescale_increment=64,
|
||||
upscale_model=None,
|
||||
):
|
||||
self.resample_mode = resample_mode
|
||||
self.rescale_increment = scale_dim(max(8, rescale_increment), increment=8)
|
||||
self.upscale_model = upscale_model
|
||||
|
||||
def __call__(self, imgbatch, scale_factor, *, pbar=None):
|
||||
if scale_factor == 1.0:
|
||||
return imgbatch
|
||||
_batch, height, width, _channels = imgbatch.shape
|
||||
target_height = scale_dim(
|
||||
height,
|
||||
scale_factor,
|
||||
increment=self.rescale_increment,
|
||||
)
|
||||
target_width = scale_dim(
|
||||
width,
|
||||
scale_factor,
|
||||
increment=self.rescale_increment,
|
||||
)
|
||||
# tqdm.write(f">> UPSCALE: {width}x{height} -> {target_width}x{target_height}")
|
||||
if (target_height, target_width) == (height, width):
|
||||
return imgbatch
|
||||
if self.upscale_model is not None:
|
||||
if pbar is not None:
|
||||
pbar.set_description(
|
||||
f"upscale with model: {width}x{height} -> {target_width}x{target_height}",
|
||||
)
|
||||
# tqdm.write("** Upscaling with model")
|
||||
imgbatch = ImageUpscaleWithModel().upscale(self.upscale_model, imgbatch)[0]
|
||||
if imgbatch.shape[1:3] == (target_height, target_width):
|
||||
return imgbatch
|
||||
# tqdm.write(
|
||||
# f"** PIL upscale {imgbatch.shape[2]}x{imgbatch.shape[1]} -> {target_width}x{target_height}",
|
||||
# )
|
||||
if pbar is not None:
|
||||
pbar.set_description(
|
||||
f"upscale: {imgbatch.shape[2]}x{imgbatch.shape[1]} -> {target_width}x{target_height}",
|
||||
)
|
||||
return pilimgbatch_to_torch(
|
||||
tuple(
|
||||
i.resize(
|
||||
(target_width, target_height),
|
||||
resample=getattr(Image.Resampling, self.resample_mode.upper()),
|
||||
)
|
||||
for i in torch_to_pilimgbatch(imgbatch)
|
||||
),
|
||||
)
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Sequence
|
||||
|
||||
import numpy as np
|
||||
@@ -36,3 +37,11 @@ def ensure_model(model):
|
||||
if model_management.LoadedModel(mp) in model_management.current_loaded_models:
|
||||
return
|
||||
model_management.load_models_gpu((mp,))
|
||||
|
||||
|
||||
def fallback(val, default, *, exclude=None, default_is_fun=False):
|
||||
return val if val is not exclude else (default() if default_is_fun else default)
|
||||
|
||||
|
||||
def scale_dim(n, factor=1.0, *, increment=64) -> int:
|
||||
return math.ceil((n * factor) / increment) * increment
|
||||
|
||||
@@ -5,8 +5,10 @@ from enum import Enum, auto
|
||||
import folder_paths
|
||||
import torch
|
||||
from comfy.taesd.taesd import TAESD
|
||||
from tqdm import tqdm
|
||||
|
||||
from .external import EXTERNAL
|
||||
from .utils import fallback
|
||||
|
||||
tiled_diffusion = EXTERNAL.get("tiled_diffusion")
|
||||
|
||||
@@ -27,8 +29,8 @@ class VAEHelper:
|
||||
device=None,
|
||||
dtype=None,
|
||||
vae=None,
|
||||
vae_encode_kwargs=None,
|
||||
vae_decode_kwargs=None,
|
||||
encode_kwargs=None,
|
||||
decode_kwargs=None,
|
||||
):
|
||||
if isinstance(mode, str):
|
||||
mode = VAEMode.__members__[mode.upper()]
|
||||
@@ -63,8 +65,8 @@ class VAEHelper:
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
self.vae = vae
|
||||
self.vae_encode_kwargs = {} if vae_encode_kwargs is None else vae_encode_kwargs
|
||||
self.vae_decode_kwargs = {} if vae_decode_kwargs is None else vae_decode_kwargs
|
||||
self.encode_kwargs = fallback(encode_kwargs, {})
|
||||
self.decode_kwargs = fallback(decode_kwargs, {})
|
||||
vae_handlers = {
|
||||
VAEMode.TAESD: (self.encode_taesd, self.decode_taesd),
|
||||
VAEMode.NORMAL: (self.encode_vae, self.decode_vae),
|
||||
@@ -79,22 +81,27 @@ class VAEHelper:
|
||||
}
|
||||
self.encode_fun, self.decode_fun = vae_handlers[mode]
|
||||
|
||||
def encode(self, imgbatch, *, fix_dims=False):
|
||||
def encode(self, imgbatch, *, fix_dims=False, disable_pbar=None):
|
||||
if fix_dims:
|
||||
imgbatch = imgbatch.moveaxis(1, -1)
|
||||
# print("ENCODING", imgbatch.min(), imgbatch.max())
|
||||
result = self.encode_fun(imgbatch[..., :3])
|
||||
with tqdm(disable=disable_pbar, total=1, desc="VAE encode") as pbar:
|
||||
result = self.encode_fun(imgbatch[..., :3])
|
||||
pbar.update()
|
||||
if self.mode != VAEMode.TAESD:
|
||||
# print("ENCODED(raw):", result.min(), result.max())
|
||||
result = self.latent_format.process_in(result)
|
||||
# print("ENCODED", result.shape, result.min(), result.max())
|
||||
return result
|
||||
|
||||
def decode(self, latent, *, skip_process_out=False):
|
||||
def decode(self, latent, *, skip_process_out=False, disable_pbar=None):
|
||||
if self.mode != VAEMode.TAESD and not skip_process_out:
|
||||
latent = self.latent_format.process_out(latent)
|
||||
# print("DECODING", latent.min(), latent.max())
|
||||
return self.decode_fun(latent)
|
||||
with tqdm(disable=disable_pbar, total=1, desc="VAE decode") as pbar:
|
||||
result = self.decode_fun(latent)
|
||||
pbar.update()
|
||||
return result
|
||||
# print("DECODED", result.shape, result.min(), result.max())
|
||||
|
||||
def encode_taesd(self, imgbatch):
|
||||
@@ -106,19 +113,19 @@ class VAEHelper:
|
||||
|
||||
def encode_vae(self, imgbatch):
|
||||
# print("VAE ENC", imgbatch.shape)
|
||||
return self.vae.encode(imgbatch, **self.vae_encode_kwargs)
|
||||
return self.vae.encode(imgbatch, **self.encode_kwargs)
|
||||
|
||||
def decode_vae(self, latent):
|
||||
return self.vae.decode(latent, **self.vae_decode_kwargs)
|
||||
return self.vae.decode(latent, **self.decode_kwargs)
|
||||
|
||||
def encode_vae_tiled(self, imgbatch):
|
||||
return self.vae.encode_tiled(imgbatch, **self.vae_encode_kwargs)
|
||||
return self.vae.encode_tiled(imgbatch, **self.encode_kwargs)
|
||||
|
||||
def decode_vae_tiled(self, latent):
|
||||
return self.vae.decode_tiled(latent, **self.vae_decode_kwargs)
|
||||
return self.vae.decode_tiled(latent, **self.decode_kwargs)
|
||||
|
||||
def encode_vae_tiled_diffusion(self, imgbatch):
|
||||
kwargs = self.td_encode_default_kwargs | self.vae_encode_kwargs
|
||||
kwargs = self.td_encode_default_kwargs | self.encode_kwargs
|
||||
return tiled_diffusion.tiled_vae.VAEEncodeTiled_TiledDiffusion().process(
|
||||
pixels=imgbatch,
|
||||
vae=self.vae,
|
||||
@@ -126,7 +133,7 @@ class VAEHelper:
|
||||
)[0]["samples"]
|
||||
|
||||
def decode_vae_tiled_diffusion(self, latent):
|
||||
kwargs = self.td_decode_default_kwargs | self.vae_decode_kwargs
|
||||
kwargs = self.td_decode_default_kwargs | self.decode_kwargs
|
||||
return tiled_diffusion.tiled_vae.VAEDecodeTiled_TiledDiffusion().process(
|
||||
samples={"samples": latent},
|
||||
vae=self.vae,
|
||||
|
||||
Reference in New Issue
Block a user