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)