Spaces:
Sleeping
Sleeping
Download app.py from misukisu/moderation-api: direct link, hf CLI and curl.
- Browser
- Download file 3.66 kB
-
https://huggingface.co/spaces/misukisu/moderation-api/resolve/main/app.py
- Command line
-
hf download hf://spaces/misukisu/moderation-api/app.py
-
curl -L -o app.py https://huggingface.co/spaces/misukisu/moderation-api/resolve/main/app.py
3.66 kB
| import json | |
| import numpy as np, onnxruntime as ort, gradio as gr | |
| from fastapi import FastAPI | |
| from pydantic import BaseModel, Field | |
| from huggingface_hub import snapshot_download | |
| from tokenizers import Tokenizer | |
| REPO, ONNX = "misukisu/moderation-modernbert-en", "onnx/model.onnx" | |
| d = snapshot_download(REPO, allow_patterns=[ONNX, "tokenizer.json", "config.json", "thresholds.json"]) | |
| tok = Tokenizer.from_file(f"{d}/tokenizer.json") | |
| tok.enable_truncation(256) | |
| tok.enable_padding() | |
| so = ort.SessionOptions() | |
| so.intra_op_num_threads = 2 | |
| sess = ort.InferenceSession(f"{d}/{ONNX}", so, providers=["CPUExecutionProvider"]) | |
| id2label = json.load(open(f"{d}/config.json"))["id2label"] | |
| LABELS = [id2label[str(i)] for i in range(len(id2label))] | |
| TH = json.load(open(f"{d}/thresholds.json")) | |
| MAX_BATCH = 32 | |
| def _run(texts): | |
| texts = [str(t)[:4000] for t in texts][:MAX_BATCH] | |
| if not texts: | |
| return [] | |
| enc = tok.encode_batch(texts) | |
| ids = np.array([e.ids for e in enc], dtype=np.int64) | |
| mask = np.array([e.attention_mask for e in enc], dtype=np.int64) | |
| probs = sess.run(None, {"input_ids": ids, "attention_mask": mask})[0] | |
| return [{"flagged": bool(any(p[i] >= TH[l] for i, l in enumerate(LABELS))), | |
| "categories": [l for i, l in enumerate(LABELS) if p[i] >= TH[l]], | |
| "scores": {l: round(float(p[i]), 4) for i, l in enumerate(LABELS)}} for p in probs] | |
| def moderate(text: str) -> dict: | |
| """Moderate one English text. Returns flagged, categories and per-label scores.""" | |
| return _run([text or ""])[0] | |
| def moderate_batch(texts: list) -> list: | |
| """Moderate up to 32 texts at once.""" | |
| return _run(texts or []) | |
| def ui(text): | |
| r = moderate(text) | |
| head = "### Flagged: " + ", ".join(r["categories"]) if r["flagged"] else "### Clean" | |
| return head, {l: s for l, s in r["scores"].items()}, r | |
| with gr.Blocks(title="Moderation API") as demo: | |
| gr.Markdown("# Moderation API\nEnglish content moderation with ModernBERT, 13 labels. " | |
| "Free to use from code, see the **Use via API** link at the bottom. " | |
| f"Model: [{REPO}](https://huggingface.co/{REPO})") | |
| with gr.Row(): | |
| with gr.Column(): | |
| inp = gr.Textbox(label="Text", lines=5, placeholder="Paste a message...") | |
| btn = gr.Button("Moderate", variant="primary") | |
| gr.Examples(["have a nice day!", "you are a worthless idiot", "i will find you and hurt you", | |
| "WIN a FREE iPhone, click http://bit.ly/x now"], inp) | |
| with gr.Column(): | |
| verdict = gr.Markdown() | |
| scores = gr.Label(label="Scores", num_top_classes=13) | |
| raw = gr.JSON(label="Response") | |
| btn.click(ui, inp, [verdict, scores, raw], api_name=False) | |
| inp.submit(ui, inp, [verdict, scores, raw], api_name=False) | |
| gr.api(moderate, api_name="moderate") | |
| gr.api(moderate_batch, api_name="moderate_batch") | |
| demo.queue(default_concurrency_limit=4) | |
| app = FastAPI(title="Moderation API") | |
| from fastapi.middleware.cors import CORSMiddleware | |
| app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"]) | |
| class Req(BaseModel): | |
| input: str | list[str] = Field(..., description="A text or a list of up to 32 texts") | |
| def rest(req: Req): | |
| texts = [req.input] if isinstance(req.input, str) else req.input | |
| return {"model": REPO, "results": _run(texts)} | |
| def labels(): | |
| return {"labels": LABELS, "thresholds": TH} | |
| app = gr.mount_gradio_app(app, demo, path="/") | |
| if __name__ == "__main__": | |
| import uvicorn | |
| uvicorn.run(app, host="0.0.0.0", port=7860) | |