added image ai-moderation
This commit is contained in:
126
ai-moderation/app/ml/image/clip_classifier.py
Normal file
126
ai-moderation/app/ml/image/clip_classifier.py
Normal 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
|
||||
|
||||
)
|
||||
Reference in New Issue
Block a user