start ai-moderation
This commit is contained in:
@@ -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
|
||||
)
|
||||
@@ -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")
|
||||
@@ -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()
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
from app.ml.model_manager import model_manager
|
||||
|
||||
def get_model_manager():
|
||||
return model_manager
|
||||
@@ -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...")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class BaseClassifier(ABC):
|
||||
@abstractmethod
|
||||
def predict(self, value):
|
||||
pass
|
||||
|
||||
@@ -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()
|
||||
@@ -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="")
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
class TextRequest(BaseModel):
|
||||
text: str = Field(..., description="The text to be moderated")
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -1 +0,0 @@
|
||||
|
||||
Reference in New Issue
Block a user