start ai-moderation
This commit is contained in:
@@ -0,0 +1,47 @@
|
||||
import torch
|
||||
|
||||
from app.ml.model_manager import model_manager
|
||||
from app.config.logging import logger
|
||||
from app.ml.prediction_result import PredictionResult
|
||||
|
||||
class TextClassifier:
|
||||
def predict(self, text: str):
|
||||
inputs = model_manager.text_tokenizer(
|
||||
text,
|
||||
return_tensors="pt",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
padding=True
|
||||
)
|
||||
|
||||
inputs = {
|
||||
key: value.to(model_manager.device)
|
||||
for key, value in inputs.items()
|
||||
}
|
||||
|
||||
print(inputs)
|
||||
|
||||
with torch.no_grad():
|
||||
outputs = model_manager.text_model(**inputs)
|
||||
|
||||
print(outputs)
|
||||
|
||||
probabilities = torch.softmax(outputs.logits, dim=1)
|
||||
print("Labels:", model_manager.text_model.config.id2label)
|
||||
print("Logits:", outputs.logits)
|
||||
print("Probabilities:", probabilities)
|
||||
|
||||
print(probabilities)
|
||||
raw_scores = probabilities[0].tolist()
|
||||
score = raw_scores[1]
|
||||
|
||||
label = (
|
||||
model_manager.text_model.config.id2label.get(1, "LABEL_1")
|
||||
)
|
||||
|
||||
return PredictionResult(
|
||||
label=label,
|
||||
score=score,
|
||||
approved=False,
|
||||
raw_scores=raw_scores
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user