diff --git a/.gitignore b/.gitignore index c9f3ec0..02b06aa 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,6 @@ .vscode .env .dev -__pycache__ \ No newline at end of file +__pycache__ + +safetychecker/*.safetensors \ No newline at end of file diff --git a/README.md b/README.md index ac72fc8..725c3f2 100644 --- a/README.md +++ b/README.md @@ -2,7 +2,15 @@ Provides nodes and API geared towards using ComfyUI as a backend for external tools. -## Sending and receiving images +* Sending and receiving images +* Regions (Attention Masking) +* Tiled image processing +* Miscellanious nodes +* Http API extensions (Model inspection) +* ⭳ Installation + + +## Sending and receiving images ComfyUI exchanges images via the filesystem. This requires a multi-step process (upload images, prompt, download images), is rather @@ -36,7 +44,7 @@ That is two 32-bit integers (big endian) with values 1 and 2 followed by the PNG {'type': 'executed', 'data': {'node': '', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '}} ``` -## Regions +## Regions These nodes implement attention masking for arbitrary number of image regions. Text prompts only apply to the masked area. In contrast to condition masking, this method is less "forceful", but leads to more natural image compositions. @@ -71,7 +79,7 @@ Copies a mask into the alpha channel of an image. * Outputs: RGBA image with mask used as transparency -## Tiles +## Tiles Splitting an image into tiles to be processed individually is a useful method to speed up diffusion and save VRAM. There are various nodes out there which provide a fixed pipeline. @@ -114,7 +122,19 @@ depending on the chosen blend size. This mask is used internally by "Merge Image Tile", but it can also be useful as input for "Set Latent Noise Mask" in upscale workflows. -## API for model inspection +## Miscellaneous Nodes + +### NSFW Filter + +Checks images for NSFW content using [Safety-Checker](https://huggingface.co/CompVis/stable-diffusion-safety-checker). Images which don't pass the check are blurred to +obfuscate contents. Model is downloaded on first use. + +Inputs: image and sensitivity (0.5 for explicit content only, 0.7+ to include partial nudity). + +**Important:** the filter isn't perfect. Some explicit content may slip through. + + +## API for model inspection There are various types of models that can be loaded as checkpoint, LoRA, ControlNet, etc. which cannot be used interchangeably. The following API helps to categorize and filter them. @@ -137,7 +157,7 @@ Lists available models with additional classification info. _Note: currently only supports checkpoints. May add other models in the future._ -## Installation +## Installation Download the repository and unpack into the `custom_nodes` folder in the ComfyUI installation directory. diff --git a/__init__.py b/__init__.py index 147ce7b..f4ecd44 100644 --- a/__init__.py +++ b/__init__.py @@ -1,4 +1,4 @@ -from . import api, nodes, tile, region +from . import api, nodes, tile, region, nsfw NODE_CLASS_MAPPINGS = { "ETN_LoadImageBase64": nodes.LoadImageBase64, @@ -15,6 +15,7 @@ NODE_CLASS_MAPPINGS = { "ETN_DefineRegion": region.DefineRegion, "ETN_ListRegionMasks": region.ListRegionMasks, "ETN_AttentionMask": region.AttentionMask, + "ETN_NSFWFilter": nsfw.NSFWFilter, } NODE_DISPLAY_NAME_MAPPINGS = { "ETN_LoadImageBase64": "Load Image (Base64)", @@ -33,4 +34,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ETN_DefineRegion": "Define Region", "ETN_ListRegionMasks": "List Region Masks", "ETN_AttentionMask": "Regions Attention Mask", + "ETN_NSFWFilter": "NSFW Filter", } diff --git a/nsfw.py b/nsfw.py new file mode 100644 index 0000000..4e9824f --- /dev/null +++ b/nsfw.py @@ -0,0 +1,145 @@ +from weakref import ref as WeakRef +from pathlib import Path +from tqdm import tqdm +import torch +import torch.nn as nn +import torch.nn.functional as F +from torch import Tensor +from transformers import CLIPImageProcessor, CLIPConfig, CLIPVisionModel, PreTrainedModel +from kornia.filters import box_blur + +from .nodes import to_bchw, to_bhwc + + +def cosine_similarity(image_embeds: Tensor, text_embeds: Tensor): + if image_embeds.dim() == 2 and text_embeds.dim() == 2: + image_embeds = image_embeds.unsqueeze(1) + return F.cosine_similarity(image_embeds, text_embeds, dim=-1) + + +class CLIPSafetyChecker(PreTrainedModel): + # https://huggingface.co/CompVis/stable-diffusion-safety-checker + # Adapted from: + # https://github.com/huggingface/diffusers/blob/main/src/diffusers/pipelines/stable_diffusion/safety_checker.py + + config_class = CLIPConfig + _no_split_modules = ["CLIPEncoderLayer"] + + def __init__(self, config: CLIPConfig): + super().__init__(config) + projdim = config.projection_dim + + self.vision_model = CLIPVisionModel(config.vision_config) + self.visual_projection = nn.Linear(config.vision_config.hidden_size, projdim, bias=False) + + self.concept_embeds = nn.Parameter(torch.ones(17, projdim), requires_grad=False) + self.special_care_embeds = nn.Parameter(torch.ones(3, projdim), requires_grad=False) + self.concept_embeds_weights = nn.Parameter(torch.ones(17), requires_grad=False) + self.special_care_embeds_weights = nn.Parameter(torch.ones(3), requires_grad=False) + + def forward(self, clip_input, images: Tensor, sensitivity: float): + with torch.no_grad(): + image_batch = self.vision_model(clip_input)[1] + image_embeds = self.visual_projection(image_batch) + sensitivity = -0.1 + 0.14 * sensitivity + + special_cos_dist = cosine_similarity(image_embeds, self.special_care_embeds) + special_scores_threshold = self.special_care_embeds_weights.unsqueeze(0) + special_scores = special_cos_dist - special_scores_threshold + sensitivity + + if torch.any(special_scores > 0): + sensitivity = sensitivity + 0.01 + + cos_dist = cosine_similarity(image_embeds, self.concept_embeds) + concept_threshold = self.concept_embeds_weights.unsqueeze(0) + concept_scores = cos_dist - concept_threshold + sensitivity + + is_nsfw = [torch.any(concept_scores[i] > 0) for i in range(concept_scores.shape[0])] + is_nsfw = [x.item() for x in is_nsfw] + return self.filter_images(images, is_nsfw) + + def filter_images(self, images: Tensor, is_nsfw: list[bool]): + if not any(is_nsfw): + return images + + images = images.clone() + images_to_filter = (i for i, nsfw in enumerate(is_nsfw) if nsfw) + orig_size = images.shape[-2:] + for idx in images_to_filter: + filtered = images[idx].unsqueeze(0) + filtered = F.interpolate(filtered, size=64, mode="nearest") + filtered = box_blur(filtered, 7, separable=True) + filtered = F.interpolate(filtered, size=orig_size, mode="bilinear") + images[idx] = filtered.squeeze(0) + return images + + +class CachedModels: + _instance: WeakRef | None = None + + def __init__(self): + model_dir = Path(__file__).parent / "safetychecker" + model_file = model_dir / "model.safetensors" + if not model_file.exists(): + self.download( + "https://huggingface.co/CompVis/stable-diffusion-safety-checker/resolve/refs%2Fpr%2F41/model.safetensors", + target=model_file, + ) + self.feature_extractor = CLIPImageProcessor.from_pretrained(model_dir) + self.safety_checker = CLIPSafetyChecker.from_pretrained(model_dir) + + @classmethod + def load(cls): + models = cls._instance and cls._instance() + if models is None: + models = cls() + cls._instance = WeakRef(models) + return models + + def download(self, url: str, target: Path): + import requests + + try: + target_temp = target.with_suffix(".download") + with requests.get(url, stream=True) as response: + text = "NSFWFilter model download" + total = int(response.headers.get("content-length", 0)) + pbar = tqdm(None, total=total, unit="b", unit_scale=True, desc=text) + with open(target_temp, "wb") as f: + for chunk in response.iter_content(chunk_size=8192): + f.write(chunk) + pbar.update(len(chunk)) + pbar.close() + target_temp.rename(target) + except Exception as e: + raise RuntimeError( + f"NSFWFilter: Failed to download safety-checker model from {url} to target location {target}: {e}" + ) from e + + +class NSFWFilter: + models: CachedModels + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "sensitivity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.10}), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "check" + CATEGORY = "external_tooling" + + def __init__(self): + self.models = CachedModels.load() + + def check(self, image, sensitivity): + image = to_bchw(image) + input = self.models.feature_extractor(image, do_rescale=False, return_tensors="pt") + filtered = self.models.safety_checker( + images=image, clip_input=input.pixel_values, sensitivity=sensitivity + ) + return (to_bhwc(filtered),) diff --git a/safetychecker/config.json b/safetychecker/config.json new file mode 100644 index 0000000..4493f5a --- /dev/null +++ b/safetychecker/config.json @@ -0,0 +1,171 @@ +{ + "_name_or_path": "clip-vit-large-patch14/", + "architectures": [ + "SafetyChecker" + ], + "initializer_factor": 1.0, + "logit_scale_init_value": 2.6592, + "model_type": "clip", + "projection_dim": 768, + "text_config": { + "_name_or_path": "", + "add_cross_attention": false, + "architectures": null, + "attention_dropout": 0.0, + "bad_words_ids": null, + "bos_token_id": 0, + "chunk_size_feed_forward": 0, + "cross_attention_hidden_size": null, + "decoder_start_token_id": null, + "diversity_penalty": 0.0, + "do_sample": false, + "dropout": 0.0, + "early_stopping": false, + "encoder_no_repeat_ngram_size": 0, + "eos_token_id": 2, + "exponential_decay_length_penalty": null, + "finetuning_task": null, + "forced_bos_token_id": null, + "forced_eos_token_id": null, + "hidden_act": "quick_gelu", + "hidden_size": 768, + "id2label": { + "0": "LABEL_0", + "1": "LABEL_1" + }, + "initializer_factor": 1.0, + "initializer_range": 0.02, + "intermediate_size": 3072, + "is_decoder": false, + "is_encoder_decoder": false, + "label2id": { + "LABEL_0": 0, + "LABEL_1": 1 + }, + "layer_norm_eps": 1e-05, + "length_penalty": 1.0, + "max_length": 20, + "max_position_embeddings": 77, + "min_length": 0, + "model_type": "clip_text_model", + "no_repeat_ngram_size": 0, + "num_attention_heads": 12, + "num_beam_groups": 1, + "num_beams": 1, + "num_hidden_layers": 12, + "num_return_sequences": 1, + "output_attentions": false, + "output_hidden_states": false, + "output_scores": false, + "pad_token_id": 1, + "prefix": null, + "problem_type": null, + "pruned_heads": {}, + "remove_invalid_values": false, + "repetition_penalty": 1.0, + "return_dict": true, + "return_dict_in_generate": false, + "sep_token_id": null, + "task_specific_params": null, + "temperature": 1.0, + "tie_encoder_decoder": false, + "tie_word_embeddings": true, + "tokenizer_class": null, + "top_k": 50, + "top_p": 1.0, + "torch_dtype": null, + "torchscript": false, + "transformers_version": "4.21.0.dev0", + "typical_p": 1.0, + "use_bfloat16": false, + "vocab_size": 49408 + }, + "text_config_dict": { + "hidden_size": 768, + "intermediate_size": 3072, + "num_attention_heads": 12, + "num_hidden_layers": 12 + }, + "torch_dtype": "float32", + "transformers_version": null, + "vision_config": { + "_name_or_path": "", + "add_cross_attention": false, + "architectures": null, + "attention_dropout": 0.0, + "bad_words_ids": null, + "bos_token_id": null, + "chunk_size_feed_forward": 0, + "cross_attention_hidden_size": null, + "decoder_start_token_id": null, + "diversity_penalty": 0.0, + "do_sample": false, + "dropout": 0.0, + "early_stopping": false, + "encoder_no_repeat_ngram_size": 0, + "eos_token_id": null, + "exponential_decay_length_penalty": null, + "finetuning_task": null, + "forced_bos_token_id": null, + "forced_eos_token_id": null, + "hidden_act": "quick_gelu", + "hidden_size": 1024, + "id2label": { + "0": "LABEL_0", + "1": "LABEL_1" + }, + "image_size": 224, + "initializer_factor": 1.0, + "initializer_range": 0.02, + "intermediate_size": 4096, + "is_decoder": false, + "is_encoder_decoder": false, + "label2id": { + "LABEL_0": 0, + "LABEL_1": 1 + }, + "layer_norm_eps": 1e-05, + "length_penalty": 1.0, + "max_length": 20, + "min_length": 0, + "model_type": "clip_vision_model", + "no_repeat_ngram_size": 0, + "num_attention_heads": 16, + "num_beam_groups": 1, + "num_beams": 1, + "num_hidden_layers": 24, + "num_return_sequences": 1, + "output_attentions": false, + "output_hidden_states": false, + "output_scores": false, + "pad_token_id": null, + "patch_size": 14, + "prefix": null, + "problem_type": null, + "pruned_heads": {}, + "remove_invalid_values": false, + "repetition_penalty": 1.0, + "return_dict": true, + "return_dict_in_generate": false, + "sep_token_id": null, + "task_specific_params": null, + "temperature": 1.0, + "tie_encoder_decoder": false, + "tie_word_embeddings": true, + "tokenizer_class": null, + "top_k": 50, + "top_p": 1.0, + "torch_dtype": null, + "torchscript": false, + "transformers_version": "4.21.0.dev0", + "typical_p": 1.0, + "use_bfloat16": false + }, + "vision_config_dict": { + "hidden_size": 1024, + "intermediate_size": 4096, + "num_attention_heads": 16, + "num_hidden_layers": 24, + "patch_size": 14 + } +} diff --git a/safetychecker/preprocessor_config.json b/safetychecker/preprocessor_config.json new file mode 100644 index 0000000..5294955 --- /dev/null +++ b/safetychecker/preprocessor_config.json @@ -0,0 +1,20 @@ +{ + "crop_size": 224, + "do_center_crop": true, + "do_convert_rgb": true, + "do_normalize": true, + "do_resize": true, + "feature_extractor_type": "CLIPFeatureExtractor", + "image_mean": [ + 0.48145466, + 0.4578275, + 0.40821073 + ], + "image_std": [ + 0.26862954, + 0.26130258, + 0.27577711 + ], + "resample": 3, + "size": 224 +}