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,24 @@
from fastapi import APIRouter
from app.models.dto.text_request import TextRequest
from app.models.response.moderation_response import ModerationResponse
from app.services.text_moderation_service import TextModerationService
router = APIRouter(
prefix="/api/v1/moderation",
tags=["Moderation"]
)
service = TextModerationService()
@router.post("/text", response_model=ModerationResponse)
def moderate_text(request: TextRequest):
prediction = service.moderate(request.text)
return ModerationResponse(
approved=prediction.approved,
score=prediction.score,
reason=prediction.reason
)

View File

@@ -0,0 +1,11 @@
import logging
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(message)s",
handlers=[
logging.StreamHandler()
]
)
logger = logging.getLogger("ai-moderation")

View File

@@ -0,0 +1,29 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
APP_NAME: str = "AI Moderation"
APP_VERSION: str = "1.0.0"
HOST: str = "0.0.0.0"
PORT: int = 8000
DEBUG: bool = True
# ===== AI =====
TEXT_MODEL: str = "textdetox/bert-multilingual-toxicity-classifier"
MODEL_CACHE_DIR: str = "./models"
DEVICE: str = "auto"
model_config = SettingsConfigDict(
env_file=".env",
extra="ignore"
)
settings = Settings()

View File

@@ -0,0 +1,4 @@
from app.ml.model_manager import model_manager
def get_model_manager():
return model_manager

View File

@@ -0,0 +1,13 @@
from contextlib import asynccontextmanager
from fastapi import FastAPI
from app.ml.model_manager import model_manager
from app.config.logging import logger
@asynccontextmanager
async def lifespan(app: FastAPI):
logger.info("Loading AI models...")
model_manager.initialize()
logger.info("AI models loaded.")
yield
logger.info("Stopping AI Service...")

View File

@@ -0,0 +1,11 @@
from fastapi import FastAPI
from app.config.settings import settings
from app.lifespan import lifespan
from app.api.text_moderation_router import router as moderation_router
app = FastAPI(
title=settings.APP_NAME,
version=settings.APP_VERSION,
lifespan=lifespan
)
app.include_router(moderation_router)

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
)

View File

@@ -0,0 +1,4 @@
from pydantic import BaseModel, Field
class TextRequest(BaseModel):
text: str = Field(..., description="The text to be moderated")

View File

@@ -0,0 +1,6 @@
from pydantic import BaseModel, Field
class ModerationResponse(BaseModel):
approved: bool = Field(..., description="Indicates if the text is approved or not")
score: float = Field(..., description="The score indicating the level of toxicity or appropriateness of the text")
reason: str = Field(..., description="The reason for the moderation decision")

View File

@@ -0,0 +1,26 @@
from app.ml.text_classifier import TextClassifier
from app.ml.prediction_result import PredictionResult
class TextModerationService:
TOXIC_THRESHOLD = 0.80
def __init__(self):
self.classifier = TextClassifier()
def moderate(self, text: str) -> PredictionResult:
prediction = self.classifier.predict(text)
prediction.approved = (
prediction.score < self.TOXIC_THRESHOLD
)
prediction.reason = (
"OK"
if prediction.approved
else "TOXIC"
)
return prediction

View File

@@ -0,0 +1,20 @@
from app.ml.model_manager import model_manager
from app.services.text_moderation_service import TextModerationService
model_manager.initialize()
service = TextModerationService()
texts = [
"Привет",
"Спасибо",
"Ты идиот",
"Хорошего дня"
]
for text in texts:
result = service.moderate(text)
print("=" * 50)
print(text)
print(result)