Team Ai
Apppublic

leygit/ITI110_Spam_Classification_Project

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py312 linesDownload Raw Back to root
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