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