hatespeech / hatespeech.py
Isa0's picture
Update README and sanitize test cases
2e3e2b8
Raw History Blame Contribute Delete
9.37 kB
#!/usr/bin/env python3
"""
hatespeech.py - ONNX Runtime Inference & Examples
Performs fast inference on text using exported ONNX (Raw or INT8 Quantized) models.
Key Features:
- Runs via ONNX Runtime (CPU or CUDA).
- Supports both raw `hatespeech.onnx` and quantized `hatespeech_int8.onnx`.
- Preprocesses input text and tokenizes inputs.
- Outputs predicted class, confidence scores across all classes, and latency in milliseconds.
- Provides interactive mode, CLI text argument, and built-in benchmark examples.
"""
import os
import sys
import re
import html
import time
import argparse
from typing import List, Dict, Union, Any
import numpy as np
import onnxruntime as ort
from transformers import AutoTokenizer
LABEL_NAMES = {
0: "Hate Speech",
1: "Offensive Language",
2: "Neither",
}
def clean_text(text: str) -> str:
"""Preprocesses input text matching the training pipeline."""
if not isinstance(text, str):
return ""
text = html.unescape(text)
text = re.sub(r"https?://\S+|www\.\S+", "", text)
text = re.sub(r"@\w+", "", text)
text = re.sub(r"\s+", " ", text).strip()
return text
def softmax(x: np.ndarray) -> np.ndarray:
"""Computes softmax over logits."""
e_x = np.exp(x - np.max(x, axis=-1, keepdims=True))
return e_x / e_x.sum(axis=-1, keepdims=True)
class HateSpeechDetector:
"""
Hate Speech & Offensive Language Detector powered by ONNX Runtime.
"""
def __init__(
self,
model_path: str = "./model/hatespeech_int8.onnx",
tokenizer_name_or_path: str = "./model",
fallback_model_name: str = "distilbert-base-uncased",
max_length: int = 128,
use_gpu: bool = False,
):
if not os.path.exists(model_path):
raise FileNotFoundError(
f"Model file not found at: '{model_path}'.\n"
f"Please train the model first by running `python train.py`."
)
print(f"Loading ONNX model from: {model_path}")
providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] if use_gpu else ["CPUExecutionProvider"]
# Configure ONNX runtime session
sess_options = ort.SessionOptions()
sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
self.session = ort.InferenceSession(model_path, sess_options=sess_options, providers=providers)
# Load tokenizer (first check local model dir, then fallback to HF hub)
if os.path.exists(os.path.join(tokenizer_name_or_path, "tokenizer_config.json")):
print(f"Loading tokenizer from local path: {tokenizer_name_or_path}")
self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name_or_path)
else:
print(f"Local tokenizer not found. Loading fallback: {fallback_model_name}")
self.tokenizer = AutoTokenizer.from_pretrained(fallback_model_name)
self.max_length = max_length
self.model_path = model_path
self.file_size_mb = os.path.getsize(model_path) / (1024 * 1024)
def predict(self, text: Union[str, List[str]]) -> List[Dict[str, Any]]:
"""
Runs inference on a single string or a list of strings.
Returns prediction details including class label, confidence, and latency.
"""
is_single = isinstance(text, str)
raw_texts = [text] if is_single else text
cleaned_texts = [clean_text(t) for t in raw_texts]
start_time = time.perf_counter()
# Tokenize
encoded = self.tokenizer(
cleaned_texts,
padding=True,
truncation=True,
max_length=self.max_length,
return_tensors="np",
)
ort_inputs = {
"input_ids": encoded["input_ids"].astype(np.int64),
"attention_mask": encoded["attention_mask"].astype(np.int64),
}
# Run ONNX inference
ort_outputs = self.session.run(None, ort_inputs)
logits = ort_outputs[0]
probs = softmax(logits)
elapsed_ms = (time.perf_counter() - start_time) * 1000
latency_per_sample = elapsed_ms / max(len(raw_texts), 1)
results = []
for i, original_text in enumerate(raw_texts):
pred_class_id = int(np.argmax(probs[i]))
class_probs = {
LABEL_NAMES[c]: float(probs[i][c]) for c in range(3)
}
results.append({
"text": original_text,
"cleaned_text": cleaned_texts[i],
"label_id": pred_class_id,
"label_name": LABEL_NAMES[pred_class_id],
"confidence": float(probs[i][pred_class_id]),
"probabilities": class_probs,
"latency_ms": latency_per_sample,
})
return results
def run_examples(detector: HateSpeechDetector):
"""Runs a series of representative test cases and prints formatted output."""
examples = [
"I really love this community, everyone is so supportive and kind!",
"Can you please send me the report by tomorrow morning?",
"Shut up, you are being so damn annoying tonight.",
"That person is crazy, get the hell out of here.",
"Those people are subhuman and don't belong in our country, kick them all out.",
]
print("\n" + "=" * 75)
print(" RUNNING BENCHMARK EXAMPLES")
print(f" Model: {detector.model_path} ({detector.file_size_mb:.2f} MB)")
print("=" * 75)
results = detector.predict(examples)
for i, res in enumerate(results, 1):
print(f"\n[Example {i}]")
print(f" Input : \"{res['text']}\"")
print(f" Prediction : {res['label_name']} (Class {res['label_id']})")
print(f" Confidence : {res['confidence'] * 100:.2f}%")
print(f" Latency : {res['latency_ms']:.2f} ms")
print(" Breakdown :")
for label, prob in res["probabilities"].items():
bar = "#" * int(prob * 20)
print(f" - {label:<20}: {prob * 100:>5.1f}% | {bar}")
print("\n" + "=" * 75)
def main():
parser = argparse.ArgumentParser(description="Hate Speech Detection with ONNX Runtime.")
parser.add_argument(
"--model",
type=str,
default=None,
help="Path to ONNX model file. If not set, automatically prefers INT8 or raw ONNX in ./model",
)
parser.add_argument(
"--raw",
action="store_true",
help="Use raw ONNX model (./model/hatespeech.onnx) instead of INT8 quantized model",
)
parser.add_argument(
"--text",
type=str,
default=None,
help="Specific text to classify via CLI",
)
parser.add_argument(
"--interactive",
action="store_true",
help="Start interactive prompt loop to test custom sentences",
)
parser.add_argument(
"--gpu",
action="store_true",
help="Use CUDA Execution Provider if available",
)
args = parser.parse_args()
# Determine model path
if args.model:
model_path = args.model
elif args.raw:
model_path = "./model/hatespeech.onnx"
else:
# Default to INT8 if present, else fallback to raw
if os.path.exists("./model/hatespeech_int8.onnx"):
model_path = "./model/hatespeech_int8.onnx"
else:
model_path = "./model/hatespeech.onnx"
try:
detector = HateSpeechDetector(
model_path=model_path,
tokenizer_name_or_path="./model",
use_gpu=args.gpu,
)
except FileNotFoundError as e:
print(f"\nError: {e}")
sys.exit(1)
# 1. Single text prediction via CLI
if args.text:
res = detector.predict(args.text)[0]
print("\n" + "=" * 50)
print(f"Text : \"{res['text']}\"")
print(f"Prediction : {res['label_name']} (Class {res['label_id']})")
print(f"Confidence : {res['confidence'] * 100:.2f}%")
print(f"Latency : {res['latency_ms']:.2f} ms")
print("Probabilities:")
for label, prob in res["probabilities"].items():
print(f" {label:<20}: {prob * 100:.1f}%")
print("=" * 50)
return
# 2. Interactive prompt loop
if args.interactive:
print("\n" + "=" * 60)
print(" INTERACTIVE HATE SPEECH DETECTION (Type 'exit' to quit)")
print(f" Model: {model_path}")
print("=" * 60)
while True:
try:
user_input = input("\nEnter text: ").strip()
if not user_input:
continue
if user_input.lower() in ("exit", "quit", "q"):
print("Exiting.")
break
res = detector.predict(user_input)[0]
print(f" -> Result : {res['label_name']} ({res['confidence'] * 100:.1f}%)")
print(f" -> Latency : {res['latency_ms']:.2f} ms")
print(f" -> Breakdown : " + ", ".join([f"{k}: {v*100:.1f}%" for k, v in res['probabilities'].items()]))
except (KeyboardInterrupt, EOFError):
print("\nExiting.")
break
return
# 3. Default: Run benchmark examples
run_examples(detector)
if __name__ == "__main__":
main()