From 6e85eef427f559d41f355cf2828bb8976b354a5d Mon Sep 17 00:00:00 2001 From: BetaDoggo Date: Thu, 18 Jul 2024 20:42:31 -0400 Subject: [PATCH] init --- __init__.py | 3 +++ nodes.py | 45 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 48 insertions(+) create mode 100644 __init__.py create mode 100644 nodes.py diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..d721463 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..174f639 --- /dev/null +++ b/nodes.py @@ -0,0 +1,45 @@ +from transformers import pipeline +from torchvision import transforms +import torch + +class YetAnotherSafetyChecker: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE",), + "threshold": ("FLOAT", { + "default": 0.8, + "min": 0.0, + "max": 1.0, + "step": 0.01 + }), + "cuda": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("IMAGE", "IMAGE", "STRING") + FUNCTION = "process_images" + + CATEGORY = "image/processing" + + def process_images(self, image, threshold, cuda): + if cuda: + device = "cuda" + else: + device = "cpu" + predict = pipeline("image-classification", model="AdamCodd/vit-base-nsfw-detector", device=device) #init pipeline + result = (predict(transforms.ToPILImage()(image[0].cpu().permute(2, 0, 1)))) #Convert to expected format + score = next(item['score'] for item in result if item['label'] == 'nsfw') + output = image + if(float(score) > threshold): + output = torch.zeros(1, 512, 512, dtype=torch.float32) #create black image tensor + return (output, image, str(score)) + +NODE_CLASS_MAPPINGS = { + "YetAnotherSafetyChecker": YetAnotherSafetyChecker +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "YetAnotherSafetyChecker": "Intercept NSFW Outputs" +}