update tile, add untile, other small fixes

This commit is contained in:
cubiq
2024-07-08 18:18:34 +02:00
parent 8ee40fd6ca
commit 442230f035
2 changed files with 212 additions and 18 deletions
+205 -14
View File
@@ -243,8 +243,8 @@ class ImageResize:
return {
"required": {
"image": ("IMAGE",),
"width": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 8, }),
"height": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 8, }),
"width": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 1, }),
"height": ("INT", { "default": 512, "min": 0, "max": MAX_RESOLUTION, "step": 1, }),
"interpolation": (["nearest", "bilinear", "bicubic", "area", "nearest-exact", "lanczos"],),
"method": (["stretch", "keep proportion", "fill / crop", "pad"],),
"condition": (["always", "downscale if bigger", "upscale if smaller", "if bigger area", "if smaller area"],),
@@ -442,22 +442,34 @@ class ImageTile:
"rows": ("INT", { "default": 2, "min": 1, "max": 256, "step": 1, }),
"cols": ("INT", { "default": 2, "min": 1, "max": 256, "step": 1, }),
"overlap": ("FLOAT", { "default": 0, "min": 0, "max": 0.5, "step": 0.01, }),
"overlap_x": ("INT", { "default": 0, "min": 0, "max": MAX_RESOLUTION//2, "step": 1, }),
"overlap_y": ("INT", { "default": 0, "min": 0, "max": MAX_RESOLUTION//2, "step": 1, }),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_TYPES = ("IMAGE", "INT", "INT", "INT", "INT")
RETURN_NAMES = ("IMAGE", "tile_width", "tile_height", "overlap_x", "overlap_y",)
FUNCTION = "execute"
CATEGORY = "essentials/image manipulation"
def execute(self, image, rows, cols, overlap):
def execute(self, image, rows, cols, overlap, overlap_x, overlap_y):
h, w = image.shape[1:3]
tile_h = h // rows
tile_w = w // cols
overlap_h = int(tile_h * overlap)
overlap_w = int(tile_w * overlap)
tile_h += overlap_h
tile_w += overlap_w
h = tile_h * rows
w = tile_w * cols
overlap_h = int(tile_h * overlap) + overlap_y
overlap_w = int(tile_w * overlap) + overlap_x
# max overlap is half of the tile size
overlap_h = min(tile_h // 2, overlap_h)
overlap_w = min(tile_w // 2, overlap_w)
if rows == 1:
overlap_h = 0
if cols == 1:
overlap_w = 0
tiles = []
for i in range(rows):
for j in range(cols):
@@ -469,21 +481,93 @@ class ImageTile:
if j > 0:
x1 -= overlap_w
y2 = y1 + tile_h
x2 = x1 + tile_w
y2 = y1 + tile_h + overlap_h
x2 = x1 + tile_w + overlap_w
if y2 > h:
y2 = h
y1 = y2 - tile_h
y1 = y2 - tile_h - overlap_h
if x2 > w:
x2 = w
x1 = x2 - tile_w
x1 = x2 - tile_w - overlap_w
tiles.append(image[:, y1:y2, x1:x2, :])
tiles = torch.cat(tiles, dim=0)
return(tiles,)
return(tiles, tile_w+overlap_w, tile_h+overlap_h, overlap_w, overlap_h,)
class ImageUntile:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"tiles": ("IMAGE",),
"overlap_x": ("INT", { "default": 0, "min": 0, "max": MAX_RESOLUTION//2, "step": 1, }),
"overlap_y": ("INT", { "default": 0, "min": 0, "max": MAX_RESOLUTION//2, "step": 1, }),
"rows": ("INT", { "default": 2, "min": 1, "max": 256, "step": 1, }),
"cols": ("INT", { "default": 2, "min": 1, "max": 256, "step": 1, }),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
CATEGORY = "essentials/image manipulation"
def execute(self, tiles, overlap_x, overlap_y, rows, cols):
tile_h, tile_w = tiles.shape[1:3]
tile_h -= overlap_y
tile_w -= overlap_x
out_w = cols * tile_w
out_h = rows * tile_h
out = torch.zeros((1, out_h, out_w, tiles.shape[3]), device=tiles.device, dtype=tiles.dtype)
for i in range(rows):
for j in range(cols):
y1 = i * tile_h
x1 = j * tile_w
if i > 0:
y1 -= overlap_y
if j > 0:
x1 -= overlap_x
y2 = y1 + tile_h + overlap_y
x2 = x1 + tile_w + overlap_x
if y2 > out_h:
y2 = out_h
y1 = y2 - tile_h - overlap_y
if x2 > out_w:
x2 = out_w
x1 = x2 - tile_w - overlap_x
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:
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:
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:
# mask[:, :, -overlap_x:] *= torch.linspace(1, 0, overlap_x, device=tiles.device, dtype=tiles.dtype).unsqueeze(0)
mask = mask.unsqueeze(-1).repeat(1, 1, 1, tiles.shape[3])
tile = tiles[i * cols + j] * mask
# apply the mask to the tile
out[:, y1:y2, x1:x2, :] = out[:, y1:y2, x1:x2, :] * (1 - mask) + tile
#out[:, y1:y2, x1:x2, :] = tiles[i * cols + j]
return(out, )
class ImageSeamCarving:
@classmethod
@@ -1138,6 +1222,109 @@ class ImagePreviewFromLatent(SaveImage):
out = {**ui, "result": (image, mask, image.shape[2], image.shape[1],)}
return out
class NoiseFromImage:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"noise_strenght": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01 }),
"noise_size": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01 }),
"color_noise": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 1.0, "step": 0.01 }),
"mask_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01 }),
"mask_scale_diff": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01 }),
"mask_contrast": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1 }),
"saturation": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.1 }),
"contrast": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1 }),
"blur": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1 }),
},
"optional": {
"noise_mask": ("IMAGE",),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
CATEGORY = "essentials"
def execute(self, image, noise_size, color_noise, mask_strength, mask_scale_diff, mask_contrast, noise_strenght, saturation, contrast, blur, noise_mask=None):
torch.manual_seed(0)
elastic_alpha = max(image.shape[1], image.shape[2])# * noise_size
elastic_sigma = elastic_alpha / 400 * noise_size
blur_size = int(6 * blur+1)
if blur_size % 2 == 0:
blur_size+= 1
if noise_mask is None:
noise_mask = image
# increase contrast of the mask
if mask_contrast != 1:
noise_mask = T.ColorJitter(contrast=(mask_contrast,mask_contrast))(noise_mask.permute([0, 3, 1, 2])).permute([0, 2, 3, 1])
# Ensure noise mask is the same size as the image
if noise_mask.shape[1:] != image.shape[1:]:
noise_mask = F.interpolate(noise_mask.permute([0, 3, 1, 2]), size=(image.shape[1], image.shape[2]), mode='bicubic', align_corners=False)
noise_mask = noise_mask.permute([0, 2, 3, 1])
# Ensure we have the same number of masks and images
if noise_mask.shape[0] > image.shape[0]:
noise_mask = noise_mask[:image.shape[0]]
else:
noise_mask = torch.cat((noise_mask, noise_mask[-1:].repeat((image.shape[0]-noise_mask.shape[0], 1, 1, 1))), dim=0)
# Convert mask to grayscale mask
noise_mask = noise_mask.mean(dim=3).unsqueeze(-1)
# add color noise
imgs = image.clone().permute([0, 3, 1, 2])
if color_noise > 0:
color_noise = torch.normal(torch.zeros_like(imgs), std=color_noise)
color_noise *= (imgs - imgs.min()) / (imgs.max() - imgs.min())
imgs = imgs + color_noise
imgs = imgs.clamp(0, 1)
# create fine and coarse noise
fine_noise = []
for n in imgs:
avg_color = n.mean(dim=[1,2])
tmp_noise = T.ElasticTransform(alpha=elastic_alpha, sigma=elastic_sigma, fill=avg_color.tolist())(n)
if blur > 0:
tmp_noise = T.GaussianBlur(blur_size, blur)(tmp_noise)
tmp_noise = T.ColorJitter(contrast=(contrast,contrast), saturation=(saturation,saturation))(tmp_noise)
fine_noise.append(tmp_noise)
imgs = None
del imgs
fine_noise = torch.stack(fine_noise, dim=0)
fine_noise = fine_noise.permute([0, 2, 3, 1])
#fine_noise = torch.stack(fine_noise, dim=0)
#fine_noise = pb(fine_noise)
mask_scale_diff = min(mask_scale_diff, 0.99)
if mask_scale_diff > 0:
coarse_noise = F.interpolate(fine_noise.permute([0, 3, 1, 2]), scale_factor=1-mask_scale_diff, mode='area')
coarse_noise = F.interpolate(coarse_noise, size=(fine_noise.shape[1], fine_noise.shape[2]), mode='bilinear', align_corners=False)
coarse_noise = coarse_noise.permute([0, 2, 3, 1])
else:
coarse_noise = fine_noise
output = (1 - noise_mask) * coarse_noise + noise_mask * fine_noise
if mask_strength < 1:
noise_mask = noise_mask.pow(mask_strength)
noise_mask = torch.nan_to_num(noise_mask).clamp(0, 1)
output = noise_mask * output + (1 - noise_mask) * image
# apply noise to image
output = output * noise_strenght + image * (1 - noise_strenght)
output = output.clamp(0, 1)
return (output, )
IMAGE_CLASS_MAPPINGS = {
# Image analysis
"ImageEnhanceDifference+": ImageEnhanceDifference,
@@ -1158,6 +1345,7 @@ IMAGE_CLASS_MAPPINGS = {
"ImageResize+": ImageResize,
"ImageSeamCarving+": ImageSeamCarving,
"ImageTile+": ImageTile,
"ImageUntile+": ImageUntile,
"RemBGSession+": RemBGSession,
# Image processing
@@ -1173,6 +1361,7 @@ IMAGE_CLASS_MAPPINGS = {
"GetImageSize+": GetImageSize,
"ImageToDevice+": ImageToDevice,
"ImagePreviewFromLatent+": ImagePreviewFromLatent,
"NoiseFromImage+": NoiseFromImage,
#"ExtractKeyframes+": ExtractKeyframes,
}
@@ -1197,6 +1386,7 @@ IMAGE_NAME_MAPPINGS = {
"ImageResize+": "🔧 Image Resize",
"ImageSeamCarving+": "🔧 Image Seam Carving",
"ImageTile+": "🔧 Image Tile",
"ImageUntile+": "🔧 Image Untile",
"RemBGSession+": "🔧 RemBG Session",
# Image processing
@@ -1212,4 +1402,5 @@ IMAGE_NAME_MAPPINGS = {
"GetImageSize+": "🔧 Get Image Size",
"ImageToDevice+": "🔧 Image To Device",
"ImagePreviewFromLatent+": "🔧 Image Preview From Latent",
"NoiseFromImage+": "🔧 Noise From Image",
}
+7 -4
View File
@@ -2,6 +2,7 @@ import math
import torch
from .utils import AnyType
import comfy.model_management
from nodes import MAX_RESOLUTION
any = AnyType("*")
@@ -203,17 +204,19 @@ class SDXLEmptyLatentSizePicker:
return {"required": {
"resolution": (["704x1408 (0.5)","704x1344 (0.52)","768x1344 (0.57)","768x1280 (0.6)","832x1216 (0.68)","832x1152 (0.72)","896x1152 (0.78)","896x1088 (0.82)","960x1088 (0.88)","960x1024 (0.94)","1024x1024 (1.0)","1024x960 (1.07)","1088x960 (1.13)","1088x896 (1.21)","1152x896 (1.29)","1152x832 (1.38)","1216x832 (1.46)","1280x768 (1.67)","1344x768 (1.75)","1344x704 (1.91)","1408x704 (2.0)","1472x704 (2.09)","1536x640 (2.4)","1600x640 (2.5)","1664x576 (2.89)","1728x576 (3.0)",], {"default": "1024x1024 (1.0)"}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096}),
"width_override": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 8}),
"height_override": ("INT", {"default": 0, "min": 0, "max": MAX_RESOLUTION, "step": 8}),
}}
RETURN_TYPES = ("LATENT","INT","INT",)
RETURN_NAMES = ("LATENT","width", "height",)
RETURN_NAMES = ("LATENT","width","height",)
FUNCTION = "execute"
CATEGORY = "essentials/utilities"
def execute(self, resolution, batch_size):
def execute(self, resolution, batch_size, width_override, height_override):
width, height = resolution.split(" ")[0].split("x")
width = int(width)
height = int(height)
width = width_override if width_override > 0 else int(width)
height = height_override if height_override > 0 else int(height)
latent = torch.zeros([batch_size, 4, height // 8, width // 8], device=self.device)