126 lines
1.9 KiB
Python
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
|
|
|
|
) |