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 DEBUG: bool = True
TEXT_TOXIC_THRESHOLD: float = 0.90
# ===== AI ===== # ===== AI =====
TEXT_MODEL: str = "textdetox/bert-multilingual-toxicity-classifier" 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" MODEL_CACHE_DIR: str = "./models"
DEVICE: str = "auto" 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(): from app.ml.text_classifier import TextClassifier
return model_manager 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 contextlib import asynccontextmanager
from fastapi import FastAPI from fastapi import FastAPI
from app.ml.model_manager import model_manager from app.ml.model_manager import model_manager
from app.config.logging import logger from app.config.logging import logger
from app.config.settings import settings
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
if settings.HF_TOKEN:
os.environ["HF_TOKEN"] = settings.HF_TOKEN
logger.info("Loading AI models...") logger.info("Loading AI models...")
model_manager.initialize() model_manager.initialize()
logger.info("AI models loaded.") logger.info("AI models loaded.")

View File

@@ -1,11 +1,44 @@
from fastapi import FastAPI from fastapi import FastAPI
from app.config.settings import settings from app.config.settings import settings
from app.lifespan import lifespan 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( app = FastAPI(
title=settings.APP_NAME, title=settings.APP_NAME,
version=settings.APP_VERSION, version=settings.APP_VERSION,
lifespan=lifespan 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.settings import settings
from app.config.logging import logger from app.config.logging import logger
from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification
import torch import torch
@@ -13,6 +15,8 @@ class ModelManager:
self.device = None self.device = None
self.text_model = None self.text_model = None
self.text_tokenizer = None self.text_tokenizer = None
self.image_model = None
self.image_processor = None
def load_device(self): def load_device(self):
if settings.DEVICE == "auto": if settings.DEVICE == "auto":
@@ -51,6 +55,35 @@ class ModelManager:
print(f"Using device: {self.device}") print(f"Using device: {self.device}")
self.load_tokenizer() self.load_tokenizer()
self.load_model() 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() model_manager = ModelManager()

View File

@@ -11,4 +11,6 @@ class PredictionResult:
raw_scores: list[float] 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): 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 fastapi import APIRouter
from typing import Annotated
from fastapi import Depends
from app.models.dto.text_request import TextRequest from app.models.dto.text_request import TextRequest
from app.models.response.moderation_response import ModerationResponse from app.models.response.moderation_response import ModerationResponse
from app.services.text_moderation_service import TextModerationService from app.services.text_moderation_service import TextModerationService
from app.core.dependencies import (
get_text_moderation_service
)
router = APIRouter( router = APIRouter(
prefix="/api/v1/moderation", prefix="/api/v1/moderation",
tags=["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) prediction = service.moderate(request.text)
return ModerationResponse( return ModerationResponse(
approved=prediction.approved, approved=prediction.approved,
score=prediction.score, score=prediction.score,
reason=prediction.reason 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.text_classifier import TextClassifier
from app.ml.prediction_result import PredictionResult from app.ml.prediction_result import PredictionResult
from app.config.settings import settings
class TextModerationService: 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 = self.classifier.predict(text)
prediction.approved = ( prediction.approved = (
prediction.score < self.TOXIC_THRESHOLD prediction.score <
settings.TEXT_TOXIC_THRESHOLD
) )
prediction.reason = ( prediction.reason = (
"OK" "OK"
if prediction.approved if prediction.approved
else "TOXIC" else "TOXIC"
) )
return prediction return prediction

View File

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