added image ai-moderation
This commit is contained in:
115
ai-moderation/app/services/image_moderation_service.py
Normal file
115
ai-moderation/app/services/image_moderation_service.py
Normal 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=[]
|
||||
|
||||
)
|
||||
Reference in New Issue
Block a user