fixed code
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -1,17 +1,14 @@
|
||||
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
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
@@ -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())
|
||||
@@ -1,2 +1,5 @@
|
||||
class InvalidImageException(Exception):
|
||||
from app.exceptions.moderation_exception import ModerationException
|
||||
|
||||
|
||||
class InvalidImageException(ModerationException):
|
||||
pass
|
||||
@@ -1,4 +1,4 @@
|
||||
class ModerationException(Exception):
|
||||
|
||||
def __init__(self, message: str):
|
||||
self.message = message
|
||||
self.message = message
|
||||
super().__init__(message)
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
)
|
||||
@@ -3,11 +3,10 @@ from dataclasses import dataclass
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ClipPredictionResult:
|
||||
|
||||
label: str
|
||||
|
||||
score: float
|
||||
|
||||
detected_labels: list[str]
|
||||
|
||||
scores: dict[str, float]
|
||||
scores: dict[str, float]
|
||||
|
||||
@@ -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,
|
||||
@@ -80,4 +66,4 @@ class ImageClassifier:
|
||||
approved=False,
|
||||
|
||||
raw_scores=raw_scores
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
@@ -3,7 +3,6 @@ from dataclasses import dataclass, field
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ImagePredictionResult:
|
||||
|
||||
label: str
|
||||
|
||||
score: float
|
||||
@@ -16,4 +15,4 @@ class ImagePredictionResult:
|
||||
|
||||
detected_labels: list[str] = field(
|
||||
default_factory=list
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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...")
|
||||
def initialize(self):
|
||||
self.load_device()
|
||||
logger.info(f"Using device: {self.device}")
|
||||
self.load_tokenizer()
|
||||
self.load_model()
|
||||
self.load_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()
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PredictionResult:
|
||||
|
||||
label: str
|
||||
|
||||
score: float
|
||||
@@ -13,4 +13,4 @@ class PredictionResult:
|
||||
|
||||
reason: str = ""
|
||||
|
||||
detected_words: list[str] = field(default_factory=list)
|
||||
detected_words: list[str] = field(default_factory=list)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -3,4 +3,4 @@ from pydantic import BaseModel
|
||||
|
||||
class ErrorResponse(BaseModel):
|
||||
code: str
|
||||
message: str
|
||||
message: str
|
||||
|
||||
@@ -4,7 +4,6 @@ from pydantic import BaseModel
|
||||
|
||||
|
||||
class ImageModerationResponse(BaseModel):
|
||||
|
||||
approved: bool
|
||||
|
||||
score: float
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -4,7 +4,6 @@ from app.exceptions.invalid_image_exception import InvalidImageException
|
||||
|
||||
|
||||
class ImageValidator:
|
||||
|
||||
MIN_WIDTH = 50
|
||||
MIN_HEIGHT = 50
|
||||
|
||||
@@ -36,4 +35,4 @@ class ImageValidator:
|
||||
if height > self.MAX_HEIGHT:
|
||||
raise InvalidImageException(
|
||||
f"Image height is too large: {height}"
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
return result[0].normal_form
|
||||
|
||||
@@ -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(
|
||||
@@ -63,4 +61,4 @@ def moderate_image(
|
||||
reason=prediction.reason,
|
||||
label=prediction.label,
|
||||
detected_labels=prediction.detected_labels
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
# =========================
|
||||
@@ -112,4 +97,4 @@ class ImageModerationService:
|
||||
|
||||
detected_labels=[]
|
||||
|
||||
)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
return prediction
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
from app.container import profanity_detector
|
||||
|
||||
|
||||
tests = [
|
||||
"ЗАлУпа",
|
||||
"Ты идиот",
|
||||
"Это идиоты",
|
||||
"Хорошего дня"
|
||||
]
|
||||
|
||||
|
||||
for text in tests:
|
||||
|
||||
result = profanity_detector.detect(text)
|
||||
|
||||
print(
|
||||
text,
|
||||
"=>",
|
||||
result
|
||||
)
|
||||
Reference in New Issue
Block a user