hatespeech / train.py
Isa0's picture
Add onnxscript dependency and set dynamo=False for classic ONNX export
5097547
Raw History Blame Contribute Delete
18.5 kB
#!/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., &amp; -> &, &lt; -> <)
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()