fixed code
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user