Improvements to contrast-adaptive sharpening mode

This commit is contained in:
blepping
2024-10-23 14:38:41 -06:00
parent 64244b00eb
commit 3dfaf18353
3 changed files with 43 additions and 27 deletions
+7 -2
View File
@@ -130,8 +130,9 @@ seed_rng: true
# as the initial noise.
seed_rng_offset: 1
# Mode used for sharpening. Can be one of: gaussian, contrast_adaptive
# If using contrast_adaptive, I'd recommend setting sharpen_strength a bit lower.
# Mode used for sharpening. Can be one of:
# gaussian, contrast_adaptive, contrast_adaptive_raw
# If using contrast_adaptive_raw, I'd recommend setting sharpen_strength a bit lower.
sharpen_mode: "gaussian"
sharpen_gaussian_kernel_size: 3
@@ -243,3 +244,7 @@ I tried to set the node defaults to align with the official implementation. Thes
## Credits
Heavily referenced from the official implementation: [DiffuseHigh](https://github.com/yhyun225/DiffuseHigh/)
Contrast-adaptive sharpening sources: [1](https://github.com/GPUOpen-Effects/FidelityFX-CAS/blob/master/ffx-cas/ffx_cas.h), [2](https://github.com/Jamy-L/Pytorch-Contrast-Adaptive-Sharpening/), [3](https://github.com/Clybius)
Thanks!
+1
View File
@@ -5,6 +5,7 @@ Note, only relatively significant changes to user-visible functionality will be
## 20241023
* Initial support for rectified flow models (Flux, SD3, SD3.5). Might slightly change seeds for other models.
* Improve contrast adaptive sharpening (hopefully). Will change seeds for workflows using `sharpen_mode: contrast_adaptive`, you can use `sharpen_mode: contrast_adaptive_raw` for the old behavior.
## 20241019
+35 -25
View File
@@ -28,6 +28,7 @@ def init_integrations():
class SharpenMode(Enum):
GAUSSIAN = auto()
CONTRAST_ADAPTIVE = auto()
CONTRAST_ADAPTIVE_RAW = auto()
def pilimgbatch_to_torch(
@@ -93,8 +94,15 @@ class Sharpen:
sigma=self.gaussian_sigma,
alpha=self.strength,
)
elif self.mode == SharpenMode.CONTRAST_ADAPTIVE:
result = contrast_adaptive_sharpening(t, amount=self.strength)
elif self.mode in {
SharpenMode.CONTRAST_ADAPTIVE,
SharpenMode.CONTRAST_ADAPTIVE_RAW,
}:
result = contrast_adaptive_sharpening(
t,
amount=self.strength,
normalize=self.mode == SharpenMode.CONTRAST_ADAPTIVE,
)
if fix_dims:
result = result.movedim(1, -1)
return result
@@ -109,33 +117,30 @@ def gaussian_blur_image_sharpening(image, kernel_size=3, sigma=(0.1, 2.0), alpha
return (alpha + 1) * image - alpha * image_blurred
# Improvements added by https://github.com/Clybius
# Improvements by https://github.com/Clybius
# 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
# 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.
def contrast_adaptive_sharpening( # noqa: PLR0914
x,
amount=0.8,
*,
normalize=True,
epsilon=1e-06,
):
if x.ndim != 4:
raise ValueError(
"Contrast-adaptive sharpening requires a tensor with 4 dimensions",
)
def on_abs_stacked(tensor_list, f, *args: list, **kwargs: dict):
return f(torch.abs(torch.stack(tensor_list)), *args, **kwargs)[0]
if normalize:
luminance = torch.linalg.vector_norm(x, dim=1, keepdim=True).add_(1e-08)
x = x / luminance
orig_mean = x.mean(dim=(-3, -2, -1), keepdim=True)
x -= orig_mean
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
@@ -187,4 +192,9 @@ def contrast_adaptive_sharpening(x, amount=0.8, *, epsilon=1e-06): # noqa: D417
div = torch.reciprocal(1 + 4 * w)
output = ((b + d + f + h) * w + e) * div
return output.real.clamp(x.min(), x.max())
output = output.real
for ob, xb in zip(x, output):
ob.clamp_(*xb.aminmax())
if normalize:
output = output.add_(orig_mean).mul_(luminance)
return output