added image ai-moderation

This commit is contained in:
SlimusMinus
2026-07-28 01:10:11 +03:00
parent 1132a774af
commit d376a7417a
34 changed files with 8302 additions and 34 deletions

View File

@@ -0,0 +1,126 @@
import torch
from PIL import Image
from transformers import (
CLIPProcessor,
CLIPModel
)
from app.ml.image.clip_prediction_result import ClipPredictionResult
from app.config.settings import settings
class ClipClassifier:
LABELS = [
"a normal photo",
"a photo containing marijuana",
"a photo containing drugs",
"a photo containing weapons",
"a pornographic photo",
"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.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:
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()
}
with torch.no_grad():
outputs = self.model(**inputs)
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
):
detected_labels.append(label)
max_index = torch.argmax(probs).item()
return ClipPredictionResult(
label=self.LABELS[max_index],
score=float(
probs[max_index]
),
detected_labels=detected_labels,
scores=scores
)

View File

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

View File

@@ -0,0 +1,83 @@
from PIL import Image
import torch
from app.ml.model_manager import model_manager
from app.ml.image.image_prediction_result import ImagePredictionResult
class ImageClassifier:
def predict(
self,
image: Image.Image
) -> ImagePredictionResult:
inputs = (
model_manager.image_processor(
image,
return_tensors="pt"
)
)
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)
print("Label:", label)
print("Score:", score)
print("Scores:", raw_scores)
return ImagePredictionResult(
label=label,
score=score,
approved=False,
raw_scores=raw_scores
)

View File

@@ -0,0 +1,14 @@
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

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

View File

@@ -3,6 +3,8 @@ from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification
import torch
@@ -13,6 +15,8 @@ class ModelManager:
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":
@@ -51,6 +55,35 @@ class ModelManager:
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...")
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()

View File

@@ -11,4 +11,6 @@ class PredictionResult:
raw_scores: list[float]
reason: str = field(default="")
reason: str = ""
detected_words: list[str] = field(default_factory=list)