Download hatespeech.py from Isa0/hatespeech: direct link, hf CLI and curl.
- Browser
- Download file 9.37 kB
-
https://huggingface.co/Isa0/hatespeech/resolve/main/hatespeech.py
- Command line
-
hf download hf://Isa0/hatespeech/hatespeech.py
-
curl -L -o hatespeech.py https://huggingface.co/Isa0/hatespeech/resolve/main/hatespeech.py
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() | |