added image ai-moderation

This commit is contained in:
SlimusMinus
2026-07-28 01:10:11 +03:00
parent 1132a774af
commit d376a7417a
34 changed files with 8302 additions and 34 deletions

View File

@@ -0,0 +1,10 @@
# app/config/image_policy.py
FORBIDDEN_IMAGE_LABELS = {
s.lower() for s in {
"a photo containing marijuana",
"a photo containing drugs",
"a photo containing weapons",
"a pornographic photo",
"a photo containing violence",
}
}

View File

@@ -12,11 +12,18 @@ class Settings(BaseSettings):
DEBUG: bool = True
TEXT_TOXIC_THRESHOLD: float = 0.90
# ===== AI =====
TEXT_MODEL: str = "textdetox/bert-multilingual-toxicity-classifier"
IMAGE_MODEL: str = "Falconsai/nsfw_image_detection"
NSFW_THRESHOLD: float = 0.85
IMAGE_CLIP_MODEL: str = "openai/clip-vit-base-patch32"
CLIP_THRESHOLD: float = 0.75
HF_TOKEN: str = "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"
MODEL_CACHE_DIR: str = "./models"
DEVICE: str = "auto"

View File

@@ -0,0 +1,17 @@
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

@@ -0,0 +1,65 @@
from functools import lru_cache
from app.ml.image.image_classifier import ImageClassifier
from app.ml.text_classifier import TextClassifier
from app.moderation.image.validator import ImageValidator
from app.moderation.profanity.detector import ProfanityDetector
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():
global _clip_classifier
if _clip_classifier is None:
_clip_classifier = ClipClassifier()
return _clip_classifier
@lru_cache
def get_text_classifier() -> TextClassifier:
return TextClassifier()
@lru_cache
def get_profanity_detector() -> ProfanityDetector:
return ProfanityDetector(
dictionary=ProfanityDictionary(
file_path="app/resources/profanity_words.txt"
),
lemmatizer=Lemmatizer()
)
@lru_cache
def get_text_moderation_service() -> TextModerationService:
return TextModerationService(
classifier=get_text_classifier(),
profanity_detector=get_profanity_detector()
)
@lru_cache
def get_image_validator():
return ImageValidator()
def get_image_classifier():
return ImageClassifier()
@lru_cache
def get_image_moderation_service():
return ImageModerationService(
classifier=get_image_classifier(),
validator=get_image_validator(),
clip_classifier=get_clip_classifier()
)

View File

@@ -1,4 +1,32 @@
from app.ml.model_manager import model_manager
from functools import lru_cache
def get_model_manager():
return model_manager
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

@@ -0,0 +1,19 @@
from fastapi import Request
from fastapi.responses import JSONResponse
from app.exceptions.invalid_text_exception import InvalidTextException
from app.models.response.error_response import ErrorResponse
async def invalid_text_exception_handler(
request: Request,
exc: InvalidTextException
):
response = ErrorResponse(
code="INVALID_TEXT",
message=exc.message
)
return JSONResponse(
status_code=400,
content=response.model_dump()
)

View File

@@ -0,0 +1,2 @@
class InvalidImageException(Exception):
pass

View File

@@ -0,0 +1,6 @@
from app.exceptions.moderation_exception import ModerationException
class InvalidTextException(ModerationException):
pass

View File

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

View File

@@ -1,11 +1,15 @@
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
@asynccontextmanager
async def lifespan(app: FastAPI):
if settings.HF_TOKEN:
os.environ["HF_TOKEN"] = settings.HF_TOKEN
logger.info("Loading AI models...")
model_manager.initialize()
logger.info("AI models loaded.")

View File

@@ -1,11 +1,44 @@
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
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
)
from app.exceptions.invalid_text_exception import (
InvalidTextException
)
app = FastAPI(
title=settings.APP_NAME,
version=settings.APP_VERSION,
lifespan=lifespan
)
app.include_router(moderation_router)
app.include_router(
moderation_router
)
app.include_router(
image_router
)
app.add_exception_handler(
InvalidTextException,
invalid_text_exception_handler
)

