Add NSFWFilter node
This commit is contained in:
+3
-1
@@ -1,4 +1,6 @@
|
||||
.vscode
|
||||
.env
|
||||
.dev
|
||||
__pycache__
|
||||
__pycache__
|
||||
|
||||
safetychecker/*.safetensors
|
||||
@@ -2,7 +2,15 @@
|
||||
|
||||
Provides nodes and API geared towards using ComfyUI as a backend for external tools.
|
||||
|
||||
## Sending and receiving images
|
||||
* <a href="#images">Sending and receiving images</a>
|
||||
* <a href="#regions">Regions (Attention Masking)
|
||||
* <a href="#tiles">Tiled image processing
|
||||
* <a href="#misc">Miscellanious nodes
|
||||
* <a href="#api">Http API extensions (Model inspection)
|
||||
* <a href="#installation">⭳ Installation</a>
|
||||
|
||||
|
||||
## <a id="images" href="#toc">Sending and receiving images</a>
|
||||
|
||||
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': '<node ID>', 'output': {'images': [{'source': 'websocket', 'content-type': 'image/png', 'type': 'output'}, ...]}, 'prompt_id': '<prompt ID>}}
|
||||
```
|
||||
|
||||
## Regions
|
||||
## <a id="regions" href="#toc">Regions</a>
|
||||
|
||||
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
|
||||
## <a id="tiles" href="#toc">Tiles</a>
|
||||
|
||||
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
|
||||
## <a id="misc" href="#toc">Miscellaneous Nodes</a>
|
||||
|
||||
### 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.
|
||||
|
||||
|
||||
## <a id="api" href="#toc">API for model inspection</a>
|
||||
|
||||
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
|
||||
## <a id="installation" href="#toc">Installation</a>
|
||||
|
||||
Download the repository and unpack into the `custom_nodes` folder in the ComfyUI installation directory.
|
||||
|
||||
|
||||
+3
-1
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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),)
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user