make new classifer code less silly
This commit is contained in:
+7
-25
@@ -1,5 +1,7 @@
|
|||||||
|
import importlib
|
||||||
import math
|
import math
|
||||||
import pathlib
|
import pathlib
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
@@ -7,31 +9,17 @@ import numpy as np
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
import torch
|
import torch
|
||||||
import tqdm
|
import tqdm
|
||||||
import transformers
|
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import model_management
|
import model_management
|
||||||
import nodes
|
import nodes
|
||||||
|
|
||||||
import sys
|
|
||||||
sys.path.append(str(pathlib.Path(__file__).parent))
|
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,
|
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]
|
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:
|
class AutoMBW:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.type = "output"
|
self.type = "output"
|
||||||
@@ -54,7 +42,7 @@ class AutoMBW:
|
|||||||
}),
|
}),
|
||||||
"search_depth": ("INT", {"default": 4, "min": 2}),
|
"search_depth": ("INT", {"default": 4, "min": 2}),
|
||||||
"sample_count": ("INT", {"default": 1, "min": 1}),
|
"sample_count": ("INT", {"default": 1, "min": 1}),
|
||||||
"classifier": (["aesthetic", "laion", "cafe_aesthetic", "cafe_waifu"],),
|
"classifier": (classifiers.__all__,),
|
||||||
}}
|
}}
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
RETURN_TYPES = ()
|
||||||
@@ -92,14 +80,7 @@ class AutoMBW:
|
|||||||
with warnings.catch_warnings():
|
with warnings.catch_warnings():
|
||||||
# several possible transformers nags
|
# several possible transformers nags
|
||||||
warnings.filterwarnings('ignore')
|
warnings.filterwarnings('ignore')
|
||||||
if self.classifier.startswith("cafe_"):
|
rating += self.classifier(image)
|
||||||
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)
|
|
||||||
return rating
|
return rating
|
||||||
|
|
||||||
def search(self, block, current, start, depth, maximum):
|
def search(self, block, current, start, depth, maximum):
|
||||||
@@ -130,7 +111,8 @@ class AutoMBW:
|
|||||||
self.negative = [[clip.encode(negative), {}]]
|
self.negative = [[clip.encode(negative), {}]]
|
||||||
self.search_depth = search_depth
|
self.search_depth = search_depth
|
||||||
self.sample_count = sample_count
|
self.sample_count = sample_count
|
||||||
self.classifier = classifier
|
self.classifier = importlib.import_module(
|
||||||
|
"." + classifier, "classifiers").score
|
||||||
|
|
||||||
# model setup
|
# model setup
|
||||||
if model_management.vram_state == model_management.VRAMState.HIGH_VRAM:
|
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 numpy as np
|
||||||
|
import torch
|
||||||
|
import safetensors
|
||||||
|
from transformers import CLIPModel, CLIPProcessor
|
||||||
|
|
||||||
use_cuda = torch.cuda.is_available()
|
use_cuda = torch.cuda.is_available()
|
||||||
|
|
||||||
@@ -28,3 +32,16 @@ class Classifier(torch.nn.Module):
|
|||||||
x = self.fc3(x)
|
x = self.fc3(x)
|
||||||
x = self.sigmoid(x)
|
x = self.sigmoid(x)
|
||||||
return 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 numpy as np
|
||||||
import clip
|
import clip
|
||||||
|
import torch
|
||||||
|
import safetensors
|
||||||
|
|
||||||
use_cuda = torch.cuda.is_available()
|
use_cuda = torch.cuda.is_available()
|
||||||
|
|
||||||
@@ -45,3 +48,13 @@ class MLP(torch.nn.Module):
|
|||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
return self.layers(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