--- license: mit pipeline_tag: text-classification tags: - ml-lab - pytorch - text-classification - spam-detection - distilbert - transformers datasets: - SetFit/enron_spam base_model: distilbert/distilbert-base-uncased metrics: - f1 - accuracy - precision - recall ml_lab: title: Email Spam Classifier order: 6 summary: Reads an email (subject + body) and says whether it is spam, using a fine-tuned DistilBERT. framework: PyTorch (transformers) architecture: DistilBERT (uncased) fine-tuned, last 2 blocks + linear head dataset: Enron-Spam space: shalev396/email-spam-classifier runtime: cpu-basic ui_kind: spam-inbox colab: https://colab.research.google.com/github/shalev396/ml-lab/blob/main/email-spam-classifier/training/notebook.ipynb github: https://github.com/shalev396/ml-lab/tree/main/email-spam-classifier metric: name: f1 value: 0.9937743190661479 split: test params: 66363649 params_human: 66.36M trained_at: '2026-07-13' model-index: - name: email-spam-classifier results: - task: type: text-classification dataset: name: Enron-Spam type: SetFit/enron_spam split: test metrics: - type: accuracy value: 0.993663 name: accuracy - type: precision value: 0.993002 name: precision - type: recall value: 0.994548 name: recall - type: f1 value: 0.993774 name: f1 - type: roc_auc value: 0.999753 name: roc_auc - type: average_precision value: 0.999764 name: average_precision --- # Email Spam Classifier Reads an email (subject + body) and returns the probability that it is **spam** or **ham** (a legitimate email). It is DistilBERT (uncased) with a new one-logit head. The last 2 of its 6 transformer blocks and the head were fine-tuned on ~23.6k emails from the Enron-Spam corpus. On the held-out 5,050-email test split it makes 32 mistakes (F1 0.9938, accuracy 0.9937). ## Model | | | |---|---| | Architecture | `distilbert/distilbert-base-uncased` encoder (6 blocks, dim 768) -> `[CLS]` hidden state -> `Dropout(0.3)` -> `Linear(768, 1)`; sigmoid = P(spam) | | Parameters | 66,363,649 total, of which 14,176,513 (the last 2 blocks + the head) were trained | | Input | one string, ideally `"Subject: \n\n"` | | Preprocessing | `model.prepare_text`: drop a leading `Subject:` header, then `model.clean_text` (HTML unescape, strip tags and URLs, lowercase, keep only `a-z 0-9 . , ! ? $ % ' -`, collapse whitespace). Same function as in training. WordPiece tokenizer, max 256 tokens (longer emails are truncated) | | Output | `{"spam": p, "ham": 1 - p}`; `Predictor.label()` says spam when p >= 0.5 | | Files | `model.safetensors` (weights), `config.json` (`backbone_config`, `dropout`, `max_len`, `threshold`, `labels`), `tokenizer.json` + `tokenizer_config.json`, `model.py` (architecture + cleaning + `load()` + `Predictor`), `handler.py` (Inference Endpoint) | `config.json` stores the full DistilBERT config, so `model.load()` rebuilds the encoder with `AutoModel.from_config` (`pretrained=False`) and then loads `model.safetensors`. Inference never downloads the base model. ## Usage ```python from huggingface_hub import snapshot_download import sys path = snapshot_download("shalev396/email-spam-classifier") sys.path.insert(0, path) import model predictor = model.load(path, device="cpu") # or "cuda" print(predictor.predict("Subject: Notes from this morning\n\nHi team, attached are the notes. Thanks, Sarah")) # {'spam': 1.2e-06, 'ham': 0.999999} print(predictor.predict(model.compose_email("You won a $1,000 gift card", "Click here to claim your prize now!"))) # {'spam': 0.9998, 'ham': 0.0002} print(predictor.predict_proba(["email one ...", "email two ..."])) # batch -> [P(spam), ...] ``` Requirements: `torch`, `transformers>=5`, `tokenizers`, `huggingface_hub`, `safetensors`. - **Space / free API:** [shalev396/email-spam-classifier](https://huggingface.co/spaces/shalev396/email-spam-classifier), `POST /gradio_api/call/predict` with `{"data": ["Subject: ...\n\n..."]}`. - **Inference Endpoint:** deploy this repo (Deploy -> Inference Endpoints). `handler.py` accepts `{"inputs": "Subject: ...\n\n..."}`, a list of such strings, or `{"inputs": {"subject": ..., "body": ...}}`, and uses a GPU when the endpoint has one. ## Training - **Data:** [Enron-Spam](https://huggingface.co/datasets/SetFit/enron_spam) (Metsis, Androutsopoulos & Paliouras, 2006), the SetFit copy: 33,716 emails (subject + body). Ham comes from the mailboxes of six Enron employees, spam from several spam traps. After cleaning and dropping empty texts, 33,665 emails were split stratified 70/15/15 (seed 42): 23,565 train / 5,050 validation / 5,050 test, about 51% spam in each. - **Recipe:** the encoder is frozen except its last 2 blocks. Loss is `BCEWithLogitsLoss` with `pos_weight = n_ham / n_spam = 0.966`. AdamW with differential learning rates (head 5e-4, encoder 2e-5, weight decay 0.01), 10% linear warmup then linear decay, batch 32, max 256 tokens, up to 5 epochs with early stopping on validation loss (patience 2), seed 42. - **This checkpoint:** from the original training run of this project (2026-07-13, CUDA GPU with mixed precision). Validation loss was lowest after epoch 5, so that epoch is the checkpoint. The original state dict was mapped 1:1 into `model.SpamClassifier` (all 102 tensors, `strict=True`) and saved as safetensors. It was then re-evaluated on CPU through `model.Predictor`, the code path the Space uses. The confusion matrix matches the original run's log exactly. Training time was not recorded. Full code: [training/](https://github.com/shalev396/ml-lab/tree/main/email-spam-classifier/training) ยท [Colab](https://colab.research.google.com/github/shalev396/ml-lab/blob/main/email-spam-classifier/training/notebook.ipynb). The notebook reproduces this recipe step by step (data -> model -> training -> evaluation -> inference -> export). ![Training curves of the original run](assets/training_curves.png) ## Experiments | test split (5,050 emails) | accuracy | precision | recall | F1 | ROC-AUC | errors | |---|---|---|---|---|---|---| | **DistilBERT fine-tuned, last 2 blocks (deployed)** | **0.9937** | **0.9930** | 0.9945 | **0.9938** | **0.9998** | **32** (18 ham -> spam, 14 spam -> ham) | | TF-IDF 1-2-grams + logistic regression (baseline) | 0.9913 | 0.9880 | **0.9949** | 0.9915 | 0.9992 | 44 (31 ham -> spam, 13 spam -> ham) | Both rows are in `metrics.json` (`comparison`). The baseline (50k TF-IDF features, `C=10`, class-balanced) trains in about 25 s on a CPU and is already strong on this corpus. DistilBERT's gain is mostly **fewer false alarms**: it sends 18 legitimate emails to spam instead of 31. Recall is about the same for both. ![Test error of each experiment](assets/comparison.png) ## Evaluation | metric (test) | value | |---|---| | accuracy | 0.9937 | | precision | 0.9930 | | recall | 0.9945 | | f1 | 0.9938 | | roc_auc | 0.9998 | | average_precision | 0.9998 | Test split, threshold 0.5. Precision, recall and F1 are for the spam class. `average_precision` is the area under the precision-recall curve. ![Confusion matrix on the test split](assets/confusion_matrix.png) ![ROC and precision-recall curves](assets/roc_pr_curves.png) ## Limitations - **One corpus, from the 2000s.** Enron-Spam's ham is the mail of a single energy company from around 2000-2002, and its spam is from the same era. Modern phishing, newsletters and marketing mail look different. Expect lower accuracy on today's inboxes, and treat the 99.4% as an in-distribution number. - **Random split.** Train and test come from the same mailboxes and the same period, so near-duplicate emails can appear on both sides. That makes the test score optimistic. - **Text only.** No headers, sender, links (URLs are removed during cleaning) or attachments. A spam filter in production uses all of these. - **English, lowercase.** The uncased English model sees lowercased text, so it cannot use capitalisation ("FREE!!!") as a signal. - **Truncation.** Only the first 256 tokens (roughly 150-200 words) of an email are read. - **Threshold.** 0.5 was not tuned. If false alarms cost more than missed spam, raise it.