fixed code
This commit is contained in:
@@ -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
|
||||
|
||||
)
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user