make new classifer code less silly
This commit is contained in:
+7
-25
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
__all__ = ["laion", "aesthetic", "cafe_waifu", "cafe_aesthetic"]
|
||||
@@ -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()
|
||||
@@ -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']
|
||||
@@ -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']
|
||||
@@ -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()
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user