Supreetha15/safe-email-classifier
0
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 