start ai-moderation

This commit is contained in:
SlimusMinus
2026-07-17 01:06:36 +03:00
parent eb6556212b
commit 1132a774af
18 changed files with 297 additions and 3 deletions

View File

@@ -0,0 +1,7 @@
from abc import ABC, abstractmethod
class BaseClassifier(ABC):
@abstractmethod
def predict(self, value):
pass

View File

@@ -0,0 +1,56 @@
from transformers import AutoTokenizer
from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
import torch
class ModelManager:
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
def load_device(self):
if settings.DEVICE == "auto":
self.device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
else:
self.device = torch.device(settings.DEVICE)
def load_tokenizer(self):
print("Loading tokenizer...")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
print("Tokenizer loaded.")
def load_model(self):
print("Loading model...")
self.text_model = AutoModelForSequenceClassification.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.text_model.to(self.device)
self.text_model.eval()
print(self.text_model.config.id2label)
print("Text model loaded.")
def initialize(self):
self.load_device()
print(f"Using device: {self.device}")
self.load_tokenizer()
self.load_model()
model_manager = ModelManager()

View File

@@ -0,0 +1,14 @@
from dataclasses import dataclass, field
@dataclass(slots=True)
class PredictionResult:
label: str
score: float
approved: bool
raw_scores: list[float]
reason: str = field(default="")

View File

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