UdaniSJ's picture
Add SHAP explanation tab
1f0470c verified
Raw
History Blame Contribute Delete
8.67 kB
import gradio as gr
import torch
import torch.nn.functional as F
import numpy as np
from transformers import BertForSequenceClassification, BertTokenizer
import transformers
from lime.lime_text import LimeTextExplainer
from textblob import TextBlob
import shap
import matplotlib.pyplot as plt
import matplotlib
matplotlib.use("Agg")
import io
from PIL import Image
MODEL_NAME = "UdaniSJ/hate-speech-severity-bert"
print("Loading model...")
tokenizer = BertTokenizer.from_pretrained(MODEL_NAME)
model = BertForSequenceClassification.from_pretrained(MODEL_NAME)
device = torch.device("cpu")
model.to(device)
model.eval()
print("Model loaded!")
class_names = ["Level 0 (Non-hate)", "Level 1 (Mild)", "Level 2 (Severe)"]
explainer_lime = LimeTextExplainer(class_names=class_names)
def bert_predict_proba(texts):
inputs = tokenizer(list(texts), truncation=True, padding=True, max_length=128, return_tensors="pt").to(device)
with torch.no_grad():
outputs = model(**inputs)
probs = F.softmax(outputs.logits, dim=1)
return probs.cpu().numpy()
def check_friendly_context(text, severity_score):
blob = TextBlob(text)
sentiment = blob.sentiment.polarity
flags = []
if severity_score > 0.5 and sentiment > 0.2:
flags.append("Positive sentiment detected")
affection_words = ["love","friend","bro","sis","mate","buddy","homie","fam","bestie","lol","haha","miss","care","heart","dawg","bruh","ily"]
if any(w in text.lower() for w in affection_words):
flags.append("Affection language detected")
if "?" in text and severity_score > 0.5:
flags.append("Questioning tone detected")
return flags, sentiment
def get_base_prediction(text, friendly_context):
probs = bert_predict_proba([text])[0]
weights = np.array([0.0, 0.5, 1.0])
severity_score = float(probs @ weights)
original_score = severity_score
auto_flags, sentiment = check_friendly_context(text, severity_score)
context_note = ""
if friendly_context:
severity_score = severity_score * 0.6
context_note = "Friendly context applied. Original: " + str(round(original_score,3)) + " Adjusted: " + str(round(severity_score,3))
elif auto_flags:
reduction = 0.15 * len(auto_flags)
severity_score = max(0, severity_score - reduction)
context_note = "Friendly signals detected. Adjusted: " + str(round(original_score,3)) + " to " + str(round(severity_score,3))
if severity_score < 0.35:
level_str = "Level 0 - Non-hate Speech"
elif severity_score < 0.65:
level_str = "Level 1 - Mild/Offensive"
else:
level_str = "Level 2 - Severe Hate Speech"
if severity_score >= 0.8:
decision = "AUTO-FLAG: High severity"
elif severity_score >= 0.5:
decision = "HUMAN REVIEW: Ambiguous"
else:
decision = "ALLOW: Low severity"
result = level_str + "\n\n"
result += "Severity Score: " + str(round(severity_score,3)) + " / 1.000\n"
result += "Sentiment: " + str(round(sentiment,2)) + "\n"
result += context_note + "\n\n"
result += "Decision: " + decision + "\n\n"
result += "Probabilities:\n"
result += " Non-hate: " + str(round(probs[0]*100,1)) + "%\n"
result += " Mild: " + str(round(probs[1]*100,1)) + "%\n"
result += " Severe: " + str(round(probs[2]*100,1)) + "%"
return result, severity_score, probs
def predict_with_lime(text, friendly_context):
if not text.strip():
return "Please enter some text.", 0.0, None
result, severity_score, probs = get_base_prediction(text, friendly_context)
try:
pred_level = int(np.argmax(probs))
exp = explainer_lime.explain_instance(text, bert_predict_proba, num_features=8, num_samples=300, labels=[pred_level])
word_weights = exp.as_list(label=pred_level)
words = [w[0] for w in word_weights]
scores = [w[1] for w in word_weights]
colors = ["#e74c3c" if s > 0 else "#2ecc71" for s in scores]
fig, ax = plt.subplots(figsize=(8, 4))
ax.barh(words, scores, color=colors)
ax.axvline(x=0, color="black", linewidth=0.8)
ax.set_xlabel("Word Importance (red=increases severity, green=decreases)")
ax.set_title("LIME Explanation - " + class_names[pred_level])
plt.tight_layout()
buf = io.BytesIO()
plt.savefig(buf, format="png", dpi=100, bbox_inches="tight")
buf.seek(0)
img = Image.open(buf)
plt.close()
except Exception as e:
print("LIME error: " + str(e))
img = None
return result, severity_score, img
def predict_with_shap(text, friendly_context):
if not text.strip():
return "Please enter some text.", 0.0, None
result, severity_score, probs = get_base_prediction(text, friendly_context)
try:
pred_level = int(np.argmax(probs))
bert_pipeline = transformers.pipeline(
"text-classification",
model=model,
tokenizer=tokenizer,
device=-1,
return_all_scores=True,
truncation=True,
max_length=128
)
explainer_shap = shap.Explainer(bert_pipeline)
shap_values = explainer_shap([text])
tokens = shap_values.data[0]
values = shap_values.values[0, :, pred_level]
valid = [(t, v) for t, v in zip(tokens, values) if t not in ["", "[PAD]", "[CLS]", "[SEP]"]]
if not valid:
return result, severity_score, None
pairs = sorted(valid, key=lambda x: abs(x[1]), reverse=True)[:10]
top_tokens = [p[0] for p in pairs]
top_values = [p[1] for p in pairs]
colors = ["#e74c3c" if v > 0 else "#2ecc71" for v in top_values]
fig, ax = plt.subplots(figsize=(8, 4))
ax.barh(top_tokens, top_values, color=colors)
ax.axvline(x=0, color="black", linewidth=0.8)
ax.set_xlabel("SHAP Value (red=increases severity, green=decreases)")
ax.set_title("SHAP Explanation - " + class_names[pred_level])
plt.tight_layout()
buf = io.BytesIO()
plt.savefig(buf, format="png", dpi=100, bbox_inches="tight")
buf.seek(0)
img = Image.open(buf)
plt.close()
except Exception as e:
print("SHAP error: " + str(e))
img = None
return result, severity_score, img
with gr.Blocks(title="Hate Speech Severity Predictor") as demo:
gr.Markdown("# Hate Speech Severity Predictor")
gr.Markdown("### Explainable AI - BERT + LIME + SHAP")
gr.Markdown("**MSc Research Project | University of Moratuwa**")
gr.Markdown("---")
with gr.Row():
with gr.Column(scale=2):
text_input = gr.Textbox(label="Enter Text", placeholder="Type here...", lines=4)
friendly_checkbox = gr.Checkbox(label="Friendly/Known context", value=False)
with gr.Row():
lime_btn = gr.Button("Analyse with LIME", variant="primary")
shap_btn = gr.Button("Analyse with SHAP", variant="secondary")
gr.Examples(
examples=[
["I love all people regardless of background", False],
["you idiot i cant believe you lol", False],
["you are my friend for life bro", True],
["those people control all the world banks", False],
["the duffers are too homophobic", False],
],
inputs=[text_input, friendly_checkbox],
label="Try these examples")
with gr.Column(scale=3):
result_output = gr.Textbox(label="Analysis Results", lines=14)
severity_slider = gr.Slider(minimum=0, maximum=1, label="Severity Score", interactive=False)
with gr.Tabs():
with gr.Tab("LIME Explanation"):
gr.Markdown("**LIME** - Fast word-level importance. Takes ~30 seconds.")
lime_plot = gr.Image(label="LIME Word Importance")
with gr.Tab("SHAP Explanation"):
gr.Markdown("**SHAP** - Theoretically grounded token importance. Takes ~2-3 minutes.")
shap_plot = gr.Image(label="SHAP Token Importance")
gr.Markdown("---")
gr.Markdown("Level 0 = Non-hate (Score < 0.35) | Level 1 = Mild (0.35-0.65) | Level 2 = Severe (> 0.65)")
gr.Markdown("Red bars = increases severity | Green bars = decreases severity")
lime_btn.click(fn=predict_with_lime, inputs=[text_input, friendly_checkbox], outputs=[result_output, severity_slider, lime_plot])
shap_btn.click(fn=predict_with_shap, inputs=[text_input, friendly_checkbox], outputs=[result_output, severity_slider, shap_plot])
demo.launch()