From 810cfd13f044420283603e0de9d074a8c36cd5b7 Mon Sep 17 00:00:00 2001 From: blepping Date: Tue, 14 May 2024 18:43:51 -0600 Subject: [PATCH] Improve tensor enhancement... maybe --- __init__.py | 2 +- py/latent_utils.py | 58 +++++++++++++++++++++++++++++----------------- 2 files changed, 38 insertions(+), 22 deletions(-) diff --git a/__init__.py b/__init__.py index 4b7a47e..e418eac 100644 --- a/__init__.py +++ b/__init__.py @@ -1,6 +1,6 @@ from .py import settings -BLEH_VERSION = 0 +BLEH_VERSION = 1 settings.load_settings() diff --git a/py/latent_utils.py b/py/latent_utils.py index d22d3d0..aa852e7 100644 --- a/py/latent_utils.py +++ b/py/latent_utils.py @@ -193,6 +193,26 @@ BIDERP_MODES |= { "revbibislerp": BLENDING_MODES["revbislerp"], } +ENHANCE_METHODS = ( + "lowpass", + "highpass", + "bandpass", + "randhilowpass", + "randmultihilowpass", + "randhibandpass", + "randlowbandpass", + "gaussianblur", + "edge", + "sharpen", + "korniabilateralblur", + "korniagaussianblur", + "korniasharpen", + "korniaedge", + "korniarevedge", + "korniarandblursharp", + "renoise1", + "renoise2", +) UPSCALE_METHODS = ( "bicubic", @@ -203,26 +223,7 @@ UPSCALE_METHODS = ( *( f"{meth}+{enh}" for meth in ("bicubic", "bislerp", "hslerp", "random") - for enh in ( - "lowpass", - "highpass", - "bandpass", - "randhilowpass", - "randmultihilowpass", - "randhibandpass", - "randlowbandpass", - "gaussianblur", - "edge", - "sharpen", - "korniabilateralblur", - "korniagaussianblur", - "korniasharpen", - "korniaedge", - "korniarevedge", - "korniarandblursharp", - "renoise1", - "renoise2", - ) + for enh in ENHANCE_METHODS ), "random", "randomaa", @@ -261,8 +262,18 @@ def antialias_tensor(x, antialias_size): return torch.nn.functional.conv2d(x, filt, groups=channels, padding="same") -def enhance_tensor(x, name, scale=1.0, sigma=None): # noqa: PLR0911 +def enhance_tensor( + x, + name, + scale=1.0, + sigma=None, + *, + skip_multiplier=1, + adjust_scale=True, +): randitems = None + orig_scale = scale + randskip = 0 match name: case "randmultihilowpass": scale *= 0.1 @@ -294,6 +305,9 @@ def enhance_tensor(x, name, scale=1.0, sigma=None): # noqa: PLR0911 return x noise = torch.randn_like(x) return noise.mul_(noise_scale).add_(x) + if not adjust_scale: + scale = orig_scale + randskip = int(randskip * skip_multiplier) if randitems: ridx = torch.randint(len(randitems) + randskip, (1,), device="cpu").item() if ridx >= len(randitems): @@ -301,6 +315,8 @@ def enhance_tensor(x, name, scale=1.0, sigma=None): # noqa: PLR0911 return enhance_tensor(x, randitems[ridx], scale=scale) fpreset = FILTER_PRESETS.get(name) if fpreset is not None: + if not adjust_scale: + scale *= 2 return ffilter(x, 1, 1.0, fpreset, 0.5 * scale) match name: case "korniabilateralblur":