fixed code

This commit is contained in:
SlimusMinus
2026-08-05 01:44:29 +03:00
parent 72159dbe4a
commit d33f0eaa78
30 changed files with 143 additions and 375 deletions

View File

@@ -1,17 +1,12 @@
import torch
from PIL import Image
from transformers import (
CLIPProcessor,
CLIPModel
)
from transformers import CLIPProcessor, CLIPModel
from app.ml.image.clip_prediction_result import ClipPredictionResult
from app.config.settings import settings
from app.ml.image.clip_prediction_result import ClipPredictionResult
class ClipClassifier:
LABELS = [
"a normal photo",
"a photo containing marijuana",
@@ -21,106 +16,46 @@ class ClipClassifier:
"a photo containing violence"
]
def __init__(self):
self.device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Loading CLIP on {self.device}")
self.device = (
"cuda"
if torch.cuda.is_available()
else "cpu"
)
print(
f"Loading CLIP on {self.device}"
)
self.model = CLIPModel.from_pretrained(
settings.IMAGE_CLIP_MODEL
)
self.processor = CLIPProcessor.from_pretrained(
settings.IMAGE_CLIP_MODEL
)
self.model = CLIPModel.from_pretrained(settings.IMAGE_CLIP_MODEL)
self.processor = CLIPProcessor.from_pretrained(settings.IMAGE_CLIP_MODEL)
self.model.to(self.device)
self.model.eval()
def predict(
self,
image: Image.Image
) -> ClipPredictionResult:
def predict(self, image: Image.Image) -> ClipPredictionResult:
inputs = self.processor(
text=self.LABELS,
images=image,
return_tensors="pt",
padding=True
)
inputs = {
k: v.to(self.device)
for k, v in inputs.items()
}
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
logits = (
outputs
.logits_per_image
)
probs = (
logits.softmax(dim=1)[0]
)
logits = outputs.logits_per_image
probs = logits.softmax(dim=1)[0]
scores = {}
detected_labels = []
THRESHOLD = 0.75
for index, probability in enumerate(probs):
label = self.LABELS[index]
score = float(probability)
scores[label] = score
if (
label != "a normal photo"
and score >= THRESHOLD
):
if label != "a normal photo" and score >= settings.CLIP_THRESHOLD:
detected_labels.append(label)
max_index = torch.argmax(probs).item()
return ClipPredictionResult(
label=self.LABELS[max_index],
score=float(
probs[max_index]
),
score=float(probs[max_index]),
detected_labels=detected_labels,
scores=scores
)

View File

@@ -3,11 +3,10 @@ from dataclasses import dataclass
@dataclass(slots=True)
class ClipPredictionResult:
label: str
score: float
detected_labels: list[str]
scores: dict[str, float]
scores: dict[str, float]

View File

@@ -1,19 +1,16 @@
from PIL import Image
import torch
from PIL import Image
from app.ml.model_manager import model_manager
from app.ml.image.image_prediction_result import ImagePredictionResult
from app.ml.model_manager import model_manager
class ImageClassifier:
def predict(
self,
image: Image.Image
) -> ImagePredictionResult:
inputs = (
model_manager.image_processor(
image,
@@ -21,48 +18,38 @@ class ImageClassifier:
)
)
inputs = {
key: value.to(model_manager.device)
for key, value in inputs.items()
}
with torch.no_grad():
outputs = (
model_manager.image_model(**inputs)
)
probabilities = torch.softmax(
outputs.logits,
dim=1
)
raw_scores = probabilities[0].tolist()
predicted_index = (
torch.argmax(probabilities, dim=1)
.item()
)
labels = (
model_manager.image_model
.config
.id2label
)
label = labels[predicted_index]
score = raw_scores[predicted_index]
print("====================")
print("IMAGE RESULT")
print("Labels:", labels)
@@ -70,7 +57,6 @@ class ImageClassifier:
print("Score:", score)
print("Scores:", raw_scores)
return ImagePredictionResult(
label=label,
@@ -80,4 +66,4 @@ class ImageClassifier:
approved=False,
raw_scores=raw_scores
)
)

View File

@@ -1,14 +0,0 @@
from dataclasses import dataclass
from dataclasses import field
@dataclass(slots=True)
class ImagePredictionResult:
approved: bool
score: float
reason: str
detected_labels: list[str] = field(default_factory=list)

View File

@@ -3,7 +3,6 @@ from dataclasses import dataclass, field
@dataclass(slots=True)
class ImagePredictionResult:
label: str
score: float
@@ -16,4 +15,4 @@ class ImagePredictionResult:
detected_labels: list[str] = field(
default_factory=list
)
)

View File

@@ -1,89 +1,69 @@
from transformers import AutoTokenizer
from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
import torch
from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification
from transformers import AutoModelForSequenceClassification
from transformers import AutoTokenizer
import torch
from app.config.logging import logger
from app.config.settings import settings
class ModelManager:
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
self.image_model = None
self.image_processor = None
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
self.image_model = None
self.image_processor = None
def load_device(self):
if settings.DEVICE == "auto":
self.device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
else:
self.device = torch.device(settings.DEVICE)
def load_device(self):
if settings.DEVICE == "auto":
self.device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
else:
self.device = torch.device(settings.DEVICE)
def load_tokenizer(self):
print("Loading tokenizer...")
def load_tokenizer(self):
logger.info("Loading tokenizer...")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
logger.info("Tokenizer loaded.")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
def load_model(self):
logger.info("Loading model...")
self.text_model = AutoModelForSequenceClassification.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.text_model.to(self.device)
self.text_model.eval()
logger.debug("id2label: %s", self.text_model.config.id2label)
logger.info("Text model loaded.")
print("Tokenizer loaded.")
def load_image_model(self):
logger.info("Loading image model...")
self.image_processor = AutoImageProcessor.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.image_model = AutoModelForImageClassification.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.image_model.to(self.device)
self.image_model.eval()
logger.info("Image model loaded.")
def load_model(self):
print("Loading model...")
self.text_model = AutoModelForSequenceClassification.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.text_model.to(self.device)
self.text_model.eval()
print(self.text_model.config.id2label)
print("Text model loaded.")
def initialize(self):
self.load_device()
print(f"Using device: {self.device}")
self.load_tokenizer()
self.load_model()
self.load_image_model()
def load_image_model(self):
print("Loading image model...")
def initialize(self):
self.load_device()
logger.info(f"Using device: {self.device}")
self.load_tokenizer()
self.load_model()
self.load_image_model()
self.image_processor = (
AutoImageProcessor.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
)
self.image_model = (
AutoModelForImageClassification.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
)
self.image_model.to(self.device)
self.image_model.eval()
print("Image model loaded.")
model_manager = ModelManager()
model_manager = ModelManager()

View File

@@ -1,8 +1,8 @@
from dataclasses import dataclass, field
@dataclass(slots=True)
class PredictionResult:
label: str
score: float
@@ -13,4 +13,4 @@ class PredictionResult:
reason: str = ""
detected_words: list[str] = field(default_factory=list)
detected_words: list[str] = field(default_factory=list)

View File

@@ -1,9 +1,9 @@
import torch
from app.ml.model_manager import model_manager
from app.config.logging import logger
from app.ml.prediction_result import PredictionResult
class TextClassifier:
def predict(self, text: str):
inputs = model_manager.text_tokenizer(