From 06ab405a49e5c8c23ef755d682b13899acc241a9 Mon Sep 17 00:00:00 2001 From: blepping Date: Thu, 1 Feb 2024 12:20:15 -0700 Subject: [PATCH] Improvements to Deep Shrink node --- README.md | 3 +++ changelog.md | 1 + py/deepshrink.py | 54 +++++++++++++++++++++++++++++++++++++----------- 3 files changed, 46 insertions(+), 12 deletions(-) diff --git a/README.md b/README.md index 92a9953..0b702dd 100644 --- a/README.md +++ b/README.md @@ -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: diff --git a/changelog.md b/changelog.md index a0d2885..20dfd23 100644 --- a/changelog.md +++ b/changelog.md @@ -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 diff --git a/py/deepshrink.py b/py/deepshrink.py index c3b5e81..f15f5ad 100644 --- a/py/deepshrink.py +++ b/py/deepshrink.py @@ -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()