diff --git a/__init__.py b/__init__.py index d7a0762..e5be676 100644 --- a/__init__.py +++ b/__init__.py @@ -1,5 +1,7 @@ +import importlib import math import pathlib +import sys import time import warnings @@ -7,31 +9,17 @@ import numpy as np from PIL import Image import torch import tqdm -import transformers import folder_paths import model_management import nodes -import sys sys.path.append(str(pathlib.Path(__file__).parent)) +import classifiers -CAFE_MODELS = {"cafe_aesthetic": 2, "cafe_waifu": 5} BLOCK_ORDER = [12, 11, 13, 10, 14, 9, 15, 8, 16, 7, 17, 6, 18, 5, 19, 4, 20, 3, 21, 2, 22, 1, 23, 0, 24] - -def run_cafe_classifier(image, classifier): - import transformers - pipe = transformers.pipeline( - "image-classification", - model=f"cafeai/{classifier}") - result = pipe(image, top_k=CAFE_MODELS[classifier]) - for data in result: - if data['label'] == classifier.split("_")[1]: - return data['score'] - - class AutoMBW: def __init__(self): self.type = "output" @@ -54,7 +42,7 @@ class AutoMBW: }), "search_depth": ("INT", {"default": 4, "min": 2}), "sample_count": ("INT", {"default": 1, "min": 1}), - "classifier": (["aesthetic", "laion", "cafe_aesthetic", "cafe_waifu"],), + "classifier": (classifiers.__all__,), }} RETURN_TYPES = () @@ -92,14 +80,7 @@ class AutoMBW: with warnings.catch_warnings(): # several possible transformers nags warnings.filterwarnings('ignore') - if self.classifier.startswith("cafe_"): - rating += run_cafe_classifier(image, self.classifier) - elif self.classifier == "laion": - import laion.score_laion_sac_logos_ava_v2 - rating += laion.score_laion_sac_logos_ava_v2.score(image) - elif self.classifier == "aesthetic": - import aesthetic.score_aes_B32_v0 - rating += aesthetic.score_aes_B32_v0.score(image) + rating += self.classifier(image) return rating def search(self, block, current, start, depth, maximum): @@ -130,7 +111,8 @@ class AutoMBW: self.negative = [[clip.encode(negative), {}]] self.search_depth = search_depth self.sample_count = sample_count - self.classifier = classifier + self.classifier = importlib.import_module( + "." + classifier, "classifiers").score # model setup if model_management.vram_state == model_management.VRAMState.HIGH_VRAM: diff --git a/aesthetic/score_aes_B32_v0.py b/aesthetic/score_aes_B32_v0.py deleted file mode 100644 index 1eb863f..0000000 --- a/aesthetic/score_aes_B32_v0.py +++ /dev/null @@ -1,19 +0,0 @@ -import os -import torch -import safetensors -from transformers import CLIPModel, CLIPProcessor -from aesthetic.aesthetic import image_embeddings_direct, Classifier - -dirname = os.path.dirname(__file__) -aesthetic_path = os.path.join(dirname, "aes-B32-v0.safetensors") -clip_name = 'openai/clip-vit-base-patch32' -clipprocessor = CLIPProcessor.from_pretrained(clip_name) -clipmodel = CLIPModel.from_pretrained(clip_name).to('cuda').eval() -aes_model = Classifier(512, 256, 1).to('cuda') -aes_model.load_state_dict(safetensors.torch.load_file(aesthetic_path)) - -def score(image): - image_embeds = image_embeddings_direct(image, clipmodel, clipprocessor) - prediction = aes_model(torch.from_numpy(image_embeds).float().to('cuda')) - return prediction.item() - diff --git a/classifiers/__init__.py b/classifiers/__init__.py new file mode 100644 index 0000000..5579bd0 --- /dev/null +++ b/classifiers/__init__.py @@ -0,0 +1 @@ +__all__ = ["laion", "aesthetic", "cafe_waifu", "cafe_aesthetic"] diff --git a/aesthetic/aes-B32-v0.safetensors b/classifiers/aes-B32-v0.safetensors similarity index 100% rename from aesthetic/aes-B32-v0.safetensors rename to classifiers/aes-B32-v0.safetensors diff --git a/aesthetic/aesthetic.py b/classifiers/aesthetic.py similarity index 60% rename from aesthetic/aesthetic.py rename to classifiers/aesthetic.py index b4d5eed..3863c59 100644 --- a/aesthetic/aesthetic.py +++ b/classifiers/aesthetic.py @@ -1,5 +1,9 @@ -import torch +import pathlib + import numpy as np +import torch +import safetensors +from transformers import CLIPModel, CLIPProcessor use_cuda = torch.cuda.is_available() @@ -28,3 +32,16 @@ class Classifier(torch.nn.Module): x = self.fc3(x) x = self.sigmoid(x) return x + +dirname = pathlib.Path(__file__).parent +aesthetic_path = dirname.joinpath("aes-B32-v0.safetensors") +clip_name = 'openai/clip-vit-base-patch32' +clipprocessor = CLIPProcessor.from_pretrained(clip_name) +clipmodel = CLIPModel.from_pretrained(clip_name).to('cuda').eval() +aes_model = Classifier(512, 256, 1).to('cuda') +aes_model.load_state_dict(safetensors.torch.load_file(aesthetic_path)) + +def score(image): + image_embeds = image_embeddings_direct(image, clipmodel, clipprocessor) + prediction = aes_model(torch.from_numpy(image_embeds).float().to('cuda')) + return prediction.item() diff --git a/classifiers/cafe_aesthetic.py b/classifiers/cafe_aesthetic.py new file mode 100644 index 0000000..1b739c8 --- /dev/null +++ b/classifiers/cafe_aesthetic.py @@ -0,0 +1,9 @@ +import transformers + +def score(image): + pipe = transformers.pipeline("image-classification", + model="cafeai/cafe_aesthetic") + result = pipe(image, top_k=2) + for data in result: + if data['label'] == "aesthetic": + return data['score'] diff --git a/classifiers/cafe_waifu.py b/classifiers/cafe_waifu.py new file mode 100644 index 0000000..ad6cf0a --- /dev/null +++ b/classifiers/cafe_waifu.py @@ -0,0 +1,9 @@ +import transformers + +def score(image): + pipe = transformers.pipeline("image-classification", + model="cafeai/cafe_waifu") + result = pipe(image, top_k=5) + for data in result: + if data['label'] == "waifu": + return data['score'] diff --git a/laion/laion-sac-logos-ava-v2.safetensors b/classifiers/laion-sac-logos-ava-v2.safetensors similarity index 100% rename from laion/laion-sac-logos-ava-v2.safetensors rename to classifiers/laion-sac-logos-ava-v2.safetensors diff --git a/laion/laion.py b/classifiers/laion.py similarity index 78% rename from laion/laion.py rename to classifiers/laion.py index f77153d..d42a90c 100644 --- a/laion/laion.py +++ b/classifiers/laion.py @@ -1,6 +1,9 @@ -import torch +import pathlib + import numpy as np import clip +import torch +import safetensors use_cuda = torch.cuda.is_available() @@ -45,3 +48,13 @@ class MLP(torch.nn.Module): def forward(self, x): return self.layers(x) + +dirname = pathlib.Path(__file__).parent +aesthetic_path = dirname.joinpath("laion-sac-logos-ava-v2.safetensors") +aes_model = MLP(768).to('cuda').eval() +aes_model.load_state_dict(safetensors.torch.load_file(aesthetic_path)) + +def score(image): + image_embeds = image_embeddings_direct_laion(image) + prediction = aes_model(torch.from_numpy(image_embeds).float().to('cuda')) + return prediction.item() diff --git a/laion/__pycache__/laion.cpython-310.pyc b/laion/__pycache__/laion.cpython-310.pyc deleted file mode 100644 index 179a8be..0000000 Binary files a/laion/__pycache__/laion.cpython-310.pyc and /dev/null differ diff --git a/laion/__pycache__/score_laion-sac-logos-ava-v2.cpython-310.pyc b/laion/__pycache__/score_laion-sac-logos-ava-v2.cpython-310.pyc deleted file mode 100644 index 3bda7a3..0000000 Binary files a/laion/__pycache__/score_laion-sac-logos-ava-v2.cpython-310.pyc and /dev/null differ diff --git a/laion/__pycache__/score_laion_sac_logos_ava_v2.cpython-310.pyc b/laion/__pycache__/score_laion_sac_logos_ava_v2.cpython-310.pyc deleted file mode 100644 index c06758d..0000000 Binary files a/laion/__pycache__/score_laion_sac_logos_ava_v2.cpython-310.pyc and /dev/null differ diff --git a/laion/score_laion_sac_logos_ava_v2.py b/laion/score_laion_sac_logos_ava_v2.py deleted file mode 100644 index eadae5d..0000000 --- a/laion/score_laion_sac_logos_ava_v2.py +++ /dev/null @@ -1,15 +0,0 @@ -import os -import torch -import safetensors -from laion.laion import image_embeddings_direct_laion, MLP - -dirname = os.path.dirname(__file__) -aesthetic_path = os.path.join(dirname, "laion-sac-logos-ava-v2.safetensors") -aes_model = MLP(768).to('cuda').eval() -aes_model.load_state_dict(safetensors.torch.load_file(aesthetic_path)) - -def score(image): - image_embeds = image_embeddings_direct_laion(image) - prediction = aes_model(torch.from_numpy(image_embeds).float().to('cuda')) - return prediction.item() -