Files
post-moderation/ai-moderation/app/ml/image/clip_classifier.py
2026-07-28 01:10:11 +03:00

126 lines
1.9 KiB
Python

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
)