Step 1: Decide what a mistake costs
Before writing any code, decide which mistake is worse. A spam filter can make two kinds:
- False positive: a real message is sent to spam. Your user misses an invoice or a message from a friend.
- False negative: a spam message gets through. Annoying, but the user can delete it.
False positives are much worse, so this guide aims for very high precision on the spam class (when the model says spam, it's right) and then catches as much spam as possible, which is recall. Write your target down now, for example: at least 98% precision, as much recall as we can get. Without a target you can't tell when you're done.
Step 2: Get labeled data
You need examples of messages, each labeled spam or not spam (often called "ham"). Two good sources:
- Your own data: messages users marked as spam or moved back to the inbox. This is the best source, because it matches what your filter will see.
- A public dataset to learn on: the SMS Spam Collection from the UCI Machine Learning Repository has about 5,500 text messages, roughly 13% of them spam. It's a single tab-separated file called
SMSSpamCollection, with the label first and the message second.
CodeLoad and inspect the data
# pip install pandas scikit-learn
import csv
import pandas as pd
df = pd.read_csv("SMSSpamCollection", sep="\t", names=["label", "text"], quoting=csv.QUOTE_NONE)
df = df.drop_duplicates("text") # the same message twice would leak into the test set
print(df.label.value_counts())
print(df.sample(5, random_state=1))Step 3: Split the data before you look closely
Set aside a test set now and don't touch it until the end. It's your honest estimate of how the filter will do on messages it has never seen. Use a stratified split so both sets have the same share of spam.
CodeTrain/test split
from sklearn.model_selection import train_test_split
X_train, X_test, y_train, y_test = train_test_split(
df["text"], df["label"], test_size=0.2, stratify=df["label"], random_state=42)Step 4: Beat a baseline first
A baseline is the simplest possible answer. Here it's "everything is ham". It looks surprisingly good on accuracy, which is exactly why accuracy is the wrong metric for this problem.
CodeThe do-nothing baseline
from sklearn.dummy import DummyClassifier
from sklearn.metrics import classification_report
baseline = DummyClassifier(strategy="most_frequent").fit(X_train, y_train)
print(classification_report(y_test, baseline.predict(X_test), zero_division=0))
# ~87% accuracy, but spam recall is 0: it never catches a single spam message.Step 5: Train a real model
A strong, fast starting point for text is TF-IDF features with logistic regression. TF-IDF turns each message into numbers based on which words and word pairs it contains. Logistic regression learns which of those point towards spam. class_weight="balanced" stops the model ignoring the rarer spam class, a common issue with class imbalance.
CodeTF-IDF + logistic regression
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.linear_model import LogisticRegression
from sklearn.pipeline import make_pipeline
model = make_pipeline(
TfidfVectorizer(ngram_range=(1, 2), min_df=2, sublinear_tf=True),
LogisticRegression(max_iter=1000, class_weight="balanced"),
)
model.fit(X_train, y_train)
print(classification_report(y_test, model.predict(X_test)))On the SMS dataset this usually gets both precision and recall for spam well above 0.9, after a few seconds of training. Because the vectorizer is inside the pipeline, it learns its vocabulary from the training data only, which keeps the test set clean.
Step 6: Read the mistakes
Scores tell you how often the model is wrong; the mistakes themselves tell you why. Look at 20 of them before changing anything.
CodeShow misclassified messages
pred = model.predict(X_test)
mistakes = pd.DataFrame({"text": X_test, "true": y_test, "predicted": pred})
pd.set_option("display.max_colwidth", 120)
print(mistakes[mistakes.true != mistakes.predicted].head(20))Typical findings: very short spam ("Call now!"), legitimate messages from businesses that sound promotional, or labels that are simply wrong. Fixing bad labels often helps more than a fancier model.
Step 7: Choose the threshold on purpose
By default the model calls a message spam when it's more than 50% sure. That's rarely the right cut-off. Since false positives are expensive, raise the threshold until precision reaches your target.
Pick the threshold using cross-validation on the training set, not the test set. Otherwise you tune to the test set and its score stops being honest.
CodeFind the lowest threshold with ≥ 98% precision
from sklearn.model_selection import cross_val_predict
from sklearn.metrics import precision_recall_curve, precision_score, recall_score
spam = list(model.classes_).index("spam")
cv_scores = cross_val_predict(model, X_train, y_train, cv=5, method="predict_proba")[:, spam]
precision, recall, thresholds = precision_recall_curve(y_train == "spam", cv_scores)
ok = precision[:-1] >= 0.98
threshold = thresholds[ok][0] if ok.any() else 0.5
print(f"Threshold: {threshold:.2f}")
# Now check it once on the untouched test set
test_scores = model.predict_proba(X_test)[:, spam]
is_spam = test_scores >= threshold
print("Test precision:", precision_score(y_test == "spam", is_spam))
print("Test recall: ", recall_score(y_test == "spam", is_spam))Step 8: Save the model and serve it
Retrain on all your data (training and test) once you're happy, save it with the chosen threshold, and wrap it in a small API so your app can call it.
CodeSave the final model
import joblib
model.fit(df["text"], df["label"])
joblib.dump({"model": model, "threshold": float(threshold)}, "spam_model.joblib")Codeapp.py: a prediction API
# pip install fastapi uvicorn joblib scikit-learn
import joblib
from fastapi import FastAPI
from pydantic import BaseModel
saved = joblib.load("spam_model.joblib")
model, threshold = saved["model"], saved["threshold"]
spam = list(model.classes_).index("spam")
app = FastAPI()
class Message(BaseModel):
text: str
@app.post("/predict")
def predict(message: Message):
score = float(model.predict_proba([message.text])[0][spam])
return {"spam_probability": round(score, 3), "is_spam": score >= threshold}
# Run with: uvicorn app:app --reload
# Try it: curl -X POST localhost:8000/predict -H "Content-Type: application/json" -d '{"text": "WIN a FREE prize, call now"}'This model is tiny and runs on a CPU in about a millisecond, so the cheapest serverless option on any cloud is plenty. See the cloud comparison for the matching service on each provider.
Step 9: Monitor it and keep improving
- Log every prediction with its score, and record when users move a message in or out of spam. Those corrections are free new labels.
- Watch the share of messages flagged. A sudden jump or drop usually means spammers changed tactics or something broke. This is data drift.
- Retrain regularly, for example monthly, on the latest labeled data, and only deploy the new model if it beats the old one on a fresh test set.
- When spammers adapt, try fine-tuning a small pretrained language model such as DistilBERT. It understands wording it has never seen, at the cost of needing a bit more compute.
Common pitfalls
- Judging the model on accuracy. With 87% ham, "never spam" scores 87%.
- Leaving duplicates in the data, so the test set contains messages the model has already memorised.
- Tuning the threshold on the test set, which makes the final score look better than reality.
- Never retraining. Spam changes constantly, and a filter left alone slowly gets worse.