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