Merge pull request #2 from SlimusMinus/fix-code

fixed code
This commit is contained in:
SlimusMinus
2026-08-05 01:45:31 +03:00
committed by GitHub
30 changed files with 143 additions and 375 deletions

View File

@@ -1,6 +1,5 @@
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
APP_NAME: str = "AI Moderation"
@@ -14,7 +13,6 @@ class Settings(BaseSettings):
TEXT_TOXIC_THRESHOLD: float = 0.90
# ===== AI =====
TEXT_MODEL: str = "textdetox/bert-multilingual-toxicity-classifier"
@@ -33,4 +31,5 @@ class Settings(BaseSettings):
extra="ignore"
)
settings = Settings()

View File

@@ -1,16 +1,13 @@
from app.moderation.profanity.detector import ProfanityDetector
from app.moderation.profanity.dictionary import ProfanityDictionary
from app.moderation.profanity.lemmatizer import Lemmatizer
from app.moderation.profanity.detector import ProfanityDetector
dictionary = ProfanityDictionary(
"app/resources/profanity_words.txt"
)
lemmatizer = Lemmatizer()
profanity_detector = ProfanityDetector(
dictionary,
lemmatizer

View File

@@ -1,5 +1,6 @@
from functools import lru_cache
from app.ml.image.clip_classifier import ClipClassifier
from app.ml.image.image_classifier import ImageClassifier
from app.ml.text_classifier import TextClassifier
from app.moderation.image.validator import ImageValidator
@@ -8,20 +9,12 @@ from app.moderation.profanity.dictionary import ProfanityDictionary
from app.moderation.profanity.lemmatizer import Lemmatizer
from app.services.image_moderation_service import ImageModerationService
from app.services.text_moderation_service import TextModerationService
from app.ml.image.clip_classifier import ClipClassifier
_clip_classifier = None
@lru_cache
def get_clip_classifier():
def get_clip_classifier() -> ClipClassifier:
return ClipClassifier()
global _clip_classifier
if _clip_classifier is None:
_clip_classifier = ClipClassifier()
return _clip_classifier
@lru_cache
def get_text_classifier() -> TextClassifier:
@@ -45,19 +38,19 @@ def get_text_moderation_service() -> TextModerationService:
profanity_detector=get_profanity_detector()
)
@lru_cache
def get_image_validator():
@lru_cache
def get_image_validator() -> ImageValidator:
return ImageValidator()
def get_image_classifier():
@lru_cache
def get_image_classifier() -> ImageClassifier:
return ImageClassifier()
@lru_cache
def get_image_moderation_service():
def get_image_moderation_service() -> ImageModerationService:
return ImageModerationService(
classifier=get_image_classifier(),
validator=get_image_validator(),

View File

@@ -1,32 +0,0 @@
from functools import lru_cache
from app.ml.text_classifier import TextClassifier
from app.services.text_moderation_service import TextModerationService
from app.ml.image.image_classifier import ImageClassifier
from app.services.image_moderation_service import ImageModerationService
from app.container import profanity_detector
@lru_cache
def get_text_classifier() -> TextClassifier:
return TextClassifier()
@lru_cache
def get_text_moderation_service() -> TextModerationService:
return TextModerationService(
classifier=get_text_classifier(),
profanity_detector=profanity_detector
)
@lru_cache
def get_image_classifier():
return ImageClassifier()
@lru_cache
def get_image_moderation_service():
return ImageModerationService(
classifier=get_image_classifier()
)

View File

@@ -2,8 +2,10 @@ from fastapi import Request
from fastapi.responses import JSONResponse
from app.exceptions.invalid_text_exception import InvalidTextException
from app.exceptions.invalid_image_exception import InvalidImageException
from app.models.response.error_response import ErrorResponse
async def invalid_text_exception_handler(
request: Request,
exc: InvalidTextException
@@ -12,8 +14,15 @@ async def invalid_text_exception_handler(
code="INVALID_TEXT",
message=exc.message
)
return JSONResponse(status_code=400, content=response.model_dump())
return JSONResponse(
status_code=400,
content=response.model_dump()
async def invalid_image_exception_handler(
request: Request,
exc: InvalidImageException
):
response = ErrorResponse(
code="INVALID_IMAGE",
message=exc.message
)
return JSONResponse(status_code=400, content=response.model_dump())

View File

@@ -1,2 +1,5 @@
class InvalidImageException(Exception):
from app.exceptions.moderation_exception import ModerationException
class InvalidImageException(ModerationException):
pass

View File

@@ -1,4 +1,4 @@
class ModerationException(Exception):
def __init__(self, message: str):
self.message = message
super().__init__(message)

View File

@@ -1,9 +1,12 @@
import os
from contextlib import asynccontextmanager
from fastapi import FastAPI
from app.ml.model_manager import model_manager
from app.config.logging import logger
from app.config.settings import settings
from app.ml.model_manager import model_manager
@asynccontextmanager
async def lifespan(app: FastAPI):

View File

@@ -1,25 +1,15 @@
from fastapi import FastAPI
from app.config.settings import settings
from app.lifespan import lifespan
from app.routers.text_moderation_router import (
router as moderation_router
)
from app.routers.image_moderation_router import (
router as image_router
)
from app.exceptions.handlers import (
invalid_text_exception_handler
invalid_text_exception_handler,
invalid_image_exception_handler
)
from app.exceptions.invalid_text_exception import (
InvalidTextException
)
from app.exceptions.invalid_text_exception import InvalidTextException
from app.exceptions.invalid_image_exception import InvalidImageException
from app.lifespan import lifespan
from app.routers.image_moderation_router import router as image_router
from app.routers.text_moderation_router import router as moderation_router
app = FastAPI(
title=settings.APP_NAME,
@@ -27,18 +17,8 @@ app = FastAPI(
lifespan=lifespan
)
app.include_router(moderation_router)
app.include_router(image_router)
app.include_router(
moderation_router
)
app.include_router(
image_router
)
app.add_exception_handler(
InvalidTextException,
invalid_text_exception_handler
)
app.add_exception_handler(InvalidTextException, invalid_text_exception_handler)
app.add_exception_handler(InvalidImageException, invalid_image_exception_handler)

View File

@@ -1,17 +1,12 @@
import torch
from PIL import Image
from transformers import (
CLIPProcessor,
CLIPModel
)
from transformers import CLIPProcessor, CLIPModel
from app.ml.image.clip_prediction_result import ClipPredictionResult
from app.config.settings import settings
from app.ml.image.clip_prediction_result import ClipPredictionResult
class ClipClassifier:
LABELS = [
"a normal photo",
"a photo containing marijuana",
@@ -21,106 +16,46 @@ class ClipClassifier:
"a photo containing violence"
]
def __init__(self):
self.device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Loading CLIP on {self.device}")
self.device = (
"cuda"
if torch.cuda.is_available()
else "cpu"
)
print(
f"Loading CLIP on {self.device}"
)
self.model = CLIPModel.from_pretrained(
settings.IMAGE_CLIP_MODEL
)
self.processor = CLIPProcessor.from_pretrained(
settings.IMAGE_CLIP_MODEL
)
self.model = CLIPModel.from_pretrained(settings.IMAGE_CLIP_MODEL)
self.processor = CLIPProcessor.from_pretrained(settings.IMAGE_CLIP_MODEL)
self.model.to(self.device)
self.model.eval()
def predict(
self,
image: Image.Image
) -> ClipPredictionResult:
def predict(self, image: Image.Image) -> ClipPredictionResult:
inputs = self.processor(
text=self.LABELS,
images=image,
return_tensors="pt",
padding=True
)
inputs = {
k: v.to(self.device)
for k, v in inputs.items()
}
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model(**inputs)
logits = (
outputs
.logits_per_image
)
probs = (
logits.softmax(dim=1)[0]
)
logits = outputs.logits_per_image
probs = logits.softmax(dim=1)[0]
scores = {}
detected_labels = []
THRESHOLD = 0.75
for index, probability in enumerate(probs):
label = self.LABELS[index]
score = float(probability)
scores[label] = score
if (
label != "a normal photo"
and score >= THRESHOLD
):
if label != "a normal photo" and score >= settings.CLIP_THRESHOLD:
detected_labels.append(label)
max_index = torch.argmax(probs).item()
return ClipPredictionResult(
label=self.LABELS[max_index],
score=float(
probs[max_index]
),
score=float(probs[max_index]),
detected_labels=detected_labels,
scores=scores
)

View File

@@ -3,7 +3,6 @@ from dataclasses import dataclass
@dataclass(slots=True)
class ClipPredictionResult:
label: str
score: float

View File

@@ -1,19 +1,16 @@
from PIL import Image
import torch
from PIL import Image
from app.ml.model_manager import model_manager
from app.ml.image.image_prediction_result import ImagePredictionResult
from app.ml.model_manager import model_manager
class ImageClassifier:
def predict(
self,
image: Image.Image
) -> ImagePredictionResult:
inputs = (
model_manager.image_processor(
image,
@@ -21,48 +18,38 @@ class ImageClassifier:
)
)
inputs = {
key: value.to(model_manager.device)
for key, value in inputs.items()
}
with torch.no_grad():
outputs = (
model_manager.image_model(**inputs)
)
probabilities = torch.softmax(
outputs.logits,
dim=1
)
raw_scores = probabilities[0].tolist()
predicted_index = (
torch.argmax(probabilities, dim=1)
.item()
)
labels = (
model_manager.image_model
.config
.id2label
)
label = labels[predicted_index]
score = raw_scores[predicted_index]
print("====================")
print("IMAGE RESULT")
print("Labels:", labels)
@@ -70,7 +57,6 @@ class ImageClassifier:
print("Score:", score)
print("Scores:", raw_scores)
return ImagePredictionResult(
label=label,

View File

@@ -1,14 +0,0 @@
from dataclasses import dataclass
from dataclasses import field
@dataclass(slots=True)
class ImagePredictionResult:
approved: bool
score: float
reason: str
detected_labels: list[str] = field(default_factory=list)

View File

@@ -3,7 +3,6 @@ from dataclasses import dataclass, field
@dataclass(slots=True)
class ImagePredictionResult:
label: str
score: float

View File

@@ -1,89 +1,69 @@
from transformers import AutoTokenizer
from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
import torch
from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification
from transformers import AutoModelForSequenceClassification
from transformers import AutoTokenizer
import torch
from app.config.logging import logger
from app.config.settings import settings
class ModelManager:
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
self.image_model = None
self.image_processor = None
def __init__(self):
self.device = None
self.text_model = None
self.text_tokenizer = None
self.image_model = None
self.image_processor = 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_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...")
def load_tokenizer(self):
logger.info("Loading tokenizer...")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
logger.info("Tokenizer loaded.")
self.text_tokenizer = AutoTokenizer.from_pretrained(
settings.TEXT_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
def load_model(self):
logger.info("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()
logger.debug("id2label: %s", self.text_model.config.id2label)
logger.info("Text model loaded.")
print("Tokenizer loaded.")
def load_image_model(self):
logger.info("Loading image model...")
self.image_processor = AutoImageProcessor.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.image_model = AutoModelForImageClassification.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
self.image_model.to(self.device)
self.image_model.eval()
logger.info("Image model 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()
self.load_image_model()
def load_image_model(self):
print("Loading image model...")
self.image_processor = (
AutoImageProcessor.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
)
self.image_model = (
AutoModelForImageClassification.from_pretrained(
settings.IMAGE_MODEL,
cache_dir=settings.MODEL_CACHE_DIR
)
)
self.image_model.to(self.device)
self.image_model.eval()
print("Image model loaded.")
def initialize(self):
self.load_device()
logger.info(f"Using device: {self.device}")
self.load_tokenizer()
self.load_model()
self.load_image_model()
model_manager = ModelManager()

View File

@@ -1,8 +1,8 @@
from dataclasses import dataclass, field
@dataclass(slots=True)
class PredictionResult:
label: str
score: float

View File

@@ -1,9 +1,9 @@
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(

View File

@@ -3,8 +3,7 @@ from pydantic import Field
class TextRequest(BaseModel):
text: str = Field(
min_length=1,
max_length=5000
)
text: str = Field(
min_length=1,
max_length=5000
)

View File

@@ -4,7 +4,6 @@ from pydantic import BaseModel
class ImageModerationResponse(BaseModel):
approved: bool
score: float

View File

@@ -1,6 +1,9 @@
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")
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

@@ -4,7 +4,6 @@ from app.exceptions.invalid_image_exception import InvalidImageException
class ImageValidator:
MIN_WIDTH = 50
MIN_HEIGHT = 50

View File

@@ -3,11 +3,8 @@ import re
from app.moderation.profanity.dictionary import ProfanityDictionary
from app.moderation.profanity.lemmatizer import Lemmatizer
class ProfanityDetector:
def __init__(
self,
dictionary: ProfanityDictionary,
@@ -17,23 +14,18 @@ class ProfanityDetector:
self.dictionary = dictionary
self.lemmatizer = lemmatizer
def detect(
self,
text: str
) -> list[str]:
words = re.findall(
r"[а-яА-ЯёЁ]+",
text.lower()
)
result = []
for word in words:
lemma = self.lemmatizer.normalize(word)

View File

@@ -1,9 +1,7 @@
from pathlib import Path
class ProfanityDictionary:
def __init__(
self,
file_path: str
@@ -20,7 +18,6 @@ class ProfanityDictionary:
path = Path(file_path)
if not path.exists():
raise FileNotFoundError(
f"Profanity dictionary not found: {file_path}"
@@ -39,8 +36,6 @@ class ProfanityDictionary:
if line.strip()
}
def contains(
self,
word: str

View File

@@ -3,18 +3,13 @@ import pymorphy3
class Lemmatizer:
def __init__(self):
self.morph = pymorphy3.MorphAnalyzer()
def normalize(
self,
word: str
) -> str:
result = self.morph.parse(word)
return result[0].normal_form

View File

@@ -1,11 +1,10 @@
from PIL import Image, UnidentifiedImageError
from fastapi import APIRouter, UploadFile, File, Depends, HTTPException
from PIL import Image, UnidentifiedImageError
from app.services.image_moderation_service import ImageModerationService
from app.core.dependencies import get_image_moderation_service
from app.models.response.image_moderation_response import ImageModerationResponse
from app.models.response.image_moderation_response import \
ImageModerationResponse
from app.services.image_moderation_service import ImageModerationService
router = APIRouter(
prefix="/api/v1/moderation",
@@ -26,7 +25,6 @@ def moderate_image(
file: UploadFile = File(...),
service: ImageModerationService = Depends(get_image_moderation_service)
):
# 1. Проверка content-type — самая дешёвая, делаем её первой
if file.content_type not in ALLOWED_CONTENT_TYPES:
raise HTTPException(

View File

@@ -6,7 +6,6 @@ from app.models.dto.text_request import TextRequest
from app.models.response.moderation_response import ModerationResponse
from app.services.text_moderation_service import TextModerationService
from app.core.dependencies import (
get_text_moderation_service
)

View File

@@ -1,12 +1,11 @@
from PIL import Image
from app.config.image_policy import FORBIDDEN_IMAGE_LABELS
from app.config.settings import settings
from app.ml.image.clip_classifier import ClipClassifier
from app.ml.image.image_classifier import ImageClassifier
from app.ml.image.image_prediction_result import ImagePredictionResult
from app.ml.image.clip_classifier import ClipClassifier
from app.moderation.image.validator import ImageValidator
from app.config.settings import settings
from app.config.image_policy import FORBIDDEN_IMAGE_LABELS
class ImageModerationService:
@@ -21,29 +20,23 @@ class ImageModerationService:
self.clip_classifier = clip_classifier
self.validator = validator
def moderate(
self,
image: Image.Image
) -> ImagePredictionResult:
self.validator.validate(image)
# =========================
# 1. NSFW MODEL
# =========================
nsfw_prediction = self.classifier.predict(image)
if (
nsfw_prediction.label.lower() == "nsfw"
and nsfw_prediction.score >= settings.NSFW_THRESHOLD
):
nsfw_prediction.approved = False
nsfw_prediction.reason = "NSFW"
@@ -54,8 +47,6 @@ class ImageModerationService:
return nsfw_prediction
# =========================
# 2. CLIP MODEL
# =========================
@@ -64,7 +55,6 @@ class ImageModerationService:
self.clip_classifier.predict(image)
)
detected_forbidden = [
label
@@ -76,10 +66,7 @@ class ImageModerationService:
]
if detected_forbidden:
return ImagePredictionResult(
label=clip_prediction.label,
@@ -94,8 +81,6 @@ class ImageModerationService:
)
# =========================
# 3. NORMAL IMAGE
# =========================

View File

@@ -16,28 +16,21 @@ class TextModerationService:
self.classifier = classifier
self.profanity_detector = profanity_detector
def moderate(
self,
text: str
) -> PredictionResult:
if text is None or not text.strip():
raise InvalidTextException(
"Text is empty"
)
detected_words = (
self.profanity_detector.detect(text)
)
if detected_words:
return PredictionResult(
label="PROFANITY",
@@ -54,23 +47,17 @@ class TextModerationService:
)
prediction = self.classifier.predict(text)
prediction.approved = (
prediction.score <
settings.TEXT_TOXIC_THRESHOLD
)
prediction.reason = (
"OK"
if prediction.approved
else "TOXIC"
)
return prediction

View File

@@ -1,20 +0,0 @@
from app.container import profanity_detector
tests = [
"ЗАлУпа",
"Ты идиот",
"Это идиоты",
"Хорошего дня"
]
for text in tests:
result = profanity_detector.detect(text)
print(
text,
"=>",
result
)