First pass implementation

This commit is contained in:
blepping
2024-10-15 08:49:31 -06:00
parent 51a8bf2e07
commit e0f1c94ea1
10 changed files with 730 additions and 3 deletions
+4 -3
View File
@@ -1,3 +1,6 @@
blehconfig.json
blehconfig.yaml
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
@@ -106,10 +109,8 @@ ipython_config.py
#pdm.lock
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
# in version control.
# https://pdm.fming.dev/latest/usage/project/#working-with-version-control
# https://pdm.fming.dev/#use-with-ide
.pdm.toml
.pdm-python
.pdm-build/
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
__pypackages__/
+5
View File
@@ -0,0 +1,5 @@
from .py import nodes
NODE_CLASS_MAPPINGS = {
"DiffuseHighSampler": nodes.DiffuseHighSamplerNode,
}
View File
+9
View File
@@ -0,0 +1,9 @@
import contextlib
import importlib
EXTERNAL = {}
with contextlib.suppress(ImportError):
EXTERNAL["tiled_diffusion"] = importlib.import_module(
"custom_nodes.ComfyUI-TiledDiffusion",
)
+84
View File
@@ -0,0 +1,84 @@
from __future__ import annotations
import yaml
from comfy.samplers import KSAMPLER
from .sampler import diffusehigh_sampler
class DiffuseHighSamplerNode:
RETURN_TYPES = ("SAMPLER",)
FUNCTION = "go"
@classmethod
def INPUT_TYPES(cls) -> dict:
return {
"required": {
"highres_sigmas": ("SIGMAS",),
"guidance_steps": ("INT", {"default": 5, "min": 0}),
"guidance_mode": (
(
"image",
"latent",
),
),
"guidance_factor": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
},
),
"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",
),
),
},
"optional": {
"reference_image_opt": ("IMAGE",),
"guidance_sampler_opt": ("SAMPLER",),
"reference_sampler_opt": ("SAMPLER",),
"vae_opt": ("VAE",),
"upscale_model_opt": ("UPSCALE_MODEL",),
"yaml_parameters": (
"STRING",
{
"tooltip": "Allows specifying custom parameters via YAML. You can also override any of the normal parameters by key. This input can be converted into a multiline text widget. Note: When specifying paramaters this way, there is very little error checking.",
"dynamicPrompts": False,
"multiline": True,
"defaultInput": True,
},
),
},
}
@classmethod
def go(cls, yaml_parameters: None | str = None, **kwargs: dict) -> tuple[KSAMPLER]:
if yaml_parameters:
extra_params = yaml.safe_load(yaml_parameters)
if extra_params is None:
pass
elif not isinstance(extra_params, dict):
raise ValueError(
"DiffuseHighSampler: yaml_parameters must either be null or an object",
)
else:
kwargs |= extra_params
return (
KSAMPLER(
diffusehigh_sampler,
extra_options={
"diffusehigh_options": kwargs,
},
),
)
+356
View File
@@ -0,0 +1,356 @@
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.auto import trange
from .utils import (
ensure_model,
pilimgbatch_to_torch,
torch_to_pilimgbatch,
)
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
class DiffuseHighSampler:
def __init__(
self,
model,
initial_x,
sigmas,
*,
callback,
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,
):
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.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.latent_format,
device=initial_x.device,
dtype=initial_x.dtype,
vae=vae_opt,
)
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 call_model(self, x, sigma):
return self.model(x, sigma * self.s_in, **self.extra_args)
def do_callback(self, idx, x, sigma, denoised):
if self.callback is None:
return
self.callback({
"i": idx,
"x": x,
"sigma": sigma,
"sigma_hat": sigma,
"denoised": denoised,
})
def apply_guidance(self, idx, denoised):
if self.guidance_waves is None or idx >= self.guidance_steps:
return denoised
mix_scale = (
self.guidance_factor
- ((self.guidance_factor / (self.guidance_steps + 1)) * 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)
denoised_waves = self.dwt(dn_img)
del dn_img
elif self.guidance_mode == "latent":
denoised_waves = self.dwt(denoised)
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]),
)
coeffs = (
(self.guidance_waves[0], denoised_waves[1])
if not self.dwt_flip_filters
else (denoised_waves[0], self.guidance_waves[1])
)
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)
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 :]
step_idx = 0
model = self.model
def model_wrapper(x, sigma, **extra_args: dict):
nonlocal step_idx
ensure_model(model)
denoised = model(x, sigma, **extra_args)
return self.apply_guidance(step_idx, denoised)
for k in (
"inner_model",
"sigmas",
):
if hasattr(model, k):
setattr(model_wrapper, k, getattr(model, k))
for repidx in range(self.guidance_restart + 1):
if repidx > 0:
noise_factor = (
guidance_sigmas[0] ** 2 - guidance_sigmas[-1] ** 2
) ** 0.5
x = x + torch.randn_like(x) * (
noise_factor * self.guidance_restart_s_noise
)
for idx in range(len(guidance_sigmas) - 1):
step_idx = idx
x = self.run_sampler(
x,
guidance_sigmas[idx : idx + 2],
model=model_wrapper,
sampler=self.guidance_sampler,
)
if len(normal_sigmas) > 1:
ensure_model(model)
x = self.run_sampler(x, normal_sigmas)
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
return sampler.sampler_function(
model,
x,
sigmas,
callback=self.callback,
extra_args=self.extra_args.copy(),
disable=self.disable_pbar,
**sampler.extra_options,
)
def __call__(self):
if self.reference_image is None:
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)
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()
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)
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]
),
)
else:
raise ValueError("ohno")
x_noise = torch.randn_like(x_new)
x_new = x_new + x_noise * (self.highres_sigmas[0] * self.renoise_factor)
# x_new = self.model.inner_model.inner_model.model_sampling.noise_scaling(
# self.highres_sigmas[0] * self.renoise_factor,
# x_noise,
# x_new,
# )
result = self.run_steps(x=x_new, sigmas=self.highres_sigmas)
if iteration == self.iterations - 1:
break
self.reference_image = self.vae.decode(result)
return result
def diffusehigh_sampler(
model,
x,
sigmas,
*,
diffusehigh_options,
disable=None,
extra_args=None,
callback=None,
):
sampler = DiffuseHighSampler(
model,
x,
sigmas,
disable_pbar=disable,
callback=callback,
extra_args=extra_args,
**diffusehigh_options,
)
return sampler()
+38
View File
@@ -0,0 +1,38 @@
from __future__ import annotations
from typing import Sequence
import numpy as np
import torch
from comfy import model_management
from PIL import Image as PILImage
def pilimgbatch_to_torch(
imgbatch: Sequence[PILImage, ...] | torch.Tensor,
) -> torch.Tensor:
if isinstance(imgbatch, torch.Tensor):
print("pibtt: skip", imgbatch.shape, imgbatch.min(), imgbatch.max())
return imgbatch
npi = np.stack(
tuple(np.array(i).astype(np.float32) / 255.0 for i in imgbatch),
axis=0,
)
return torch.from_numpy(npi)
return torch.from_numpy(npi.transpose(0, 3, 1, 2))
def torch_to_pilimgbatch(t: torch.Tensor) -> tuple[PILImage, ...]:
return tuple(
PILImage.fromarray(
np.clip((255.0 * i).cpu().numpy(), 0, 255).astype(np.uint8),
)
for i in t
)
def ensure_model(model):
mp = model.inner_model.model_patcher
if model_management.LoadedModel(mp) in model_management.current_loaded_models:
return
model_management.load_models_gpu((mp,))
+189
View File
@@ -0,0 +1,189 @@
from __future__ import annotations
from enum import Enum, auto
import folder_paths
import torch
from comfy.taesd.taesd import TAESD
from .external import EXTERNAL
tiled_diffusion = EXTERNAL.get("tiled_diffusion")
class VAEMode(Enum):
TAESD = auto()
NORMAL = auto()
TILED = auto()
TILED_DIFFUSION = auto()
class VAEHelper:
def __init__(
self,
mode: VAEMode | str,
latent_format,
*,
device=None,
dtype=None,
vae=None,
vae_encode_kwargs=None,
vae_decode_kwargs=None,
):
if isinstance(mode, str):
mode = VAEMode.__members__[mode.upper()]
if mode == VAEMode.TILED_DIFFUSION:
if tiled_diffusion is None:
raise ValueError(
"Cannot use tiled_diffusion VAE mode without ComfyUI-TiledDiffusion!",
)
self.td_encode_default_kwargs, self.td_decode_default_kwargs = (
{
k: v[1]["default"]
for k, v in td_node.INPUT_TYPES()
.get(
"required",
{},
)
.items()
if k not in {"pixels", "samples", "vae"}
and len(v) == 2
and isinstance(v[1], dict)
and "default" in v[1]
}
for td_node in (
tiled_diffusion.tiled_vae.VAEEncodeTiled_TiledDiffusion,
tiled_diffusion.tiled_vae.VAEDecodeTiled_TiledDiffusion,
)
)
if mode != VAEMode.TAESD and vae is None:
raise ValueError("Must pass a VAE when using non-TAESD VAE modes!")
self.mode = mode
self.latent_format = latent_format
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
vae_handlers = {
VAEMode.TAESD: (self.encode_taesd, self.decode_taesd),
VAEMode.NORMAL: (self.encode_vae, self.decode_vae),
VAEMode.TILED: (
self.encode_vae_tiled,
self.decode_vae_tiled,
),
VAEMode.TILED_DIFFUSION: (
self.encode_vae_tiled_diffusion,
self.decode_vae_tiled_diffusion,
),
}
self.encode_fun, self.decode_fun = vae_handlers[mode]
def encode(self, imgbatch, *, fix_dims=False):
if fix_dims:
imgbatch = imgbatch.moveaxis(1, -1)
# print("ENCODING", imgbatch.min(), imgbatch.max())
result = self.encode_fun(imgbatch[..., :3])
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):
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)
# print("DECODED", result.shape, result.min(), result.max())
def encode_taesd(self, imgbatch):
dummy = torch.zeros((), device=self.device, dtype=self.dtype)
return OCSTAESD.encode(self.latent_format, imgbatch, dummy)
def decode_taesd(self, latent):
return OCSTAESD.decode(self.latent_format, latent)
def encode_vae(self, imgbatch):
# print("VAE ENC", imgbatch.shape)
return self.vae.encode(imgbatch, **self.vae_encode_kwargs)
def decode_vae(self, latent):
return self.vae.decode(latent, **self.vae_decode_kwargs)
def encode_vae_tiled(self, imgbatch):
return self.vae.encode_tiled(imgbatch, **self.vae_encode_kwargs)
def decode_vae_tiled(self, latent):
return self.vae.decode_tiled(latent, **self.vae_decode_kwargs)
def encode_vae_tiled_diffusion(self, imgbatch):
kwargs = self.td_encode_default_kwargs | self.vae_encode_kwargs
return tiled_diffusion.tiled_vae.VAEEncodeTiled_TiledDiffusion().process(
pixels=imgbatch,
vae=self.vae,
**kwargs,
)[0]["samples"]
def decode_vae_tiled_diffusion(self, latent):
kwargs = self.td_decode_default_kwargs | self.vae_decode_kwargs
return tiled_diffusion.tiled_vae.VAEDecodeTiled_TiledDiffusion().process(
samples={"samples": latent},
vae=self.vae,
**kwargs,
)[0]
class OCSTAESD:
@classmethod
def get_encoder_name(cls, latent_format):
result = latent_format.taesd_decoder_name
if not result.endswith("_decoder"):
msg = f"Could not determine TAESD encoder name from {result!r}"
raise RuntimeError(
msg,
)
return f"{result[:-7]}encoder"
@classmethod
def get_taesd_path(cls, name):
taesd_path = next(
(
fn
for fn in folder_paths.get_filename_list("vae_approx")
if fn.startswith(name)
),
"",
)
if not taesd_path:
msg = f"Could not get TAESD path for {name!r}"
raise RuntimeError(msg)
return folder_paths.get_full_path("vae_approx", taesd_path)
@classmethod
def decode(cls, latent_format, latent):
filename = cls.get_taesd_path(latent_format.taesd_decoder_name)
model = TAESD(
decoder_path=filename,
latent_channels=latent_format.latent_channels,
).to(latent.device)
return (
model.taesd_decoder(
(latent - model.vae_shift).mul_(model.vae_scale),
)
.clamp_(0, 1)
.movedim(1, -1)
)
@classmethod
def encode(cls, latent_format, imgbatch, latent) -> torch.Tensor:
filename = cls.get_taesd_path(cls.get_encoder_name(latent_format))
model = TAESD(
encoder_path=filename,
latent_channels=latent_format.latent_channels,
).to(device=latent.device)
return (
model.taesd_encoder(imgbatch.to(latent.device).moveaxis(-1, 1))
.div_(model.vae_scale)
.add_(model.vae_shift)
)
+2
View File
@@ -0,0 +1,2 @@
pywavelets
pytorch-wavelets
+43
View File
@@ -0,0 +1,43 @@
[lint]
ignore = [
"ANN001",
"ANN101",
"ANN102",
"ANN201",
"ANN202",
"ANN204",
"ANN206",
"C901",
"CPY001",
"D100",
"D101",
"D102",
"D103",
"D104",
"D105",
"D107",
"D211",
"D213",
"E402",
"E501",
"EM101",
"ERA001",
"F403",
"F405",
"FBT001",
"FBT002",
"G004",
"PLR0912",
"PLR0913",
"PLR0915",
"PLR2004",
"PLR6104",
"T201",
"TD001",
"TD002",
"TD003",
"TRY003",
"N802",
"N999",
]
select = ["ALL"]