ai-course / evaluate.py
Soscolni's picture
Initial commit: hate speech & conspiracy detection project
83eae8f
Raw
History Blame Contribute Delete
7.67 kB
"""Full evaluation across all test sets with generalization analysis."""
import os
import numpy as np
import pandas as pd
import torch
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.metrics import confusion_matrix, roc_curve, auc
from transformers import AutoModelForSequenceClassification, AutoTokenizer
from peft import PeftModel, PeftConfig
from config import CONFIG
from data.dataset import ConspiracyDataset
from model.utils import print_evaluation_report
from generalization_tests import GENERALIZATION_EXAMPLES
def load_trained_model(model_path=None):
"""Load the best trained model and tokenizer. Supports LoRA adapters and full models."""
if model_path is None:
model_path = os.path.join(CONFIG["output_dir"], "best_model")
tokenizer = AutoTokenizer.from_pretrained(model_path)
if os.path.exists(os.path.join(model_path, "adapter_config.json")):
peft_config = PeftConfig.from_pretrained(model_path)
base_model = AutoModelForSequenceClassification.from_pretrained(
peft_config.base_model_name_or_path,
num_labels=CONFIG["num_labels"],
)
model = PeftModel.from_pretrained(base_model, model_path)
model = model.merge_and_unload()
print(f" LoRA adapter loaded and merged from {model_path}")
else:
model = AutoModelForSequenceClassification.from_pretrained(model_path)
print(f" Full model loaded from {model_path}")
model.eval()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)
return model, tokenizer, device
THRESHOLD = 0.5
def predict(model, tokenizer, texts, device, batch_size=16):
"""Run inference on a list of texts."""
all_preds = []
all_probs = []
for i in range(0, len(texts), batch_size):
batch_texts = texts[i:i + batch_size]
encoding = tokenizer(
batch_texts,
max_length=CONFIG["max_length"],
padding=True,
truncation=True,
return_tensors="pt",
).to(device)
with torch.no_grad():
outputs = model(**encoding)
logits = outputs.logits
probs = torch.softmax(logits, dim=-1)
conspiracy_probs = probs[:, 1].cpu().numpy()
preds = (conspiracy_probs >= THRESHOLD).astype(int)
all_preds.extend(preds)
all_probs.extend(conspiracy_probs)
return np.array(all_preds), np.array(all_probs)
def plot_confusion_matrix(labels, preds, title, save_path):
"""Plot and save a confusion matrix."""
cm = confusion_matrix(labels, preds)
plt.figure(figsize=(6, 5))
sns.heatmap(
cm, annot=True, fmt="d", cmap="Blues",
xticklabels=["non-conspiracy", "conspiracy"],
yticklabels=["non-conspiracy", "conspiracy"],
)
plt.title(title)
plt.xlabel("Predicted")
plt.ylabel("Actual")
plt.tight_layout()
plt.savefig(save_path, dpi=150)
plt.close()
print(f" Saved: {save_path}")
def plot_roc_curve(labels, probs, title, save_path):
"""Plot and save an ROC curve."""
# Guard against NaN/Inf probabilities (can happen with broken training)
probs = np.array(probs, dtype=np.float64)
if np.any(np.isnan(probs)) or np.any(np.isinf(probs)):
print(f" Warning: probabilities contain NaN/Inf — skipping ROC curve for '{title}'")
return
# Need both classes present
if len(set(labels)) < 2:
print(f" Warning: only one class in labels — skipping ROC curve for '{title}'")
return
fpr, tpr, _ = roc_curve(labels, probs)
roc_auc = auc(fpr, tpr)
plt.figure(figsize=(6, 5))
plt.plot(fpr, tpr, color="darkorange", lw=2, label=f"AUC = {roc_auc:.3f}")
plt.plot([0, 1], [0, 1], color="navy", lw=1, linestyle="--")
plt.xlabel("False Positive Rate")
plt.ylabel("True Positive Rate")
plt.title(title)
plt.legend(loc="lower right")
plt.tight_layout()
plt.savefig(save_path, dpi=150)
plt.close()
print(f" Saved: {save_path}")
def evaluate_test_set(model, tokenizer, device, df, title):
"""Evaluate on a DataFrame test set."""
texts = df["text"].tolist()
labels = df["label"].values
preds, probs = predict(model, tokenizer, texts, device)
print_evaluation_report(labels, preds, probs, title=title)
return labels, preds, probs
def evaluate_generalization(model, tokenizer, device):
"""Evaluate on hand-crafted generalization examples."""
texts = [t for t, _ in GENERALIZATION_EXAMPLES]
labels = np.array([l for _, l in GENERALIZATION_EXAMPLES])
preds, probs = predict(model, tokenizer, texts, device)
print_evaluation_report(labels, preds, probs, title="Generalization Test (Non-COVID)")
# Print individual predictions for inspection
print("Individual predictions:")
print(f" {'Label':>5} {'Pred':>5} {'Prob':>6} Text")
print(f" {'-'*5:>5} {'-'*5:>5} {'-'*6:>6} {'-'*50}")
for i, (text, label) in enumerate(GENERALIZATION_EXAMPLES):
marker = "X" if preds[i] != label else " "
print(f" {label:>5} {preds[i]:>5} {probs[i]:>6.3f} {marker} {text[:80]}...")
return labels, preds, probs
def main():
print("Loading trained model...")
model, tokenizer, device = load_trained_model()
plots_dir = os.path.join(CONFIG["output_dir"], "plots")
os.makedirs(plots_dir, exist_ok=True)
# 1. Mixed test set
test_path = os.path.join(CONFIG["output_dir"], "test_split.csv")
if os.path.exists(test_path):
test_df = pd.read_csv(test_path)
labels_mixed, preds_mixed, probs_mixed = evaluate_test_set(
model, tokenizer, device, test_df, "Mixed Test Set",
)
plot_confusion_matrix(
labels_mixed, preds_mixed, "Mixed Test Set",
os.path.join(plots_dir, "cm_mixed.png"),
)
plot_roc_curve(
labels_mixed, probs_mixed, "ROC - Mixed Test Set",
os.path.join(plots_dir, "roc_mixed.png"),
)
else:
print(f"Warning: {test_path} not found. Run train.py first.")
# 2. COVID-only test set
covid_path = os.path.join(CONFIG["output_dir"], "covid_test_split.csv")
if os.path.exists(covid_path):
covid_df = pd.read_csv(covid_path)
labels_covid, preds_covid, probs_covid = evaluate_test_set(
model, tokenizer, device, covid_df, "COVID-Only Test Set",
)
plot_confusion_matrix(
labels_covid, preds_covid, "COVID-Only Test Set",
os.path.join(plots_dir, "cm_covid.png"),
)
else:
print(f"Warning: {covid_path} not found. Run train.py first.")
# 3. Generalization test
labels_gen, preds_gen, probs_gen = evaluate_generalization(model, tokenizer, device)
plot_confusion_matrix(
labels_gen, preds_gen, "Generalization Test (Non-COVID)",
os.path.join(plots_dir, "cm_generalization.png"),
)
# 4. Generalization gap analysis
print("\n" + "=" * 60)
print(" Generalization Gap Analysis")
print("=" * 60)
if os.path.exists(covid_path):
from sklearn.metrics import f1_score
f1_covid = f1_score(labels_covid, preds_covid)
f1_gen = f1_score(labels_gen, preds_gen)
gap = f1_covid - f1_gen
print(f" COVID test F1: {f1_covid:.4f}")
print(f" Generalization test F1: {f1_gen:.4f}")
print(f" Generalization gap: {gap:.4f} {'(good!)' if abs(gap) < 0.1 else '(needs improvement)'}")
print("\nEvaluation complete! Check plots in:", plots_dir)
if __name__ == "__main__":
main()