fixed code

This commit is contained in:
SlimusMinus
2026-08-05 01:44:29 +03:00
parent 72159dbe4a
commit d33f0eaa78
30 changed files with 143 additions and 375 deletions

View File

@@ -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
)

View File

@@ -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]

View File

@@ -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
)
)

View File

@@ -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)

View File

@@ -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
)
)