Team Ai
Apppublic

leygit/ITI110_Spam_Classification_Project

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
apptest.py251 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.metrics import classification_report11from transformers import BertTokenizer12 13 14# Load dataset15file_path = 'spam_ham_dataset.csv'16df = pd.read_csv(file_path)17 18# Convert labels to numeric19df['label_num'] = df['label'].map({'ham': 0, 'spam': 1})20 21# Load tokenizer22tokenizer = DistilBertTokenizer.from_pretrained('distilbert-base-uncased')23 24# Tokenize dataset25encodings = tokenizer(df['text'].tolist(), padding=True, truncation=True, max_length=128, return_tensors="pt")26labels = torch.tensor(df['label_num'].values)27 28# Custom Dataset29class SpamDataset(Dataset):30    def __init__(self, encodings, labels):31        self.encodings = encodings32        self.labels = labels33 34    def __len__(self):35        return len(self.labels)36 37    def __getitem__(self, idx):38        item = {key: val[idx] for key, val in self.encodings.items()}39        item['labels'] = torch.tensor(self.labels[idx], dtype=torch.long)40        return item41 42# Create dataset43dataset = SpamDataset(encodings, labels)44 45# Split dataset (80% train, 20% validation)46train_size = int(0.8 * len(dataset))47val_size = len(dataset) - train_size48train_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size])49 50# DataLoader with batch size51def collate_fn(batch):52    keys = batch[0].keys()53    return {key: torch.stack([b[key] for b in batch]) for key in keys}54 55train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True, collate_fn=collate_fn)56val_loader = DataLoader(val_dataset, batch_size=16, shuffle=False, collate_fn=collate_fn)57 58# Load DistilBERT model59device = torch.device("cuda" if torch.cuda.is_available() else "cpu")60model = DistilBertForSequenceClassification.from_pretrained("distilbert-base-uncased", num_labels=2)61model.to(device)62 63# Define optimizer and loss function64optimizer = optim.AdamW(model.parameters(), lr=5e-5, weight_decay=0.01)65loss_fn = nn.CrossEntropyLoss()66 67# Training Loop68EPOCHS = 1069for epoch in range(EPOCHS):70    model.train()71    total_loss = 072 73    for batch in train_loader:74        optimizer.zero_grad()75 76        inputs = {key: val.to(device) for key, val in batch.items()}77        labels = inputs.pop("labels").to(device)78 79        outputs = model(**inputs)80        loss = loss_fn(outputs.logits, labels)81 82        loss.backward()83        optimizer.step()84 85        total_loss += loss.item()86 87    avg_loss = total_loss / len(train_loader)88    print(f"Epoch {epoch+1}, Loss: {avg_loss:.4f}")89 90# Save trained model91torch.save(model.state_dict(), "distilbert_spam_model.pt")92 93# Evaluation94model.eval()95correct = 096total = 097with torch.no_grad():98    for batch in val_loader:99        inputs = {key: val.to(device) for key, val in batch.items()}100        labels = inputs.pop("labels").to(device)101 102        outputs = model(**inputs)103        predictions = torch.argmax(outputs.logits, dim=1)104        correct += (predictions == labels).sum().item()105        total += labels.size(0)106 107accuracy = correct / total108print(f"Validation Accuracy: {accuracy:.4f}")109 110 111 112# Classification function113def classify_email(email_text):114    model.eval()  # Set model to evaluation mode115 116    with torch.no_grad():117        # Tokenize and convert input text to tensor118        inputs = tokenizer(email_text, padding=True, truncation=True, max_length=256, return_tensors="pt")119 120        # Move inputs to the appropriate device121        inputs = {key: val.to(device) for key, val in inputs.items()}122 123        # Get model predictions124        outputs = model(**inputs)125        logits = outputs.logits126 127        # Convert logits to predicted class128        predictions = torch.argmax(logits, dim=1)129 130        # Convert logits to probabilities using softmax131        probs = F.softmax(logits, dim=1)132        confidence = torch.max(probs).item() * 100  # Convert to percentage133 134    # Convert numeric prediction to label135    result = "Spam" if predictions.item() == 1 else "Ham"136 137    return {138        "result": result,139        "confidence": f"{confidence:.2f}%",140    }141 142# Evaluation function with detailed classification report143def evaluate_model_with_report(val_loader):144    model.eval()  # Set model to evaluation mode145    y_true = []146    y_pred = []147    correct = 0148    total = 0149 150    with torch.no_grad():151        for batch in val_loader:152            inputs = {key: val.to(device) for key, val in batch.items()}153            labels = inputs.pop("labels").to(device)154 155            outputs = model(**inputs)156            predictions = torch.argmax(outputs.logits, dim=1)157 158            # Collect labels and predictions159            y_true.extend(labels.cpu().numpy())160            y_pred.extend(predictions.cpu().numpy())161 162            # Calculate accuracy163            correct += (predictions == labels).sum().item()164            total += labels.size(0)165 166    # Calculate accuracy167    accuracy = correct / total if total > 0 else 0168    print(f"Validation Accuracy: {accuracy:.4f}")169 170    # Print classification report171    print("\nClassification Report:")172    print(classification_report(y_true, y_pred, target_names=["Ham", "Spam"]))173 174    return accuracy175 176# Run evaluation with classification report177accuracy = evaluate_model_with_report(val_loader)178print(f"Model Validation Accuracy: {accuracy:.4f}")179 180## Gradio Interface181 182import gradio as gr183 184# Create Gradio Interface185def create_interface():186    performance_metrics = generate_performance_metrics()187 188    # Introduction - Title + Brief Description189    with gr.Blocks(css=custom_css) as interface:190        gr.Markdown("Spam Email Classification")191        gr.Markdown(192            """193            Brief description of the project here194 195            """196        )197 198        # Email Text Input199        with gr.Row():200            email_input = gr.Textbox(201                lines=8, placeholder="Type or paste your email content here...", label="Email Content"202            )203 204        # Email Text Results and Analysis205        with gr.Row():206            result_output = gr.HTML(label="Classification Result") # label = [function that prints classification result]207            confidence_output = gr.Textbox(label="Confidence Score", interactive=False)208            accuracy_output = gr.Textbox(label="Accuracy", interactive=False)209 210 211        analyze_button = gr.Button("Analyze Email ๐Ÿ•ต๏ธโ€โ™‚๏ธ")212 213        analyze_button.click(214            fn=email_analysis_pipeline,215            inputs=email_input,216            outputs=[result_output, confidence_output, accuracy_output]217        )218 219        # Analysis220        gr.Markdown("## ๐Ÿ“Š Model Performance Analytics")221        with gr.Row():222            with gr.Column():223                gr.Textbox(value=performance_metrics["accuracy"], label="Accuracy", interactive=False, elem_classes=["metric"])224                gr.Textbox(value=performance_metrics["precision"], label="Precision", interactive=False, elem_classes=["metric"])225                gr.Textbox(value=performance_metrics["recall"], label="Recall", interactive=False, elem_classes=["metric"])226                gr.Textbox(value=performance_metrics["f1_score"], label="F1 Score", interactive=False, elem_classes=["metric"])227            with gr.Column():228                gr.Markdown("### Confusion Matrix")229                gr.HTML(f"<img src='data:image/png;base64,{performance_metrics['confusion_matrix_plot']}' style='max-width: 100%; height: auto;' />")230 231        gr.Markdown("## ๐Ÿ“˜ Glossary and Explanation of Labels")232        gr.Markdown(233            """234            ### Labels:235            - **Spam:** Unwanted or harmful emails flagged by the system.236            - **Ham:** Legitimate, safe emails.237 238            ### Metrics:239            - **Accuracy:** The percentage of correct classifications.240            - **Precision:** Out of predicted Spam, how many are actually Spam.241            - **Recall:** Out of all actual Spam emails, how many are predicted as Spam.242            - **F1 Score:** Harmonic mean of Precision and Recall.243            """244        )245 246    return interface247 248# Launch the interface249interface = create_interface()250interface.launch(share=True)251