From c26bc23d53c1953dca5697b211d6f9a51c4d607e Mon Sep 17 00:00:00 2001 From: matt3o Date: Tue, 14 May 2024 11:56:14 +0200 Subject: [PATCH] add clipseg node --- essentials.py | 83 ++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 82 insertions(+), 1 deletion(-) diff --git a/essentials.py b/essentials.py index 3f0b2b8..c662a98 100644 --- a/essentials.py +++ b/essentials.py @@ -415,7 +415,7 @@ class MaskBlur: CATEGORY = "essentials" def execute(self, mask, amount): - size = int(6 * amount +1) + size = int(6 * amount + 1) if size % 2 == 0: size+= 1 @@ -1905,6 +1905,81 @@ class ImageBatchMultiple: return (out,) +class LoadCLIPSegModels: + @classmethod + def INPUT_TYPES(s): + return { + "required": {}, + } + + RETURN_TYPES = ("CLIP_SEG",) + FUNCTION = "execute" + CATEGORY = "essentials" + + def execute(self): + from transformers import CLIPSegProcessor, CLIPSegForImageSegmentation + processor = CLIPSegProcessor.from_pretrained("CIDAS/clipseg-rd64-refined") + model = CLIPSegForImageSegmentation.from_pretrained("CIDAS/clipseg-rd64-refined") + + return ((processor, model),) + +class ApplyCLIPSeg: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "clip_seg": ("CLIP_SEG",), + "prompt": ("STRING", { "multiline": False, "default": "" }), + "threshold": ("FLOAT", { "default": 0.4, "min": 0.0, "max": 1.0, "step": 0.05 }), + "smooth": ("INT", { "default": 9, "min": 0, "max": 32, "step": 1 }), + "dilate": ("INT", { "default": 0, "min": -32, "max": 32, "step": 1 }), + "blur": ("INT", { "default": 0, "min": 0, "max": 64, "step": 1 }), + }, + } + + RETURN_TYPES = ("MASK",) + FUNCTION = "execute" + CATEGORY = "essentials" + + def execute(self, image, clip_seg, prompt, threshold, smooth, dilate, blur): + processor, model = clip_seg + + imagenp = image.mul(255).clamp(0, 255).byte().cpu().numpy() + + outputs = [] + for i in imagenp: + inputs = processor(text=prompt, images=[i], return_tensors="pt") + out = model(**inputs) + out = out.logits.unsqueeze(1) + out = torch.sigmoid(out[0][0]) + out = (out > threshold) + outputs.append(out) + + del imagenp + + outputs = torch.stack(outputs, dim=0) + + if smooth > 0: + if smooth % 2 == 0: + smooth += 1 + outputs = T.functional.gaussian_blur(outputs, smooth) + + outputs = outputs.float() + + if dilate != 0: + outputs = expand_mask(outputs, dilate, True) + + if blur > 0: + if blur % 2 == 0: + blur += 1 + outputs = T.functional.gaussian_blur(outputs, blur) + + # resize to original size + outputs = F.interpolate(outputs.unsqueeze(1), size=(image.shape[1], image.shape[2]), mode='bicubic').squeeze(1) + + return (outputs,) + NODE_CLASS_MAPPINGS = { "GetImageSize+": GetImageSize, @@ -1959,6 +2034,9 @@ NODE_CLASS_MAPPINGS = { "ConditioningCombineMultiple+": ConditioningCombineMultiple, "ImageBatchMultiple+": ImageBatchMultiple, + "LoadCLIPSegModels+": LoadCLIPSegModels, + "ApplyCLIPSeg+": ApplyCLIPSeg, + #"NoiseFromImage~": NoiseFromImage, } @@ -2016,5 +2094,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ConditioningCombineMultiple+": "🔧 Conditionings Combine Multiple ", "ImageBatchMultiple+": "🔧 Images Batch Multiple", + "LoadCLIPSegModels+": "🔧 Load CLIPSeg Models", + "ApplyCLIPSeg+": "🔧 Apply CLIPSeg", + #"NoiseFromImage~": "🔧 Noise From Image", }