claw / app.py
devserenerocks's picture
Update app.py
bfc1409 verified
Raw History Blame Contribute Delete
4.02 kB
import os
from functools import lru_cache
from io import BytesIO
import requests
import torch
from fastapi import FastAPI, HTTPException
from huggingface_hub import InferenceClient
from PIL import Image
from pydantic import BaseModel, HttpUrl
from transformers import AutoProcessor, ShieldGemma2ForImageClassification, pipeline
app = FastAPI()
client = InferenceClient(api_key=os.environ.get("HF_TOKEN"))
NSFW_MODEL = "Falconsai/nsfw_image_detection"
SHIELD_MODEL = "google/shieldgemma-2-4b-it"
THRESHOLD = 0.5
# ShieldGemma 2 checks one image against each policy separately.
POLICIES = {
"dangerous": "dangerous content",
"sexual": "sexually explicit content",
"violence": "violence or gore",
}
class ChatRequest(BaseModel):
message: str
class ImageRequest(BaseModel):
image_url: HttpUrl
@app.get("/")
def greet_json():
return {"Hello": "World!"}
@app.post("/chat")
def chat(request: ChatRequest):
response = client.chat.completions.create(
model="Qwen/Qwen3-VL-30B-A3B-Instruct",
messages=[{"role": "user", "content": request.message}],
)
return {"response": response.choices[0].message.content}
# ---------- Image moderation ----------
@lru_cache(maxsize=1)
def get_nsfw_classifier():
return pipeline("image-classification", model=NSFW_MODEL)
@lru_cache(maxsize=1)
def get_shield():
token = os.environ.get("HF_TOKEN") # ShieldGemma 2 is a gated model
processor = AutoProcessor.from_pretrained(SHIELD_MODEL, token=token)
model = ShieldGemma2ForImageClassification.from_pretrained(
SHIELD_MODEL,
token=token,
torch_dtype=torch.bfloat16,
device_map="auto",
).eval()
return processor, model
def load_image(url: str) -> Image.Image:
try:
r = requests.get(url, timeout=15)
r.raise_for_status()
return Image.open(BytesIO(r.content)).convert("RGB")
except Exception as e:
raise HTTPException(status_code=400, detail=f"Could not load image: {e}")
def run_nsfw_check(image: Image.Image) -> float:
"""Returns the probability the image is NSFW (Falconsai ViT classifier)."""
results = get_nsfw_classifier()(image)
return next(r["score"] for r in results if r["label"].lower() == "nsfw")
def run_shield_check(image: Image.Image) -> dict:
"""Returns the probability of a violation per ShieldGemma 2 policy."""
processor, model = get_shield()
inputs = processor(
images=[image], policies=list(POLICIES), return_tensors="pt"
).to(model.device)
with torch.inference_mode():
out = model(**inputs)
# probabilities columns are [Yes (violates), No]
probs = out.probabilities[:, 0].float().tolist()
return dict(zip(POLICIES, probs))
@app.post("/moderate-image")
def moderate_image(request: ImageRequest):
try:
image = load_image(str(request.image_url))
nsfw_score = run_nsfw_check(image)
shield_scores = run_shield_check(image)
reasons = []
if nsfw_score >= THRESHOLD:
reasons.append(
f"NSFW classifier flagged the image as explicit ({nsfw_score:.0%} confidence)"
)
for policy, score in shield_scores.items():
if score >= THRESHOLD:
reasons.append(
f"ShieldGemma 2 flagged {POLICIES[policy]} ({score:.0%} confidence)"
)
is_safe = not reasons
return {
"rating": "safe" if is_safe else "not safe",
"reason": (
"No policy violations detected by either model."
if is_safe
else "; ".join(reasons)
),
"scores": {
"nsfw": round(nsfw_score, 4),
**{k: round(v, 4) for k, v in shield_scores.items()},
},
}
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=f"{type(e).__name__}: {e}\n{traceback.format_exc()}")