fixed code

This commit is contained in:
SlimusMinus
2026-08-05 01:44:29 +03:00
parent 72159dbe4a
commit d33f0eaa78
30 changed files with 143 additions and 375 deletions

View File

@@ -1,6 +1,5 @@
from pydantic_settings import BaseSettings, SettingsConfigDict from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings): class Settings(BaseSettings):
APP_NAME: str = "AI Moderation" APP_NAME: str = "AI Moderation"
@@ -14,7 +13,6 @@ class Settings(BaseSettings):
TEXT_TOXIC_THRESHOLD: float = 0.90 TEXT_TOXIC_THRESHOLD: float = 0.90
# ===== AI ===== # ===== AI =====
TEXT_MODEL: str = "textdetox/bert-multilingual-toxicity-classifier" TEXT_MODEL: str = "textdetox/bert-multilingual-toxicity-classifier"
@@ -33,4 +31,5 @@ class Settings(BaseSettings):
extra="ignore" extra="ignore"
) )
settings = Settings() 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.dictionary import ProfanityDictionary
from app.moderation.profanity.lemmatizer import Lemmatizer from app.moderation.profanity.lemmatizer import Lemmatizer
from app.moderation.profanity.detector import ProfanityDetector
dictionary = ProfanityDictionary( dictionary = ProfanityDictionary(
"app/resources/profanity_words.txt" "app/resources/profanity_words.txt"
) )
lemmatizer = Lemmatizer() lemmatizer = Lemmatizer()
profanity_detector = ProfanityDetector( profanity_detector = ProfanityDetector(
dictionary, dictionary,
lemmatizer lemmatizer

View File

@@ -1,5 +1,6 @@
from functools import lru_cache from functools import lru_cache
from app.ml.image.clip_classifier import ClipClassifier
from app.ml.image.image_classifier import ImageClassifier from app.ml.image.image_classifier import ImageClassifier
from app.ml.text_classifier import TextClassifier from app.ml.text_classifier import TextClassifier
from app.moderation.image.validator import ImageValidator 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.moderation.profanity.lemmatizer import Lemmatizer
from app.services.image_moderation_service import ImageModerationService from app.services.image_moderation_service import ImageModerationService
from app.services.text_moderation_service import TextModerationService from app.services.text_moderation_service import TextModerationService
from app.ml.image.clip_classifier import ClipClassifier
_clip_classifier = None
@lru_cache @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 @lru_cache
def get_text_classifier() -> TextClassifier: def get_text_classifier() -> TextClassifier:
@@ -45,19 +38,19 @@ def get_text_moderation_service() -> TextModerationService:
profanity_detector=get_profanity_detector() profanity_detector=get_profanity_detector()
) )
@lru_cache
def get_image_validator():
@lru_cache
def get_image_validator() -> ImageValidator:
return ImageValidator() return ImageValidator()
def get_image_classifier(): @lru_cache
def get_image_classifier() -> ImageClassifier:
return ImageClassifier() return ImageClassifier()
@lru_cache @lru_cache
def get_image_moderation_service(): def get_image_moderation_service() -> ImageModerationService:
return ImageModerationService( return ImageModerationService(
classifier=get_image_classifier(), classifier=get_image_classifier(),
validator=get_image_validator(), 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 fastapi.responses import JSONResponse
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.models.response.error_response import ErrorResponse from app.models.response.error_response import ErrorResponse
async def invalid_text_exception_handler( async def invalid_text_exception_handler(
request: Request, request: Request,
exc: InvalidTextException exc: InvalidTextException
@@ -12,8 +14,15 @@ async def invalid_text_exception_handler(
code="INVALID_TEXT", code="INVALID_TEXT",
message=exc.message message=exc.message
) )
return JSONResponse(status_code=400, content=response.model_dump())
return JSONResponse(
status_code=400, async def invalid_image_exception_handler(
content=response.model_dump() 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 pass

View File

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

View File

@@ -1,9 +1,12 @@
import os 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.config.logging import logger from app.config.logging import logger
from app.config.settings import settings from app.config.settings import settings
from app.ml.model_manager import model_manager
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):

View File

@@ -1,25 +1,15 @@
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.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 ( 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 ( from app.exceptions.invalid_image_exception import InvalidImageException
InvalidTextException 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( app = FastAPI(
title=settings.APP_NAME, title=settings.APP_NAME,
@@ -27,18 +17,8 @@ app = FastAPI(
lifespan=lifespan lifespan=lifespan
) )
app.include_router(moderation_router)
app.include_router(image_router)
app.include_router( app.add_exception_handler(InvalidTextException, invalid_text_exception_handler)
moderation_router app.add_exception_handler(InvalidImageException, invalid_image_exception_handler)
)
app.include_router(
image_router
)
app.add_exception_handler(
InvalidTextException,
invalid_text_exception_handler
)

View File

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

View File

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

View File

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

View File

@@ -1,89 +1,69 @@
from transformers import AutoTokenizer import torch
from transformers import AutoModelForSequenceClassification
from app.config.settings import settings
from app.config.logging import logger
from transformers import AutoImageProcessor from transformers import AutoImageProcessor
from transformers import AutoModelForImageClassification 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: class ModelManager:
def __init__(self): def __init__(self):
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_model = None
self.image_processor = None self.image_processor = None
def load_device(self): def load_device(self):
if settings.DEVICE == "auto": if settings.DEVICE == "auto":
self.device = torch.device( self.device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu" "cuda" if torch.cuda.is_available() else "cpu"
) )
else: else:
self.device = torch.device(settings.DEVICE) self.device = torch.device(settings.DEVICE)
def load_tokenizer(self): def load_tokenizer(self):
print("Loading tokenizer...") 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( def load_model(self):
settings.TEXT_MODEL, logger.info("Loading model...")
cache_dir=settings.MODEL_CACHE_DIR 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): def initialize(self):
print("Loading model...") self.load_device()
logger.info(f"Using device: {self.device}")
self.text_model = AutoModelForSequenceClassification.from_pretrained( self.load_tokenizer()
settings.TEXT_MODEL, self.load_model()
cache_dir=settings.MODEL_CACHE_DIR self.load_image_model()
)
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.")
model_manager = ModelManager() model_manager = ModelManager()

View File

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

View File

@@ -1,9 +1,9 @@
import torch import torch
from app.ml.model_manager import model_manager from app.ml.model_manager import model_manager
from app.config.logging import logger
from app.ml.prediction_result import PredictionResult from app.ml.prediction_result import PredictionResult
class TextClassifier: class TextClassifier:
def predict(self, text: str): def predict(self, text: str):
inputs = model_manager.text_tokenizer( inputs = model_manager.text_tokenizer(

View File

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

View File

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

View File

@@ -1,6 +1,9 @@
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
class ModerationResponse(BaseModel): class ModerationResponse(BaseModel):
approved: bool = Field(..., description="Indicates if the text is approved or not") approved: bool = Field(...,
score: float = Field(..., description="The score indicating the level of toxicity or appropriateness of the text") description="Indicates if the text is approved or not")
reason: str = Field(..., description="The reason for the moderation decision") 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: class ImageValidator:
MIN_WIDTH = 50 MIN_WIDTH = 50
MIN_HEIGHT = 50 MIN_HEIGHT = 50

View File

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

View File

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

View File

@@ -3,18 +3,13 @@ import pymorphy3
class Lemmatizer: class Lemmatizer:
def __init__(self): def __init__(self):
self.morph = pymorphy3.MorphAnalyzer() self.morph = pymorphy3.MorphAnalyzer()
def normalize( def normalize(
self, self,
word: str word: str
) -> str: ) -> str:
result = self.morph.parse(word) result = self.morph.parse(word)
return result[0].normal_form 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 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.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( router = APIRouter(
prefix="/api/v1/moderation", prefix="/api/v1/moderation",
@@ -26,7 +25,6 @@ def moderate_image(
file: UploadFile = File(...), file: UploadFile = File(...),
service: ImageModerationService = Depends(get_image_moderation_service) service: ImageModerationService = Depends(get_image_moderation_service)
): ):
# 1. Проверка content-type — самая дешёвая, делаем её первой # 1. Проверка content-type — самая дешёвая, делаем её первой
if file.content_type not in ALLOWED_CONTENT_TYPES: if file.content_type not in ALLOWED_CONTENT_TYPES:
raise HTTPException( 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.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 ( from app.core.dependencies import (
get_text_moderation_service get_text_moderation_service
) )

View File

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

View File

@@ -16,28 +16,21 @@ class TextModerationService:
self.classifier = classifier self.classifier = classifier
self.profanity_detector = profanity_detector self.profanity_detector = profanity_detector
def moderate( def moderate(
self, self,
text: str text: str
) -> PredictionResult: ) -> PredictionResult:
if text is None or not text.strip(): if text is None or not text.strip():
raise InvalidTextException( raise InvalidTextException(
"Text is empty" "Text is empty"
) )
detected_words = ( detected_words = (
self.profanity_detector.detect(text) self.profanity_detector.detect(text)
) )
if detected_words: if detected_words:
return PredictionResult( return PredictionResult(
label="PROFANITY", label="PROFANITY",
@@ -54,23 +47,17 @@ class TextModerationService:
) )
prediction = self.classifier.predict(text) prediction = self.classifier.predict(text)
prediction.approved = ( prediction.approved = (
prediction.score < prediction.score <
settings.TEXT_TOXIC_THRESHOLD 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 +0,0 @@
from app.container import profanity_detector
tests = [
"ЗАлУпа",
"Ты идиот",
"Это идиоты",
"Хорошего дня"
]
for text in tests:
result = profanity_detector.detect(text)
print(
text,
"=>",
result
)