From 2bcbd9fd8c878a8f9fabdfb19bfe19fbbe87598c Mon Sep 17 00:00:00 2001 From: Mackerel Date: Fri, 21 Apr 2023 16:17:16 -0400 Subject: [PATCH] make new classifer code less silly --- __init__.py | 32 ++++-------------- aesthetic/score_aes_B32_v0.py | 19 ----------- classifiers/__init__.py | 1 + .../aes-B32-v0.safetensors | Bin {aesthetic => classifiers}/aesthetic.py | 19 ++++++++++- classifiers/cafe_aesthetic.py | 9 +++++ classifiers/cafe_waifu.py | 9 +++++ .../laion-sac-logos-ava-v2.safetensors | Bin {laion => classifiers}/laion.py | 15 +++++++- laion/__pycache__/laion.cpython-310.pyc | Bin 2073 -> 0 bytes ...ore_laion-sac-logos-ava-v2.cpython-310.pyc | Bin 792 -> 0 bytes ...ore_laion_sac_logos_ava_v2.cpython-310.pyc | Bin 735 -> 0 bytes laion/score_laion_sac_logos_ava_v2.py | 15 -------- 13 files changed, 58 insertions(+), 61 deletions(-) delete mode 100644 aesthetic/score_aes_B32_v0.py create mode 100644 classifiers/__init__.py rename {aesthetic => classifiers}/aes-B32-v0.safetensors (100%) rename {aesthetic => classifiers}/aesthetic.py (60%) create mode 100644 classifiers/cafe_aesthetic.py create mode 100644 classifiers/cafe_waifu.py rename {laion => classifiers}/laion-sac-logos-ava-v2.safetensors (100%) rename {laion => classifiers}/laion.py (78%) delete mode 100644 laion/__pycache__/laion.cpython-310.pyc delete mode 100644 laion/__pycache__/score_laion-sac-logos-ava-v2.cpython-310.pyc delete mode 100644 laion/__pycache__/score_laion_sac_logos_ava_v2.cpython-310.pyc delete mode 100644 laion/score_laion_sac_logos_ava_v2.py 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 179a8bee523e7f5fc587ce84a9901848693d9e55..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 2073 zcmZ`(TZz|(B#Pl72Mgi0vyfDJloMTM$!zz5N>%)F+e zf~)SJtG!{YV%=2>hvHi?SQzwFPc7oq;%7`PspSKvmT_unxV-NTRusR_`X`QXJ6lgZ z`prjxKUsGhWY+d5we z#}h+$yPz~m8z=H>ZB%A(rd2}$PV#J5D^hFm7T77?<15^}j;W!wSpe)tUu9ohzO#1q zIueYGZDpxW$*SW)jz*x?(eE&+;K)8l9r7ukh>G0gRr_9kUgWF_CmlG*_r%wDS3bUd zkFz4W%Iw8QoHP2q0aZ~WCS6o*3|g2%BWl`1P{jHy9z7S2jh+{R>(<4*SWFi-m`eZ` zRHZ?Eb&67Mnnambh}qLciV~PrhbeWPYR!#x4=k2FARz+4f`fQvjJBO9jNn|ZE)o&|DZd- z>sBH9i2oQ-90c@+op<-nJ?FtT6w217G)aXLPLFkwKAMx89?1R#sMtPAbBdTZykSyM zUy+D4(653a^3`J0pm>o}(XP5mO!O501g}W&J$G1A)FxTm-cM0ADJmIb@ zWtu0}O8M_YGVefk5oF1zn1}Oksm(9tmxE}Qdl@-n!F~f$scG4hfQUsr;<1QC+*u7)c_hws{sB5c^QZs- 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 3bda7a3b3736eff55daafe3a3fb784da4b548c20..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 792 zcmZWny^ho{5VoCUll@<=qCqE7u|+l#H9`nl?u0~vlqT}VlU=Wpcx5}sftKoc03IPJ zc^q1|RJ;Nelk62MF_FhJlllF|YCfM3q+dU-#flN~yH_qLK;;b{qmca2 z<5KOEiTNFm?_O>UFTg@=ls1xhgg8HM1kY)o`;pK}d_3>)m>0+_IgtnYjaq71>z%N* z*6&Cg2$ayyPOKeu1Y2acx({538UZA3%)w-p5>UFpDxL4$q|mC$>ZUq=aY0ckZXL7A z`O-qQ_FOOw=4Kc388nz<@n_gJGJ*Sk zNYo8Fp(*B#Rq@T+Pl?I3JX(|FB{xPE60}JMtniQ0aK`X|-1>$GhbcB-M#v)iltxsq zP+UK8Oc~u<-Sh(C^!l;d=AznWsCMQ9PPzx&`A2T|F2qgNyaE?xSs_cvGUp$ZtkVTz zY4>1dp7oA;JjB2qFLnPg-4ekR8KPkUUqf!)crf%ih873lF4(dUieBM%ni5|3>o-@( RA{;BCD>`EfdWGMN{R4p6--`eM 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 c06758dfa22cc282f642399a2e7bc7e11ba21567..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 735 zcmYjPJ&)8d5Vf6TH~YO@MMFcy7JHEpB|@O9kl;WU9)T7w{oqzx$_EG=nV zvy!1@O`_5=Rau!)0%SkvhLHEP97cZxLz(of97874d)iH8*3)jLZVqfuy9JDU z+EYs8Kn}l?aw+L2vdVwsx#W7q{F7RVEpSjbAZ6WdEtj={%JD|jy7kHXcOSn#jsDt- zYTf9qwrg<|>*MoFD>mSuwc6O4B!u|1I!I9=;r|1pB*l4!(>_CR1C(59;>LXPZU~=b6(! ztsSUU;&hXuAH>YD@ShKb!jLGeEUvWL+ zSK&_)WhiW=4Lr!hQ4$^{;V215NkT7i!vt&F0YXk6)0~OFCt9dV$}R{R5UL%4Prn 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() -