#!/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()