moderation-api / app.py
misukisu's picture
Moderation API
1048d56 verified
Raw History Blame Contribute Delete
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")
@app.post("/v1/moderate")
def rest(req: Req):
texts = [req.input] if isinstance(req.input, str) else req.input
return {"model": REPO, "results": _run(texts)}
@app.get("/v1/labels")
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)