Improvements to Deep Shrink node

This commit is contained in:
blepping
2024-02-01 12:20:15 -07:00
parent db408b6624
commit 06ab405a49
3 changed files with 46 additions and 12 deletions
+3
View File
@@ -59,6 +59,9 @@ following differences:
1. Instead of choosing a block to apply the downscale effect to, you can enter a comma-separated list of blocks. This may or not actually be useful but it seems like you can get interesting effects applying it to multiple blocks. Try `2,3` or `1,2,3`.
2. Adds a `start_fadeout_percent` input. When this is less than `end_percent` the downscale will be scaled to end at `end_percent`. For example, if `downscale_factor=2.0`, `start_percent=0.0`, `end_percent=0.5` and `start_fadeout_percent=0.0` then at 25% you could expect `downscale_factor` to be around `1.5`. This is because we are deep shrinking between 0 and 50% and we are halfway through the effect range. (`downscale_factor=1.0` would of course be a no-op and values below 1 don't seem to work.)
3. Expands the options for upscale and downscale types, you can also turn on antialiasing for `bicubic` and `bilinear` modes.
*Notes*: It seems like when shrinking multiple blocks, blocks downstream are also affected. So if you do x2 downscaling on 3 blocks, you are going to be applying `x2 * 3` downscaling to the lowest block (and maybe downstream ones?). I am not 100% sure how it works, but the takeway is you want to reduce the downscale amount when you are downscaling multiple blocks. For example, using blocks `2,3,4` and a downscale factor of `2.0` or `2.5` generating at 3072x3072 seems to work pretty well. Another note is schedulers that move at a steady pace seem to produce better results when fading out the deep shrink effect. In other words, exponential or Karras schedulers don't work well (and may produce complete nonsense). `ddim_uniform` and `sgm_uniform` seem to work pretty well and `normal` appear to be decent.
Deep Shrink credits:
+1
View File
@@ -5,6 +5,7 @@ Note, only relatively significant changes to user-visible functionality will be
## 20240201
* Added `BlehDeepShrink` node (see README for usage and description)
* Add more upscale/downscale methods to the Deep Shrink node, allow setting a higher downscale factor, allow enabling antialiasing for `bilinear` and `bicubic` modes.
## 20240128
+42 -12
View File
@@ -2,11 +2,19 @@
import bisect
from comfy.utils import common_upscale
import torch
from comfy.utils import bislerp
class DeepShrinkBleh:
upscale_methods = ("bicubic", "nearest-exact", "bilinear", "area", "bislerp")
upscale_methods = (
"bicubic",
"nearest-exact",
"bilinear",
"area",
"bislerp",
"linear",
)
@classmethod
def INPUT_TYPES(cls):
@@ -21,7 +29,7 @@ class DeepShrinkBleh:
),
"downscale_factor": (
"FLOAT",
{"default": 2.0, "min": 1.0, "max": 9.0, "step": 0.001},
{"default": 2.0, "min": 1.0, "max": 32.0, "step": 0.1},
),
"start_percent": (
"FLOAT",
@@ -38,6 +46,8 @@ class DeepShrinkBleh:
"downscale_after_skip": ("BOOLEAN", {"default": True}),
"downscale_method": (cls.upscale_methods,),
"upscale_method": (cls.upscale_methods,),
"antialias_downscale": ("BOOLEAN", {"default": False}),
"antialias_upscale": ("BOOLEAN", {"default": False}),
},
}
@@ -56,6 +66,8 @@ class DeepShrinkBleh:
downscale_after_skip,
downscale_method,
upscale_method,
antialias_downscale,
antialias_upscale,
):
block_numbers = tuple(
int(x) for x in commasep_block_numbers.split(",") if x.strip()
@@ -65,6 +77,14 @@ class DeepShrinkBleh:
raise ValueError(
"BlehDeepShrink: Bad value for block numbers: must be comma-separated list of numbers between 1-32",
)
antialias_downscale = antialias_downscale and downscale_method in (
"bicubic",
"bilinear",
)
antialias_upscale = antialias_upscale and upscale_method in (
"bicubic",
"bilinear",
)
if start_fadeout_percent < start_percent:
start_fadeout_percent = start_percent
elif start_fadeout_percent > end_percent:
@@ -109,23 +129,33 @@ class DeepShrinkBleh:
/ (end_percent - start_fadeout_percent)
)
scaled_scale = 1.0 - ((1.0 - downscale_factor) * downscale_pct)
return common_upscale(
h,
width, height = (
round(h.shape[-1] * scaled_scale),
round(h.shape[-2] * scaled_scale),
downscale_method,
"disabled",
)
if downscale_method == "bislerp":
return bislerp(h, width, height)
return torch.nn.functional.interpolate(
h,
size=(height, width),
mode=downscale_method,
antialias=antialias_downscale,
)
def output_block_patch(h, hsp, _transformer_options):
if h.shape[2] == hsp.shape[2]:
return h, hsp
return common_upscale(
if upscale_method == "bislerp":
return bislerp(
h,
hsp.shape[-1],
hsp.shape[-2],
), hsp
return torch.nn.functional.interpolate(
h,
hsp.shape[-1],
hsp.shape[-2],
upscale_method,
"disabled",
size=(hsp.shape[-2], hsp.shape[-1]),
mode=upscale_method,
antialias=antialias_upscale,
), hsp
m = model.clone()