View File

@@ -0,0 +1,126 @@
import torch
from PIL import Image
from transformers import (
CLIPProcessor,
CLIPModel
)
from app.ml.image.clip_prediction_result import ClipPredictionResult
from app.config.settings import settings
class ClipClassifier:
LABELS = [
"a normal photo",
"a photo containing marijuana",
"a photo containing drugs",
"a photo containing weapons",
"a pornographic photo",
"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.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:
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()
}
with torch.no_grad():
outputs = self.model(**inputs)
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
):
detected_labels.append(label)
max_index = torch.argmax(probs).item()
return ClipPredictionResult(
label=self.LABELS[max_index],
score=float(
probs[max_index]
),
detected_labels=detected_labels,
scores=scores
)

View File

@@ -0,0 +1,13 @@
from dataclasses import dataclass
@dataclass(slots=True)
class ClipPredictionResult:
label: str
score: float
detected_labels: list[str]
scores: dict[str, float]

View File

@@ -0,0 +1,83 @@
from PIL import Image
import torch
from app.ml.model_manager import model_manager
from app.ml.image.image_prediction_result import ImagePredictionResult
class ImageClassifier:
def predict(
self,
image: Image.Image
) -> ImagePredictionResult:
inputs = (
model_manager.image_processor(
image,
return_tensors="pt"
)
)
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)
print("Label:", label)
print("Score:", score)
print("Scores:", raw_scores)
return ImagePredictionResult(
label=label,
score=score,
approved=False,
raw_scores=raw_scores
)

View File

@@ -0,0 +1,14 @@
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

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

View File

@@ -3,6 +3,8 @@ from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification
import torch
@@ -13,6 +15,8 @@ class ModelManager:
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":
@@ -51,6 +55,35 @@ class ModelManager:
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.")
model_manager = ModelManager()

View File

@@ -11,4 +11,6 @@ class PredictionResult:
raw_scores: list[float]
reason: str = field(default="")
reason: str = ""
detected_words: list[str] = field(default_factory=list)

View File

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

View File

@@ -0,0 +1,6 @@
from pydantic import BaseModel
class ErrorResponse(BaseModel):
code: str
message: str

View File

@@ -0,0 +1,16 @@
from typing import List
from pydantic import BaseModel
class ImageModerationResponse(BaseModel):
approved: bool
score: float
reason: str
label: str
detected_labels: List[str]

View File

@@ -0,0 +1,39 @@
from PIL import Image
from app.exceptions.invalid_image_exception import InvalidImageException
class ImageValidator:
MIN_WIDTH = 50
MIN_HEIGHT = 50
MAX_WIDTH = 10000
MAX_HEIGHT = 10000
def validate(
self,
image: Image.Image
) -> None:
width, height = image.size
if width < self.MIN_WIDTH:
raise InvalidImageException(
f"Image width is too small: {width}"
)
if height < self.MIN_HEIGHT:
raise InvalidImageException(
f"Image height is too small: {height}"
)
if width > self.MAX_WIDTH:
raise InvalidImageException(
f"Image width is too large: {width}"
)
if height > self.MAX_HEIGHT:
raise InvalidImageException(
f"Image height is too large: {height}"
)

View File

@@ -0,0 +1,51 @@
import re
from app.moderation.profanity.dictionary import ProfanityDictionary
from app.moderation.profanity.lemmatizer import Lemmatizer
class ProfanityDetector:
def __init__(
self,
dictionary: ProfanityDictionary,
lemmatizer: Lemmatizer
):
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)
for bad_word in self.dictionary.words:
if lemma.startswith(bad_word):
result.append(word)
break
return result

View File

