Spaces:
Sleeping
Sleeping
Download app.py from devserenerocks/claw: direct link, hf CLI and curl.
- Browser
- Download file 4.02 kB
-
https://huggingface.co/spaces/devserenerocks/claw/resolve/main/app.py
- Command line
-
hf download hf://spaces/devserenerocks/claw/app.py
-
curl -L -o app.py https://huggingface.co/spaces/devserenerocks/claw/resolve/main/app.py
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 | |
| def greet_json(): | |
| return {"Hello": "World!"} | |
| 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 ---------- | |
| def get_nsfw_classifier(): | |
| return pipeline("image-classification", model=NSFW_MODEL) | |
| 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)) | |
| 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()}") |