#!/usr/bin/env python3 """ train.py - Hate Speech & Offensive Language Classifier Trains a PyTorch transformer model on Davidson et al. labeled_data.csv. Exports both raw ONNX and INT8 quantized (8q) ONNX models. Key Features: - Downloads dataset to a temporary directory and cleans it up after loading. - Preprocessing: unescapes HTML, strips URLs and Twitter mentions. - Prevents Overfitting & Underfitting: * Stratified 80/10/10 train/val/test split. * Balanced class weights for CrossEntropyLoss (combats ~5% hate speech imbalance). * Weight decay (AdamW) + Linear learning rate warmup & decay. * Early stopping tracking validation Macro F1 score. - Epoch Optimization: * Sequence length truncated to 128 (optimal for tweets, fast epochs). * PyTorch Automatic Mixed Precision (AMP / FP16) on CUDA. - Model Exports: * Raw ONNX (FP32) * Quantized INT8 ONNX (dynamic quantization for ~75% size reduction & fast inference) * Tokenizer files saved alongside ONNX models for standalone portability. """ from __future__ import annotations import os import sys import re import html import time import shutil import tempfile import urllib.request import argparse from typing import Tuple, Dict, Any, List import numpy as np import pandas as pd import torch import torch.nn as nn import torch.optim as optim from torch.optim import AdamW from torch.utils.data import Dataset, DataLoader from transformers import ( AutoConfig, AutoTokenizer, AutoModelForSequenceClassification, get_linear_schedule_with_warmup, ) from sklearn.model_selection import train_test_split from sklearn.metrics import classification_report, f1_score, accuracy_score import onnx import onnxruntime from onnxruntime.quantization import quantize_dynamic, QuantType # Label mappings LABEL_NAMES = { 0: "Hate Speech", 1: "Offensive Language", 2: "Neither", } DATASET_URL = ( "https://raw.githubusercontent.com/t-davidson/hate-speech-and-offensive-language/" "master/data/labeled_data.csv" ) def clean_text(text: str) -> str: """Cleans tweet text by unescaping HTML, normalizing handles, and stripping URLs.""" if not isinstance(text, str): return "" # Decode HTML entities (e.g., & -> &, < -> <) text = html.unescape(text) # Remove URLs text = re.sub(r"https?://\S+|www\.\S+", "", text) # Remove user mentions text = re.sub(r"@\w+", "", text) # Normalize excessive whitespaces text = re.sub(r"\s+", " ", text).strip() return text def download_dataset_temp(url: str = DATASET_URL) -> pd.DataFrame: """ Downloads labeled_data.csv into a temporary directory, loads it into pandas, and cleans up the temporary directory immediately after. """ print(f"\n[1/6] Fetching dataset from: {url}") with tempfile.TemporaryDirectory() as temp_dir: temp_file_path = os.path.join(temp_dir, "labeled_data.csv") print(f" -> Temporary download location: {temp_file_path}") req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0"}) with urllib.request.urlopen(req) as response, open(temp_file_path, "wb") as out_file: shutil.copyfileobj(response, out_file) file_size_mb = os.path.getsize(temp_file_path) / (1024 * 1024) print(f" -> Download completed successfully ({file_size_mb:.2f} MB).") df = pd.read_csv(temp_file_path) # Required columns: 'class' (0, 1, 2) and 'tweet' df = df[["class", "tweet"]].dropna() df["tweet"] = df["tweet"].apply(clean_text) df = df[df["tweet"].str.strip() != ""].reset_index(drop=True) print(f" -> Loaded {len(df):,} valid samples.") class_counts = df["class"].value_counts().sort_index() for cls_id, count in class_counts.items(): pct = count / len(df) * 100 print(f" Class {cls_id} ({LABEL_NAMES.get(cls_id, 'Unknown')}): {count:,} ({pct:.1f}%)") print(" -> Temporary directory and file automatically cleaned up.") return df class HateSpeechDataset(Dataset): """PyTorch Dataset for tokenized hate speech data.""" def __init__(self, texts, labels, tokenizer, max_length: int = 128): self.encodings = tokenizer( texts, truncation=True, padding=True, max_length=max_length, return_tensors="pt", ) self.labels = torch.tensor(labels, dtype=torch.long) def __len__(self): return len(self.labels) def __getitem__(self, idx): item = {key: val[idx] for key, val in self.encodings.items()} item["labels"] = self.labels[idx] return item def compute_class_weights(labels: np.ndarray, num_classes: int = 3) -> torch.Tensor: """ Computes balanced class weights inversely proportional to class frequencies. Crucial to prevent under-learning the minority Hate Speech class (~5.7%). """ counts = np.bincount(labels, minlength=num_classes) total_samples = len(labels) # Balanced weighting: total / (num_classes * count) weights = total_samples / (num_classes * counts.astype(np.float32)) # Normalize so mean weight is 1.0 to preserve learning rate stability weights = weights / np.mean(weights) return torch.tensor(weights, dtype=torch.float) def train_epoch( model: nn.Module, dataloader: DataLoader, optimizer: optim.Optimizer, scheduler: Any, scaler: Any, criterion: nn.Module, device: torch.device, ) -> float: """Trains the model for 1 epoch with optional AMP mixed-precision.""" model.train() total_loss = 0.0 use_amp = (device.type == "cuda") for batch in dataloader: optimizer.zero_grad() input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) labels = batch["labels"].to(device) if use_amp: try: autocast_ctx = torch.amp.autocast(device_type="cuda", dtype=torch.float16) except (AttributeError, TypeError): autocast_ctx = torch.cuda.amp.autocast() with autocast_ctx: outputs = model(input_ids=input_ids, attention_mask=attention_mask) loss = criterion(outputs.logits, labels) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update() else: outputs = model(input_ids=input_ids, attention_mask=attention_mask) loss = criterion(outputs.logits, labels) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() scheduler.step() total_loss += loss.item() return total_loss / len(dataloader) def evaluate( model: nn.Module, dataloader: DataLoader, criterion: nn.Module, device: torch.device, ) -> Tuple[float, float, float, np.ndarray, np.ndarray]: """Evaluates the model and computes validation loss, accuracy, and Macro F1.""" model.eval() total_loss = 0.0 all_preds = [] all_labels = [] with torch.no_grad(): for batch in dataloader: input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) labels = batch["labels"].to(device) outputs = model(input_ids=input_ids, attention_mask=attention_mask) loss = criterion(outputs.logits, labels) total_loss += loss.item() preds = torch.argmax(outputs.logits, dim=1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.cpu().numpy()) avg_loss = total_loss / len(dataloader) all_preds = np.array(all_preds) all_labels = np.array(all_labels) acc = accuracy_score(all_labels, all_preds) macro_f1 = f1_score(all_labels, all_preds, average="macro") return avg_loss, acc, macro_f1, all_preds, all_labels def export_to_onnx( model: nn.Module, tokenizer: Any, output_dir: str, device: torch.device, max_length: int = 128, ) -> Tuple[str, str]: """ Exports PyTorch model to: 1. Raw ONNX format (`hatespeech.onnx`) with dynamic batch and sequence length axes. 2. Quantized INT8 ONNX (`hatespeech_int8.onnx`) for CPU/edge speed and small file size. """ os.makedirs(output_dir, exist_ok=True) raw_onnx_path = os.path.join(output_dir, "hatespeech.onnx") int8_onnx_path = os.path.join(output_dir, "hatespeech_int8.onnx") model.eval() model.to("cpu") # Export on CPU for universal ONNX portability print("\n[5/6] Exporting to Raw ONNX...") dummy_text = "Sample sentence for hate speech detection ONNX export." dummy_encodings = tokenizer( dummy_text, max_length=max_length, padding="max_length", truncation=True, return_tensors="pt", ) dummy_inputs = ( dummy_encodings["input_ids"], dummy_encodings["attention_mask"], ) export_kwargs = { "input_names": ["input_ids", "attention_mask"], "output_names": ["logits"], "dynamic_axes": { "input_ids": {0: "batch_size", 1: "sequence_length"}, "attention_mask": {0: "batch_size", 1: "sequence_length"}, "logits": {0: "batch_size"}, }, "opset_version": 14, "do_constant_folding": True, } import inspect sig = inspect.signature(torch.onnx.export) if "dynamo" in sig.parameters: export_kwargs["dynamo"] = False torch.onnx.export( model, dummy_inputs, raw_onnx_path, **export_kwargs, ) # Verify raw ONNX model validity onnx_model = onnx.load(raw_onnx_path) onnx.checker.check_model(onnx_model) raw_size_mb = os.path.getsize(raw_onnx_path) / (1024 * 1024) print(f" -> Raw ONNX saved to: {raw_onnx_path} ({raw_size_mb:.2f} MB)") print("\n[6/6] Quantizing ONNX model to INT8 (8q dynamic)...") quantize_dynamic( model_input=raw_onnx_path, model_output=int8_onnx_path, weight_type=QuantType.QInt8, ) int8_size_mb = os.path.getsize(int8_onnx_path) / (1024 * 1024) compression = (1.0 - (int8_size_mb / raw_size_mb)) * 100 print(f" -> INT8 ONNX saved to: {int8_onnx_path} ({int8_size_mb:.2f} MB, {compression:.1f}% smaller)") # Save tokenizer files alongside models for offline independence tokenizer.save_pretrained(output_dir) print(f" -> Tokenizer assets saved in: {output_dir}") return raw_onnx_path, int8_onnx_path def main(): parser = argparse.ArgumentParser(description="Train Hate Speech Classifier and Export to ONNX & INT8.") parser.add_argument("--model_name", type=str, default="distilbert-base-uncased", help="Pretrained Hugging Face model architecture (default: distilbert-base-uncased)") parser.add_argument("--epochs", type=int, default=3, help="Maximum training epochs (default: 3, prevents overfitting)") parser.add_argument("--batch_size", type=int, default=32, help="Training batch size (default: 32)") parser.add_argument("--lr", type=float, default=2e-5, help="Learning rate for AdamW (default: 2e-5)") parser.add_argument("--max_length", type=int, default=128, help="Max sequence token length for tweets (default: 128 for fast epoch time)") parser.add_argument("--patience", type=int, default=2, help="Early stopping patience epochs based on Macro F1 (default: 2)") parser.add_argument("--output_dir", type=str, default="./model", help="Directory to save ONNX models and tokenizer (default: ./model)") args = parser.parse_args() # Hardware detection device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print("=" * 65) print(" HATE SPEECH & OFFENSIVE LANGUAGE CLASSIFIER TRAINING") print("=" * 65) print(f"Using Compute Device : {device}") if device.type == "cuda": print(f"GPU Model : {torch.cuda.get_device_name(0)}") print(f"Automatic Mixed Prec : Enabled (FP16)") print(f"Base Architecture : {args.model_name}") print(f"Max Sequence Length : {args.max_length} tokens") print(f"Batch Size / Max Epochs: {args.batch_size} / {args.epochs}") print("=" * 65) # 1. Download & Clean Dataset df = download_dataset_temp(DATASET_URL) texts = df["tweet"].tolist() labels = df["class"].to_numpy() # 2. Stratified Data Split (80% Train, 10% Validation, 10% Test) print("\n[2/6] Preparing stratified data splits (80% Train, 10% Val, 10% Test)...") train_texts, temp_texts, train_labels, temp_labels = train_test_split( texts, labels, test_size=0.20, random_state=42, stratify=labels ) val_texts, test_texts, val_labels, test_labels = train_test_split( temp_texts, temp_labels, test_size=0.50, random_state=42, stratify=temp_labels ) print(f" -> Train split : {len(train_labels):,} samples") print(f" -> Val split : {len(val_labels):,} samples") print(f" -> Test split : {len(test_labels):,} samples") # 3. Tokenizer & DataLoaders print(f"\n[3/6] Loading tokenizer '{args.model_name}' and tokenizing...") tokenizer = AutoTokenizer.from_pretrained(args.model_name) train_dataset = HateSpeechDataset(train_texts, train_labels, tokenizer, max_length=args.max_length) val_dataset = HateSpeechDataset(val_texts, val_labels, tokenizer, max_length=args.max_length) test_dataset = HateSpeechDataset(test_texts, test_labels, tokenizer, max_length=args.max_length) train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=args.batch_size * 2, shuffle=False) test_loader = DataLoader(test_dataset, batch_size=args.batch_size * 2, shuffle=False) # 4. Model Setup & Anti-Overfitting Controls print("\n[4/6] Initializing model & regularization controls...") config = AutoConfig.from_pretrained(args.model_name, num_labels=3) if hasattr(config, "seq_classif_dropout"): config.seq_classif_dropout = 0.2 elif hasattr(config, "classifier_dropout"): config.classifier_dropout = 0.2 model = AutoModelForSequenceClassification.from_pretrained( args.model_name, config=config, ) model.to(device) # Balanced class weights to counteract the severe minority of class 0 class_weights = compute_class_weights(train_labels, num_classes=3).to(device) print(f" -> Computed balanced class weights: {class_weights.cpu().numpy().round(3)}") criterion = nn.CrossEntropyLoss(weight=class_weights) # Optimizer with weight decay to block overfitting optimizer = AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) # Linear warmup scheduler total_steps = len(train_loader) * args.epochs warmup_steps = int(0.1 * total_steps) scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_steps ) try: scaler = torch.amp.GradScaler("cuda", enabled=(device.type == "cuda")) except (AttributeError, TypeError): scaler = torch.cuda.amp.GradScaler(enabled=(device.type == "cuda")) # Training Loop with Early Stopping monitoring Macro F1 best_val_macro_f1 = -1.0 best_model_state = None patience_counter = 0 print("\n" + "-" * 75) print(f"{'Epoch':<8}{'Train Loss':<14}{'Val Loss':<12}{'Val Acc':<12}{'Val Macro F1':<16}{'Time (s)':<10}") print("-" * 75) for epoch in range(1, args.epochs + 1): start_time = time.time() train_loss = train_epoch( model, train_loader, optimizer, scheduler, scaler, criterion, device ) val_loss, val_acc, val_f1, _, _ = evaluate(model, val_loader, criterion, device) elapsed = time.time() - start_time print(f"{epoch:<8}{train_loss:<14.4f}{val_loss:<12.4f}{val_acc*100:<12.2f}%{val_f1:<16.4f}{elapsed:<10.1f}") # Check for best model based on Macro F1 (prevents skew towards majority class) if val_f1 > best_val_macro_f1: best_val_macro_f1 = val_f1 best_model_state = {k: v.cpu().clone() for k, v in model.state_dict().items()} patience_counter = 0 improved_mark = " (Best Saved)" else: patience_counter += 1 improved_mark = f" (Patience {patience_counter}/{args.patience})" print(f" Status: Macro F1 = {val_f1:.4f}{improved_mark}") if patience_counter >= args.patience: print(f"\n[!] Early stopping triggered at epoch {epoch} to prevent overfitting.") break print("-" * 75) # Load best checkpoint if best_model_state is not None: model.load_state_dict(best_model_state) print(f"Loaded best checkpoint with Validation Macro F1: {best_val_macro_f1:.4f}") # Final Evaluation on Independent Test Set print("\n" + "=" * 65) print(" FINAL EVALUATION ON TEST SET (10% HOLDOUT)") print("=" * 65) _, test_acc, test_macro_f1, test_preds, test_targets = evaluate( model, test_loader, criterion, device ) print(f"Overall Test Accuracy: {test_acc * 100:.2f}%") print(f"Test Macro F1 Score : {test_macro_f1:.4f}\n") print("Detailed Classification Report:") target_names = [f"{i}: {LABEL_NAMES[i]}" for i in range(3)] print(classification_report(test_targets, test_preds, target_names=target_names, digits=4)) # 5 & 6. Export to Raw ONNX & 8q Quantized ONNX export_to_onnx( model=model, tokenizer=tokenizer, output_dir=args.output_dir, device=device, max_length=args.max_length, ) print("\n" + "=" * 65) print(" TRAINING COMPLETE!") print(f"Saved artifacts inside '{args.output_dir}':") print(f" 1. Raw ONNX Model : {os.path.join(args.output_dir, 'hatespeech.onnx')}") print(f" 2. Quantized INT8 ONNX: {os.path.join(args.output_dir, 'hatespeech_int8.onnx')}") print(f" 3. Tokenizer Assets : {args.output_dir}/") print("=" * 65) if __name__ == "__main__": main()