lower memory usage of ImageSmartSharpen

This commit is contained in:
cubiq
2024-08-10 12:23:52 +02:00
parent 0f34459116
commit a94f25989f
+31 -31
View File
@@ -631,13 +631,13 @@ class ImageUntile:
mask = torch.ones((1, tile_h+overlap_y, tile_w+overlap_x), device=tiles.device, dtype=tiles.dtype)
# feather the overlap on top
if i > 0:
if i > 0 and overlap_y > 0:
mask[:, :overlap_y, :] *= torch.linspace(0, 1, overlap_y, device=tiles.device, dtype=tiles.dtype).unsqueeze(1)
# feather the overlap on bottom
#if i < rows - 1:
# mask[:, -overlap_y:, :] *= torch.linspace(1, 0, overlap_y, device=tiles.device, dtype=tiles.dtype).unsqueeze(1)
# feather the overlap on left
if j > 0:
if j > 0 and overlap_x > 0:
mask[:, :, :overlap_x] *= torch.linspace(0, 1, overlap_x, device=tiles.device, dtype=tiles.dtype).unsqueeze(0)
# feather the overlap on right
#if j < cols - 1:
@@ -1053,45 +1053,45 @@ class ImageSmartSharpen:
return {
"required": {
"image": ("IMAGE",),
"remove_noise": ("FLOAT", { "default": 0.5, "min": 0.0, "max": 1.0, "step": 0.05, }),
"preserve_edges": ("FLOAT", { "default": 0.5, "min": 0.0, "max": 1.0, "step": 0.05 }),
"noise_radius": ("INT", { "default": 7, "min": 1, "max": 25, "step": 1, }),
"preserve_edges": ("FLOAT", { "default": 0.75, "min": 0.0, "max": 1.0, "step": 0.05 }),
"sharpen": ("FLOAT", { "default": 5.0, "min": 0.0, "max": 25.0, "step": 0.5 }),
"ratio": ("FLOAT", { "default": 0.5, "min": 0.0, "max": 1.0, "step": 0.1 }),
"device": (["auto", "cpu", "gpu"],),
}}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "essentials/image processing"
FUNCTION = "execute"
def execute(self, image, remove_noise, preserve_edges, sharpen, ratio, device):
if "gpu" == device:
device = comfy.model_management.get_torch_device()
elif "auto" == device:
device = comfy.model_management.intermediate_device()
else:
device = 'cpu'
def execute(self, image, noise_radius, preserve_edges, sharpen, ratio):
import cv2
output = []
#diagonal = np.sqrt(image.shape[1]**2 + image.shape[2]**2)
if preserve_edges > 0:
preserve_edges = max(1 - preserve_edges, 0.05)
for img in image:
if noise_radius > 1:
sigma = 0.3 * ((noise_radius - 1) * 0.5 - 1) + 0.8 # this is what pytorch uses for blur
#sigma_color = preserve_edges * (diagonal / 2048)
blurred = cv2.bilateralFilter(img.cpu().numpy(), noise_radius, preserve_edges, sigma)
blurred = torch.from_numpy(blurred)
else:
blurred = img
if sharpen > 0:
sharpened = kornia.enhance.sharpness(img.permute(2,0,1), sharpen).permute(1,2,0)
else:
sharpened = img
img = ratio * sharpened + (1 - ratio) * blurred
img = torch.clamp(img, 0, 1)
output.append(img)
image = image.permute(0,3,1,2).to(device)
del blurred, sharpened
output = torch.stack(output)
if remove_noise > 0:
remove_noise = int(remove_noise * 15)
sigma = 0.3 * ((remove_noise - 1) * 0.5 - 1) + 0.8
preserve_edges = max(1.0 - preserve_edges, 0.1)
if remove_noise % 2 == 0:
remove_noise += 1
blurred = kornia.filters.bilateral_blur(image, remove_noise, preserve_edges, (sigma, sigma))
else:
blurred = image
if sharpen > 0:
sharpened = kornia.enhance.sharpness(image, sharpen)
else:
sharpened = image
output = ratio * sharpened + (1 - ratio) * blurred
output = torch.clamp(output, 0, 1)
output = output.permute(0,2,3,1).to(comfy.model_management.intermediate_device())
return (output,)