leygit/ITI110_Spam_Classification_Project
0
1# DISTILLBERT RUN 3 , added weight_decay=0.012import pandas as pd3import torch4import torch.nn as nn5import torch.optim as optim6import torch.nn.functional as F7from torch.utils.data import Dataset, DataLoader8from transformers import DistilBertTokenizer, DistilBertForSequenceClassification9from sklearn.model_selection import train_test_split10from sklearn.feature_extraction.text import CountVectorizer # Converts text into a matrix of token counts11from sklearn.metrics import classification_report, accuracy_score12import gradio as gr13 14# Load dataset15file_path = 'spam_ham_dataset.csv'16df = pd.read_csv(file_path)17 18# Convert label column to numeric (0 for ham, 1 for spam)19df['label_num'] = df['label'].astype('category').cat.codes20 21# Define device22device = torch.device("cuda" if torch.cuda.is_available() else "cpu")23 24# Load tokenizer25tokenizer = DistilBertTokenizer.from_pretrained("distilbert-base-uncased")26 27# Tokenize dataset28encodings = tokenizer(df['text'].tolist(), padding=True, truncation=True, max_length=128, return_tensors="pt")29labels = torch.tensor(df['label_num'].values)30 31# Custom Dataset32class SpamDataset(Dataset):33 def __init__(self, encodings, labels):34 self.encodings = encodings35 self.labels = labels36 37 def __len__(self):38 return len(self.labels)39 40 def __getitem__(self, idx):41 item = {key: val[idx] for key, val in self.encodings.items()} # Keep as PyTorch tensors42 item['labels'] = torch.tensor(self.labels[idx], dtype=torch.long) # Ensure labels are `long`43 return item44 45# Create dataset46dataset = SpamDataset(encodings, labels)47 48# Split dataset (80% train, 20% validation)49train_size = int(0.8 * len(dataset))50val_size = len(dataset) - train_size51train_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size])52 53def get_top_words(corpus, n=None):54 vec = CountVectorizer(stop_words='english').fit(corpus)55 bag_of_words = vec.transform(corpus)56 sum_words = bag_of_words.sum(axis=0)57 words_freq = [(word, sum_words[0, idx]) for word, idx in vec.vocabulary_.items()]58 words_freq = sorted(words_freq, key=lambda x: x[1], reverse=True)59 return words_freq[:n]60 61# DataLoader Function (Fix Collate)62def collate_fn(batch):63 keys = batch[0].keys()64 collated = {key: torch.stack([b[key] for b in batch]) for key in keys}65 return collated66 67# Create DataLoader68train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True, collate_fn=collate_fn)69val_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, collate_fn=collate_fn)70 71# Load the trained model72def load_model(model_path="distilbert_spam_model.pt"):73 model = DistilBertForSequenceClassification.from_pretrained("distilbert-base-uncased", num_labels=2)74 model.load_state_dict(torch.load(model_path, map_location=device)) # Load model weights75 model.to(device)76 model.eval() # Set model to evaluation mode77 return model78 79# Load model globally80model = load_model()81 82# Classification function83def classify_email(email_text):84 model.eval()85 86 with torch.no_grad():87 inputs = tokenizer(email_text, padding=True, truncation=True, max_length=256, return_tensors="pt")88 inputs = {key: val.to(device) for key, val in inputs.items()}89 outputs = model(**inputs)90 logits = outputs.logits91 predictions = torch.argmax(logits, dim=1)92 probs = F.softmax(logits, dim=1)93 confidence = torch.max(probs).item() * 10094 95 result = "Spam" if predictions.item() == 1 else "Ham"96 return result, f"{confidence:.2f}%"97 98# Evaluation function with detailed classification report99def evaluate_model_with_report(val_loader):100 model.eval() # Set model to evaluation mode101 y_true = []102 y_pred = []103 correct = 0104 total = 0105 106 with torch.no_grad():107 for batch in val_loader:108 inputs = {key: val.to(device) for key, val in batch.items()}109 labels = inputs.pop("labels").to(device)110 111 outputs = model(**inputs)112 predictions = torch.argmax(outputs.logits, dim=1)113 114 # Collect labels and predictions115 y_true.extend(labels.cpu().numpy())116 y_pred.extend(predictions.cpu().numpy())117 118 # Calculate accuracy119 correct += (predictions == labels).sum().item()120 total += labels.size(0)121 122 # Calculate accuracy123 accuracy = correct / total if total > 0 else 0124 print(f"Validation Accuracy: {accuracy:.4f}")125 126 # Print classification report127 print("\nClassification Report:")128 print(classification_report(y_true, y_pred, target_names=["Ham", "Spam"]))129 130 return accuracy131 132# Performance metrics133def generate_performance_metrics():134 model.eval() # Set model to evaluation mode135 136 y_true = [] # True labels137 y_pred = [] # Predicted labels138 139 with torch.no_grad():140 for batch in val_loader:141 inputs = {key: val.to(device) for key, val in batch.items()}142 labels = inputs.pop("labels").to(device) # Extract labels143 144 outputs = model(**inputs)145 predictions = torch.argmax(outputs.logits, dim=1)146 147 y_true.extend(labels.cpu().numpy())148 y_pred.extend(predictions.cpu().numpy())149 150 # Compute accuracy and classification report151 accuracy = accuracy_score(y_true, y_pred)152 report = classification_report(y_true, y_pred, output_dict=True)153 154 return {155 "accuracy": f"{accuracy:.2%}",156 "precision": f"{report['1']['precision']:.2%}",157 "recall": f"{report['1']['recall']:.2%}",158 "f1_score": f"{report['1']['f1-score']:.2%}",159 }160 161 162 163# Gradio Interface164 165def create_interface():166 performance_metrics = generate_performance_metrics()167 with gr.Blocks() as interface:168 with gr.Tab(" π¨ Demo"):169 gr.Markdown(" # π§π Spam and Phishing Email Detection")170 gr.Markdown(171 """172 Welcome to the Spam and Phishing Email Detection Demo! This tool leverages DistilBERT, a lightweight yet powerful transformer model, to classify emails as ham (legitimate), spam, or phishing based on their content.173 This project aims to enhance email security by identifying malicious messages with high accuracy, reducing the risk of scams and fraud. Feel free to explore the demo and see how AI can provide a safer environment for everyone.174 """) 175 176 177 # Email Text Input178 email_input = gr.Textbox(179 lines=8, placeholder="Type or paste your email content here...", label="Email Content"180 )181 182 # Email Text Results and Analysis183 result_output = gr.Textbox(label="Classification Result")184 confidence_output = gr.Textbox(label="Confidence Score", interactive=False)185 186 analyze_button = gr.Button("Analyze Email")187 188 def email_analysis_pipeline(email_text):189 results = classify_email(email_text)190 return (191 results["result"],192 results["confidence"]193 )194 195 analyze_button.click(196 fn=classify_email,197 inputs=email_input,198 outputs=[result_output, confidence_output]199 )200 201 with gr.Tab(" π Analysis"):202 gr.Markdown("## Dataset Overview")203 gr.Markdown("### Dataet Headers")204 gr.DataFrame(df)205 206 # Top 10 words for spam207 gr.Markdown("### Top Spam Words")208 top_spam_words = get_top_words(df[df['label'] == "spam"]['text'], n=10)209 gr.DataFrame(top_spam_words)210 211 # Top 10 words for ham212 gr.Markdown("### Top Ham Words")213 top_ham_words = get_top_words(df[df['label'] == "ham"]['text'], n=10)214 gr.DataFrame(top_ham_words)215 216 gr.Markdown("## π Model Performance Analytics")217 with gr.Row():218 gr.Textbox(value=performance_metrics["accuracy"], label="Accuracy", interactive=False)219 gr.Textbox(value=performance_metrics["precision"], label="Precision", interactive=False)220 gr.Textbox(value=performance_metrics["recall"], label="Recall", interactive=False)221 gr.Textbox(value=performance_metrics["f1_score"], label="F1 Score", interactive=False)222 223 with gr.Tab("π Glossary"):224 with gr.Column():225 gr.Markdown(226 """227 ## Label Definitions228 - Spam: Unwanted or potentially harmful emails detected by the system.229 - Ham: Legitimate and safe emails.230 231 ## Evaluation Metrics232 - Accuracy: Measures the percentage of correctly classified emails.233 - Precision: Out of all emails classified as spam, how many were actually spam?234 - Recall: Out of all actual spam emails, how many were identified correctly?235 - F1 Score: A balance between precision and recall for overall performance assessment.236 237 """ 238 )239 with gr.Column():240 gr.Markdown(" ## π Libraries Used and Their Objectives")241 gr.Markdown(242 """243 ### 1. Pandas (import pandas as pd)244 245 Objective: Data manipulation and preprocessing.246 Justification: Used for loading, cleaning, and structuring the email dataset for analysis and model training.247 248 ### 2. NumPy (import numpy as np)249 250 Objective: Efficient numerical operations.251 Justification: Facilitates handling large datasets and computations, such as text vectorization and matrix operations.252 253 ### 3. Torch & Torch-related Libraries254 255 import torch β Core deep learning framework for model training.256 import torch.nn as nn β Defines deep learning model architecture.257 import torch.optim as optim β Implements optimization algorithms.258 import torch.nn.functional as F β Provides additional functions like activation and loss functions.259 from torch.utils.data import Dataset, DataLoader β Handles data batching and loading for model training.260 Justification: Essential for training and fine-tuning DistilBERT on email classification.261 262 ### 4. Transformers (from transformers import DistilBertTokenizer, DistilBertForSequenceClassification)263 264 Objective: Tokenization and model training using DistilBERT.265 Justification: DistilBERT offers a lighter yet powerful alternative to BERT, improving efficiency while maintaining accuracy.266 267 ### 5. Scikit-learn (sklearn)268 269 Feature Extraction:270 CountVectorizer: Converts text into a matrix of token counts.271 TfidfVectorizer: Converts text into TF-IDF features, which measure the importance of words in documents.272 Model Training & Evaluation:273 MultinomialNB: Implements the NaΓ―ve Bayes classifier for a baseline model.274 train_test_split: Splits the dataset for training and testing.275 classification_report, accuracy_score, precision_score, recall_score, f1_score: Computes evaluation metrics.276 Justification: Used for feature extraction, baseline modeling, and performance evaluation of different models.277 278 ### 6. Matplotlib & Seaborn (import matplotlib.pyplot as plt, import seaborn as sns)279 280 Objective: Data visualization.281 Justification: Used to visualize word distributions, spam vs. ham comparisons, and model performance metrics.282 283 ### 7. Gradio (import gradio as gr)284 285 Objective: Building an interactive web-based demo.286 Justification: Allows users to test the spam detection system by inputting emails and viewing real-time predictions.287 """)288 with gr.Column():289 gr.Markdown("## π Thanks & Acknowledgments π")290 gr.Markdown("""291 ### π Special Thanks to Our Contributors292 293 **πΉ Remus** 294 - Led **Data Collection & Preprocessing**, ensuring a clean dataset for training. 295 - Developed the **Baseline Model**, which served as the foundation for further improvements. 296 - Fine-tuned **BERT**, optimizing hyperparameters to enhance accuracy. 297 298 **πΉ Ashley** 299 - Played a key role in **Data Collection & Preprocessing**, improving dataset quality. 300 - Successfully handled the **Deployment on Hugging Face**, making the model accessible to users. 301 - Implemented and optimized **DistilBERT**, achieving a balance between speed and performance. 302 303 This project was a collaborative effort, and we appreciate the hard work put into making it a success! π 304 """)305 306 307 return interface308 309# Launch the interface310interface = create_interface()311interface.launch(share=True)312 