add remove latent mask

This commit is contained in:
matt3o
2024-02-17 13:07:18 +01:00
parent b3d606f416
commit 702ecc5baf
+137 -1
View File
@@ -1167,7 +1167,7 @@ class SDXLEmptyLatentSizePicker:
latent = torch.zeros([batch_size, 4, height // 8, width // 8], device=self.device)
return (latent, width, height,)
return ({"samples":latent}, width, height,)
LUTS_DIR = os.path.join(os.path.dirname(os.path.realpath(__file__)), "luts")
# From https://github.com/yoonsikp/pycubelut/blob/master/pycubelut.py (MIT license)
@@ -1352,6 +1352,134 @@ class ImageRemoveBackground:
return(output[:, :, :, :3], mask,)
class NoiseFromImage:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"noise_size": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01 }),
"color_noise": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01 }),
"mask_strength": ("FLOAT", {"default": 1.0, "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 }),
"noise_strenght": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01 }),
"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","IMAGE",)
FUNCTION = "execute"
CATEGORY = "essentials"
def execute(self, image, noise_size, color_noise, mask_strength, mask_scale_diff, 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
# Ensure noise mask is the same size as the image
if noise_mask.shape[1:] != image.shape[1:]:
noise_mask = F.interpolate(p(noise_mask), size=(image.shape[1], image.shape[2]), mode='bicubic', align_corners=False)
noise_mask = pb(noise_mask)
# 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 image to grayscale mask
noise_mask = noise_mask.mean(dim=3).unsqueeze(-1)
# add color noise
imgs = p(image.clone())
if color_noise > 0:
color_noise = torch.normal(torch.zeros_like(imgs), std=color_noise)
#color_noise = torch.rand_like(imgs) * (color_noise * 2) - color_noise
color_noise *= (imgs - imgs.min()) / (imgs.max() - imgs.min())
imgs = imgs + color_noise
imgs = imgs.clamp(0, 1)
# create fine 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)
#tmp_noise = T.functional.adjust_saturation(tmp_noise, 2.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)
#tmp_noise = F.interpolate(tmp_noise, scale_factor=.1, mode='bilinear', align_corners=False)
#tmp_noise = F.interpolate(tmp_noise, size=(tmp_noise.shape[1], tmp_noise.shape[2]), mode='bilinear', align_corners=False)
#tmp_noise = T.ElasticTransform(alpha=elastic_alpha, sigma=elastic_sigma/3, fill=avg_color.tolist())(n)
#tmp_noise = T.GaussianBlur(blur_size, blur)(tmp_noise)
#tmp_noise = T.functional.adjust_saturation(tmp_noise, saturation)
#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 = pb(fine_noise)
#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(p(fine_noise), 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 = pb(coarse_noise)
else:
coarse_noise = fine_noise
#noise_mask = noise_mask * mask_strength + (1 - mask_strength)
# merge fine and coarse noise
output = (1 - noise_mask) * coarse_noise + noise_mask * fine_noise
#noise_mask = noise_mask * mask_strength
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,noise_mask.repeat(1,1,1,3),)
class RemoveLatentMask:
@classmethod
def INPUT_TYPES(s):
return {"required": { "samples": ("LATENT",),}}
RETURN_TYPES = ("LATENT",)
FUNCTION = "execute"
CATEGORY = "essentials"
def execute(self, samples):
s = samples.copy()
if "noise_mask" in s:
del s["noise_mask"]
return (s,)
NODE_CLASS_MAPPINGS = {
"GetImageSize+": GetImageSize,
@@ -1394,6 +1522,10 @@ NODE_CLASS_MAPPINGS = {
"DrawText+": DrawText,
"RemBGSession+": RemBGSession,
"ImageRemoveBackground+": ImageRemoveBackground,
"RemoveLatentMask+": RemoveLatentMask,
#"NoiseFromImage~": NoiseFromImage,
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -1437,4 +1569,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DrawText+": "🔧 Draw Text",
"RemBGSession+": "🔧 RemBG Session",
"ImageRemoveBackground+": "🔧 Image Remove Background",
"RemoveLatentMask+": "🔧 Remove Latent Mask",
#"NoiseFromImage~": "🔧 Noise From Image",
}