Ak47-model-ml's picture
Upload 4 files
931e82f verified
Raw
History Blame Contribute Delete
5.88 kB
import os
import re
import html
import numpy as np
import torch
import torch.nn as nn
from transformers import BertTokenizerFast, BertForSequenceClassification
import gradio as gr
# ----------------------------------------------------
# Config
# ----------------------------------------------------
MODEL_NAME = "bert-base-uncased"
MAX_LEN = 128
SENTIMENT_MODEL_PATH = "sentiment_bert_best.pt"
SARCASM_MODEL_PATH = "sarcasm_bert_best.pt"
SARCASM_THRESHOLD = 0.6
CORRECTION_CONFIDENCE_THRESHOLD = 0.55
id2sentiment = {0: "Negative", 1: "Neutral", 2: "Positive"}
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
softmax = nn.Softmax(dim=1)
# ----------------------------------------------------
# Minimal cleaning for BERT
# ----------------------------------------------------
def clean_text(text: str) -> str:
if not isinstance(text, str):
return ""
text = html.unescape(text)
text = re.sub(r"<.*?>", " ", text)
text = re.sub(r"http\S+|www\.\S+", " ", text)
text = re.sub(r"\s+", " ", text).strip()
return text
# ----------------------------------------------------
# Load tokenizer & models
# ----------------------------------------------------
tokenizer = BertTokenizerFast.from_pretrained(MODEL_NAME)
sentiment_model = BertForSequenceClassification.from_pretrained(
MODEL_NAME, num_labels=3
)
sarcasm_model = BertForSequenceClassification.from_pretrained(
MODEL_NAME, num_labels=2
)
if os.path.exists(SENTIMENT_MODEL_PATH):
sentiment_model.load_state_dict(
torch.load(SENTIMENT_MODEL_PATH, map_location=device)
)
if os.path.exists(SARCASM_MODEL_PATH):
sarcasm_model.load_state_dict(
torch.load(SARCASM_MODEL_PATH, map_location=device)
)
sentiment_model.to(device).eval()
sarcasm_model.to(device).eval()
# ----------------------------------------------------
# Prediction Helpers
# ----------------------------------------------------
def get_probs(model, text):
encoding = tokenizer(
text,
max_length=MAX_LEN,
padding="max_length",
truncation=True,
return_tensors="pt",
)
input_ids = encoding["input_ids"].to(device)
attention_mask = encoding["attention_mask"].to(device)
with torch.no_grad():
outputs = model(input_ids=input_ids, attention_mask=attention_mask)
probs = softmax(outputs.logits).cpu().numpy()[0]
return probs
# ----------------------------------------------------
# Fusion Logic (Sarcasm-aware correction)
# ----------------------------------------------------
def fusion_polarity_correction(sarc_probs, sent_probs):
sarc_pred = int(np.argmax(sarc_probs))
sarc_conf = float(sarc_probs[sarc_pred])
sent_pred = int(np.argmax(sent_probs))
sent_conf = float(sent_probs[sent_pred])
final_pred = sent_pred
corrected = False
if sarc_pred == 1 and sarc_conf >= SARCASM_THRESHOLD:
if sent_conf >= CORRECTION_CONFIDENCE_THRESHOLD:
if sent_pred == 2:
final_pred = 0
corrected = True
elif sent_pred == 0:
final_pred = 2
corrected = True
elif sent_pred == 1:
final_pred = 0
corrected = True
return final_pred, corrected, sarc_conf, sent_conf
# ----------------------------------------------------
# Main Prediction Function
# ----------------------------------------------------
def predict_review(text: str):
if not isinstance(text, str) or not text.strip():
return "⚠ Please enter valid input text."
text_clean = clean_text(text)
sarc_probs = get_probs(sarcasm_model, text_clean)
sent_probs = get_probs(sentiment_model, text_clean)
final_pred, corrected, sarc_conf, sent_conf = fusion_polarity_correction(
sarc_probs, sent_probs
)
return {
"Input": text,
"Sarcasm Probability": round(float(sarc_probs[1]), 4),
"Original Sentiment": id2sentiment[int(np.argmax(sent_probs))],
"Original Confidence": round(sent_conf, 4),
"Final Sentiment": id2sentiment[final_pred],
"Correction Applied": corrected,
"Sentiment Probabilities": {
"Negative": round(float(sent_probs[0]), 4),
"Neutral": round(float(sent_probs[1]), 4),
"Positive": round(float(sent_probs[2]), 4),
},
}
# ----------------------------------------------------
# Gradio UI
# ----------------------------------------------------
def gradio_predict(text):
result = predict_review(text)
if isinstance(result, str):
return result
return (
f"Input: {result['Input']}\n\n"
f"Sarcasm Probability: {result['Sarcasm Probability']}\n\n"
f"Original Sentiment: {result['Original Sentiment']} "
f"(conf={result['Original Confidence']})\n\n"
f"Final Sentiment: {result['Final Sentiment']}\n\n"
f"Correction Applied: {result['Correction Applied']}\n\n"
f"Sentiment Probabilities:\n"
f" Negative: {result['Sentiment Probabilities']['Negative']}\n"
f" Neutral: {result['Sentiment Probabilities']['Neutral']}\n"
f" Positive: {result['Sentiment Probabilities']['Positive']}"
)
with gr.Blocks() as demo:
gr.Markdown("## 🔍 Sarcasm-Aware Sentiment Analysis (BERT Fusion Model)")
textbox = gr.Textbox(
lines=4,
placeholder="Yeah great, the battery died in 1 hour...",
label="Enter review text",
)
output = gr.Textbox(lines=12, label="Analysis Result")
btn = gr.Button("Analyze")
btn.click(gradio_predict, textbox, output)
if __name__ == "__main__":
demo.launch()