start ai-moderation
This commit is contained in:
17
.env
17
.env
@@ -0,0 +1,17 @@
|
|||||||
|
APP_NAME=AI Moderation Service
|
||||||
|
|
||||||
|
APP_VERSION=1.0.0
|
||||||
|
|
||||||
|
HOST=0.0.0.0
|
||||||
|
|
||||||
|
PORT=8000
|
||||||
|
|
||||||
|
DEBUG=true
|
||||||
|
|
||||||
|
# ===== AI =====
|
||||||
|
|
||||||
|
TEXT_MODEL=textdetox/bert-multilingual-toxicity-classifier
|
||||||
|
|
||||||
|
MODEL_CACHE_DIR=./models
|
||||||
|
|
||||||
|
DEVICE=auto
|
||||||
|
|||||||
3
.idea/misc.xml
generated
3
.idea/misc.xml
generated
@@ -1,5 +1,8 @@
|
|||||||
<?xml version="1.0" encoding="UTF-8"?>
|
<?xml version="1.0" encoding="UTF-8"?>
|
||||||
<project version="4">
|
<project version="4">
|
||||||
|
<component name="Black">
|
||||||
|
<option name="sdkName" value="Python 3.10 (moderation-post)" />
|
||||||
|
</component>
|
||||||
<component name="ProjectRootManager">
|
<component name="ProjectRootManager">
|
||||||
<output url="file://$PROJECT_DIR$/out" />
|
<output url="file://$PROJECT_DIR$/out" />
|
||||||
</component>
|
</component>
|
||||||
|
|||||||
@@ -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 @@
|
|||||||
|
|
||||||
@@ -2,8 +2,11 @@
|
|||||||
<module type="PYTHON_MODULE" version="4">
|
<module type="PYTHON_MODULE" version="4">
|
||||||
<component name="NewModuleRootManager" inherit-compiler-output="true">
|
<component name="NewModuleRootManager" inherit-compiler-output="true">
|
||||||
<exclude-output />
|
<exclude-output />
|
||||||
<content url="file://$MODULE_DIR$" />
|
<content url="file://$MODULE_DIR$">
|
||||||
<orderEntry type="jdk" jdkName="Python 3.10 (moderation-post)" jdkType="Python SDK" />
|
<sourceFolder url="file://$MODULE_DIR$/ai-moderation" isTestSource="false" />
|
||||||
|
<excludeFolder url="file://$MODULE_DIR$/venv" />
|
||||||
|
</content>
|
||||||
|
<orderEntry type="jdk" jdkName="Python 3.12 (moderation-post)" jdkType="Python SDK" />
|
||||||
<orderEntry type="sourceFolder" forTests="false" />
|
<orderEntry type="sourceFolder" forTests="false" />
|
||||||
</component>
|
</component>
|
||||||
</module>
|
</module>
|
||||||
Reference in New Issue
Block a user