@@ -0,0 +1,49 @@
from pathlib import Path
class ProfanityDictionary:
def __init__(
self,
file_path: str
):
self.words = self.load(file_path)
def load(
self,
file_path: str
) -> set[str]:
path = Path(file_path)
if not path.exists():
raise FileNotFoundError(
f"Profanity dictionary not found: {file_path}"
)
with open(
path,
"r",
encoding="utf-8"
) as file:
return {
line.strip().lower().replace("\ufeff", "")
for line in file
if line.strip()
}
def contains(
self,
word: str
) -> bool:
return word in self.words

View File

@@ -0,0 +1,20 @@
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

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,66 @@
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
router = APIRouter(
prefix="/api/v1/moderation",
tags=["Moderation"]
)
MAX_SIZE = 5 * 1024 * 1024 # например, 5 МБ
ALLOWED_CONTENT_TYPES = (
"image/jpeg",
"image/png",
"image/webp",
)
@router.post("/image", response_model=ImageModerationResponse)
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(
status_code=400,
detail="Unsupported image format"
)
# 2. Проверка размера — читаем файл целиком
contents = file.file.read()
if len(contents) > MAX_SIZE:
raise HTTPException(
status_code=413,
detail="Image too large"
)
file.file.seek(0) # обязательно вернуть указатель в начало!
# 3. Только теперь пытаемся открыть изображение
try:
image = Image.open(file.file)
image = image.convert("RGB")
except UnidentifiedImageError:
raise HTTPException(
status_code=400,
detail="Invalid image"
)
prediction = service.moderate(image)
return ImageModerationResponse(
approved=prediction.approved,
score=prediction.score,
reason=prediction.reason,
label=prediction.label,
detected_labels=prediction.detected_labels
)

View File

@@ -1,24 +1,37 @@
from fastapi import APIRouter
from typing import Annotated
from fastapi import Depends
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
)
router = APIRouter(
prefix="/api/v1/moderation",
tags=["Moderation"]
)
service = TextModerationService()
@router.post("/text", response_model=ModerationResponse)
def moderate_text(request: TextRequest):
@router.post(
"/text",
response_model=ModerationResponse
)
def moderate_text(
request: TextRequest,
service: Annotated[
TextModerationService,
Depends(get_text_moderation_service)
]
):
prediction = service.moderate(request.text)
return ModerationResponse(
approved=prediction.approved,
score=prediction.score,
reason=prediction.reason
)
)

View File

@@ -0,0 +1,115 @@
from PIL import Image
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:
def __init__(
self,
classifier: ImageClassifier,
validator: ImageValidator,
clip_classifier: ClipClassifier
):
self.classifier = classifier
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"
nsfw_prediction.detected_labels = [
nsfw_prediction.label
]
return nsfw_prediction
# =========================
# 2. CLIP MODEL
# =========================
clip_prediction = (
self.clip_classifier.predict(image)
)
detected_forbidden = [
label
for label in clip_prediction.detected_labels
if label.lower()
in FORBIDDEN_IMAGE_LABELS
]
if detected_forbidden:
return ImagePredictionResult(
label=clip_prediction.label,
score=clip_prediction.score,
approved=False,
reason="FORBIDDEN_CONTENT",
detected_labels=detected_forbidden
)
# =========================
# 3. NORMAL IMAGE
# =========================
return ImagePredictionResult(
label="normal",
score=1.0,
approved=True,
reason="OK",
detected_labels=[]
)

View File

@@ -1,26 +1,76 @@
from app.exceptions.invalid_text_exception import InvalidTextException
from app.moderation.profanity.detector import ProfanityDetector
from app.ml.text_classifier import TextClassifier
from app.ml.prediction_result import PredictionResult
from app.config.settings import settings
class TextModerationService:
TOXIC_THRESHOLD = 0.80
def __init__(
self,
classifier: TextClassifier,
profanity_detector: ProfanityDetector
):
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",
score=1.0,
approved=False,
raw_scores=[],
reason="PROFANITY",
detected_words=detected_words
)
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.score <
settings.TEXT_TOXIC_THRESHOLD
)
prediction.reason = (
"OK"
if prediction.approved
else "TOXIC"
)
return prediction

View File

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