Merge pull request #1 from blepping/initial_implementation
Initial implementation
This commit is contained in:
+4
-3
@@ -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__/
|
||||
|
||||
@@ -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
|
||||
@@ -1,6 +1,162 @@
|
||||
# ComfyUI jank DiffuseHigh
|
||||
Janky implementation of [DiffuseHigh](https://github.com/yhyun225/DiffuseHigh/) for ComfyUI
|
||||
Janky implementation of [DiffuseHigh](https://github.com/yhyun225/DiffuseHigh/) for ComfyUI.
|
||||
|
||||
Facilitates generating directly to resolutions higher than the model was trained for, similar to Kohya Deep Shrink, HiDiffusion, etc.
|
||||
|
||||
This is a best-effort attempt at implementation. If you experience poor results, please don't let it reflect on the official version. There's a good chance it's something I did wrong.
|
||||
|
||||
## Current Status
|
||||
|
||||
Alpha - early implementation. Many rough edges but the core functionality is there. Mainly targetted at advanced users who can deal with some weird stuff and frequent workflow-breaking changes.
|
||||
|
||||
**Known issues/caveats**
|
||||
|
||||
* There will be frequent workflow-breaking changes for a while yet.
|
||||
* Progress and previews are pretty wonky (you can look at the log for some progress information).
|
||||
* Using VAE or upscale models may result in the main model getting repeatedly unloaded/reloaded. Try using `latent` as the `guidance_mode`. If you actually have enough VRAM, maybe disabling smart memory (via ComfyUI commandline parameter) would help.
|
||||
* Currently only tested on SD15 and SDXL, may not work with models like Flux. (Not much testing in general as of yet.)
|
||||
|
||||
## Description
|
||||
|
||||
The DiffuseHigh approach is similar to an iterative upscale/run some more steps at low denoise approach with a twist: it mixes in guidance from a reference image for a number of steps at the beginning of each sampling iteration. The guidance is derived from the low frequency parts of the reference and also gets sharpened first to increase detail.
|
||||
|
||||
My approach implements it as a sampler which means it's _mostly_ model-agnostic and avoids some common issues with alternative approaches like Deep Shrink and HiDiffusion that require model patches. It's also possible to generate a low or mid-resolution image to see if you like the results and then increase the number of iterations to get a similar result where with Deep Shrink/HiDiffusion type effects enabling/disabling the patch will effectively change the seed.
|
||||
|
||||
The main disadvantage compared to the alternatives I mentioned is that it is relatively slow and VRAM hungry since it requires multiple iterations at high res while Deep Shrink/HiDiffusion actually speed up generation while the scaling effect is active.
|
||||
|
||||
## Nodes
|
||||
|
||||
### `DiffuseHighSampler`
|
||||
|
||||
#### Inputs
|
||||
|
||||
* `highres_sigmas`: Sigmas used for everything other than the initial reference image. **Note**: Should be around 0.3-0.5 denoise. You won't get good results connecting something like `KarrasScheduler` here without splitting the sigmas.
|
||||
* `sampler`: Default sampler used for steps. If not specified the sampler will default to non-ancestral Euler.
|
||||
* `reference_image_opt`: 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 (i.e. the sigmas attached to `SamplerCustom` or whatever actual sampler node you're using).
|
||||
* `guidance_sampler_opt`: 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`: Optional: Sampler used to generate the initial low-resolution reference. Only used if reference_image_opt is not connected.
|
||||
* `vae_opt`: Optional when vae_mode is set to `taesd`, otherwise this is the VAE that will be used for encoding/decoding images. If using TAESD, you will require the corresponding encoder (which I believe ComfyUI does not install by default). TAESD models available here: https://github.com/madebyollin/taesd
|
||||
* `upscale_model_opt`: 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`: Optional: 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. See below for some information about advanced parameters.
|
||||
|
||||
#### Parameters
|
||||
|
||||
* `guidance_steps`: Number of guidance steps after an upscale.
|
||||
* `guidance_mode`: 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. Personally I recommend setting this to `latent`.
|
||||
* `guidance_factor`: Mix factor used on guidance steps. 1.0 means use 100% DiffuseHigh guidance for those steps (like the original implementation).
|
||||
* `fadeout_factor`: 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_factor`s for the guidance steps: 1.00, 0.75, 0.50, 0.25
|
||||
* `scale_factor`: Upscale factor per iteration. The scaled size will be rounded to increments of 64 by default (can be adjusted via YAML parameters).
|
||||
* `renoise_factor`: 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. Something like `1.02` seems pretty good.
|
||||
* `iterations`: 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.
|
||||
* `vae_mode`: 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](https://github.com/shiimizu/ComfyUI-TiledDiffusion) installed you can use `tiled_diffusion` here.
|
||||
|
||||
#### YAML Parameters
|
||||
|
||||
<details>
|
||||
|
||||
<summary>Expand for advanced parameters</summary>
|
||||
|
||||
Note: JSON is also valid YAML so you can use that instead if you prefer.
|
||||
|
||||
You can also override normal parameters from the node. For example:
|
||||
```yaml
|
||||
iterations: 3
|
||||
scale_factor: 1.5
|
||||
```
|
||||
|
||||
*Note*: A lot of these parameters are experimental/just stuff to try for a different effect. Their existence doesn't necessarily mean enabling/changing the parameter will be better than the default.
|
||||
|
||||
Default advanced parameter values:
|
||||
|
||||
```yaml
|
||||
# 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
|
||||
# "image" can only be used when guidance_mode is also "image" - will fall back to "wavelets" in that case.
|
||||
blend_by_mode: "image"
|
||||
|
||||
# Multiplier on the denoised wavelets. This would be the high frequency component by default.
|
||||
denoised_wavelet_multiplier: 1.0
|
||||
|
||||
# See: https://pytorch-wavelets.readthedocs.io/en/latest/index.html
|
||||
# dtcwt_mode enables using DTCWT rather than the default DWT.
|
||||
dtcwt_biort: "near_sym_a"
|
||||
dtcwt_mode: false
|
||||
dtcwt_qshift: "qshift_a"
|
||||
dwt_level: 1
|
||||
dwt_mode: "symmetric"
|
||||
dwt_wave: "db4"
|
||||
|
||||
# Flips the highpass/lowpass filters. Normally the reference lowpass and denoised highpass parts
|
||||
# get used. If you flip them, you'll be using denoised for structural guidance and the reference
|
||||
# for the high-frequency part.
|
||||
dwt_flip_filters: false
|
||||
|
||||
# Number of times to restart guidance steps. (Does a restart back like restart sampling.)
|
||||
guidance_restart: 0
|
||||
# Factor for noise added during guidance restarts.
|
||||
guidance_restart_s_noise: 1.0
|
||||
|
||||
# Multiplier on the reference wavelets. This would be the low frequency component by default.
|
||||
reference_wavelet_multiplier: 1.0
|
||||
|
||||
# Mode used for simple image rescales. Probably the main alternative here is setting it to lanczos.
|
||||
# See: https://pillow.readthedocs.io/en/stable/handbook/concepts.html#filters-comparison-table
|
||||
resample_mode: "bicubic"
|
||||
|
||||
# Increment image sizes are rounded to. Must be at least 8 and a multiple of 8.
|
||||
rescale_increment: 64
|
||||
|
||||
# Mode used for sharpening. Can be one of: gaussian, contrast_adaptive
|
||||
# If using contrast_adaptive, I'd recommend setting sharpen_strength a bit lower.
|
||||
sharpen_mode: "gaussian"
|
||||
|
||||
# Allows disabling sharpening. Setting it to false is effectively the same as sharpening_strength: 0
|
||||
sharpen_reference: true
|
||||
|
||||
sharpen_gaussian_kernel_size: 3
|
||||
sharpen_gaussian_sigma: [0.1, 2.0]
|
||||
sharpen_strength: 1.0
|
||||
|
||||
# Disables the callback function (basically disables previews).
|
||||
skip_callback: false
|
||||
|
||||
# Allows specifying an offset into highres_sigmas.
|
||||
sigma_offset: 0
|
||||
|
||||
# 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 }
|
||||
vae_decode_kwargs: null
|
||||
vae_encode_kwargs: null
|
||||
|
||||
# Either null or an object.
|
||||
# Allows overriding parameters per iteration. See description below.
|
||||
iteration_override: null
|
||||
```
|
||||
|
||||
**Iteration Overrides**
|
||||
|
||||
Example:
|
||||
|
||||
```yaml
|
||||
iteration_override:
|
||||
0:
|
||||
scale_factor: 2.0
|
||||
1:
|
||||
scale_factor: 1.5
|
||||
skip_callback: true
|
||||
```
|
||||
|
||||
You can override most parameters this way. Exceptions: Node inputs, `iteration_override` itself and `iterations`.
|
||||
|
||||
The `iteration_overrides` should either be `null` (disabled) or a YAML object with the iteration number (note: zero-based) as the key which contains an object with parameters in the same format as the main YAML parameters. Can be used to vary `scale_factor` across iterations, switched to tiled VAE only when the image is large enough for it to be worthwhile, disable previews (via `skip_callback: false`) if you're running out of memory at high res, etc.
|
||||
|
||||
|
||||
</details>
|
||||
|
||||
***
|
||||
|
||||
**Coming Soon**
|
||||
## Credits
|
||||
|
||||
Heavily referenced from the official implementation: [DiffuseHigh](https://github.com/yhyun225/DiffuseHigh/)
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from .py import nodes
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DiffuseHighSampler": nodes.DiffuseHighSamplerNode,
|
||||
}
|
||||
+213
@@ -0,0 +1,213 @@
|
||||
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",
|
||||
"reference_wavelet_multiplier",
|
||||
"renoise_factor",
|
||||
"resample_mode",
|
||||
"rescale_increment",
|
||||
"scale_factor",
|
||||
"sharpen_gaussian_kernel_size",
|
||||
"sharpen_gaussian_sigma",
|
||||
"sharpen_mode",
|
||||
"sharpen_reference",
|
||||
"sharpen_strength",
|
||||
"skip_callback",
|
||||
"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,
|
||||
skip_callback=False,
|
||||
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.skip_callback = skip_callback
|
||||
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")
|
||||
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
|
||||
@@ -0,0 +1,16 @@
|
||||
import contextlib
|
||||
import importlib
|
||||
|
||||
EXTERNAL = {}
|
||||
|
||||
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
|
||||
+161
@@ -0,0 +1,161 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> dict:
|
||||
return {
|
||||
"required": {
|
||||
"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",
|
||||
{
|
||||
"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.",
|
||||
},
|
||||
),
|
||||
"vae_mode": (
|
||||
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": {
|
||||
"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",
|
||||
{
|
||||
"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,
|
||||
},
|
||||
),
|
||||
)
|
||||
+276
@@ -0,0 +1,276 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
from tqdm.auto import trange
|
||||
|
||||
from .config import Config
|
||||
from .tensor_image_ops import (
|
||||
blend_wavelets,
|
||||
scale_wavelets,
|
||||
)
|
||||
from .utils import ensure_model, fallback
|
||||
|
||||
|
||||
class DiffuseHighSampler:
|
||||
def __init__(
|
||||
self,
|
||||
model,
|
||||
initial_x,
|
||||
sigmas,
|
||||
*,
|
||||
callback,
|
||||
extra_args,
|
||||
disable_pbar,
|
||||
highres_sigmas,
|
||||
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 = fallback(extra_args, {})
|
||||
self.model = model
|
||||
self.latent_format = model.inner_model.inner_model.latent_format
|
||||
self.config = self.base_config = Config(
|
||||
initial_x.device,
|
||||
initial_x.dtype,
|
||||
self.latent_format,
|
||||
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.guidance_waves = None
|
||||
|
||||
def __getattr__(self, key):
|
||||
return getattr(self.config, key)
|
||||
|
||||
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) * idx) * self.fadeout_factor
|
||||
)
|
||||
if mix_scale == 0:
|
||||
return denoised
|
||||
if self.guidance_mode not in {"image", "latent"}:
|
||||
raise ValueError("Bad guidance mode")
|
||||
if self.guidance_mode == "image":
|
||||
dn_img = (
|
||||
self.vae.decode(denoised, disable_pbar=self.disable_pbar)
|
||||
.to(denoised)
|
||||
.movedim(-1, 1)
|
||||
)
|
||||
denoised_waves = self.dwt(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 = 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":
|
||||
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
|
||||
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
|
||||
|
||||
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 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
|
||||
) ** 0.5
|
||||
x = x + torch.randn_like(x) * (
|
||||
noise_factor * self.guidance_restart_s_noise
|
||||
)
|
||||
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)
|
||||
with tqdm(disable=self.disable_pbar, total=1, desc="normal steps") as pbar:
|
||||
x = self.run_sampler(x, normal_sigmas)
|
||||
pbar.update()
|
||||
return x
|
||||
|
||||
def run_sampler(self, x, sigmas, *, model=None, sampler=None, disable_pbar=False):
|
||||
sampler = fallback(sampler, self.sampler)
|
||||
return sampler.sampler_function(
|
||||
fallback(model, self.model),
|
||||
x,
|
||||
sigmas,
|
||||
callback=self.callback if not self.skip_callback else None,
|
||||
extra_args=self.extra_args.copy(),
|
||||
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:
|
||||
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, disable_pbar=self.disable_pbar)
|
||||
elif self.iterations < 1:
|
||||
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)
|
||||
if self.config.sigma_offset >= len(self.highres_sigmas) - 1:
|
||||
raise ValueError(
|
||||
"Bad sigma_offset: posts to sigma past penultimate sigma",
|
||||
)
|
||||
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":
|
||||
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(x_new.clone())
|
||||
if self.reference_wavelet_multiplier != 1:
|
||||
self.guidance_waves = scale_wavelets(
|
||||
self.guidance_waves,
|
||||
self.reference_wavelet_multiplier,
|
||||
)
|
||||
else:
|
||||
raise ValueError("Bad guidance_mode")
|
||||
x_noise = torch.randn_like(x_new)
|
||||
x_new = x_new + x_noise * (
|
||||
self.highres_sigmas[self.sigma_offset] * self.renoise_factor
|
||||
)
|
||||
del x_noise
|
||||
# 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,
|
||||
disable_pbar=self.disable_pbar,
|
||||
)
|
||||
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()
|
||||
@@ -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 and self.upscale_model is None:
|
||||
pbar.set_description(
|
||||
f"upscale (simple): {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)
|
||||
),
|
||||
)
|
||||
+47
@@ -0,0 +1,47 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
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,))
|
||||
|
||||
|
||||
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
|
||||
@@ -0,0 +1,196 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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")
|
||||
|
||||
|
||||
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,
|
||||
encode_kwargs=None,
|
||||
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.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),
|
||||
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, disable_pbar=None):
|
||||
if fix_dims:
|
||||
imgbatch = imgbatch.moveaxis(1, -1)
|
||||
# print("ENCODING", imgbatch.min(), imgbatch.max())
|
||||
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, 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())
|
||||
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):
|
||||
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.encode_kwargs)
|
||||
|
||||
def decode_vae(self, latent):
|
||||
return self.vae.decode(latent, **self.decode_kwargs)
|
||||
|
||||
def encode_vae_tiled(self, imgbatch):
|
||||
return self.vae.encode_tiled(imgbatch, **self.encode_kwargs)
|
||||
|
||||
def decode_vae_tiled(self, latent):
|
||||
return self.vae.decode_tiled(latent, **self.decode_kwargs)
|
||||
|
||||
def encode_vae_tiled_diffusion(self, imgbatch):
|
||||
kwargs = self.td_encode_default_kwargs | self.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.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)
|
||||
)
|
||||
@@ -0,0 +1,2 @@
|
||||
pywavelets
|
||||
pytorch-wavelets
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user