CoolFace
Apppublic

Supreetha15/safe-email-classifier

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
models.py87 linesDownload Raw Back to root
1# models.py
2
3from sentence_transformers import SentenceTransformer
4from sklearn.svm import LinearSVC
5from sklearn.preprocessing import LabelEncoder
6from sklearn.model_selection import train_test_split
7from sklearn.metrics import classification_report
8from joblib import dump, load
9import pandas as pd
10import re
11
12
13class SBERT_SVM_Classifier:
14    """
15    SBERT + Linear SVM classifier for email subject classification.
16    """
17
18    def __init__(self, model_name="paraphrase-MiniLM-L6-v2"):
19        """
20        Initialize sentence transformer and SVM model.
21        """
22        self.encoder = SentenceTransformer(model_name)
23        self.label_encoder = LabelEncoder()
24        self.model = LinearSVC(C=1.0)
25
26    def preprocess(self, text):
27        """
28        Lowercase and clean subject prefix.
29        """
30        text = text.lower()
31        return re.sub(r'^subject:\s*', '', text.strip())
32
33    def train(self, csv_path):
34        """
35        Train model on email dataset.
36        """
37        df = pd.read_csv(csv_path)
38        df['email'] = df['email'].apply(self.preprocess)
39
40        X = df['email'].tolist()
41        y = self.label_encoder.fit_transform(df['type'].tolist())
42
43        print("[*] Generating sentence embeddings...")
44        X_embed = self.encoder.encode(
45            X,
46            batch_size=32,
47            show_progress_bar=True,
48            convert_to_numpy=True
49        )
50
51        print("[*] Training Linear SVM (FAST mode)...")
52        self.model.fit(X_embed, y)
53
54        # Optional evaluation
55        X_train, X_test, y_train, y_test = train_test_split(
56            X_embed, y, test_size=0.2, random_state=42
57        )
58        y_pred = self.model.predict(X_test)
59        print(
60            classification_report(
61                y_test,
62                y_pred,
63                target_names=self.label_encoder.classes_
64            )
65        )
66
67    def predict(self, email_texts):
68        """
69        Predict category for list of email strings.
70        """
71        emails = [self.preprocess(e) for e in email_texts]
72        X = self.encoder.encode(emails, convert_to_numpy=True)
73        y_pred = self.model.predict(X)
74        return self.label_encoder.inverse_transform(y_pred)
75
76    def save(self, path="sbert_linear_model.joblib"):
77        """
78        Save model + label encoder to disk.
79        """
80        dump((self.model, self.label_encoder), path)
81
82    def load(self, path="sbert_linear_model.joblib"):
83        """
84        Load model + label encoder from disk.
85        """
86        self.model, self.label_encoder = load(path)
87