"""Email Spam Classifier: DistilBERT fine-tuned to tell spam from ham (legitimate email). Single source of truth for the architecture, the text cleaning and inference. It is used by the Hugging Face Space (`space/app.py`), the Inference Endpoint (`handler.py`), the training code (`training/src/`) and anyone who downloads this repo from the Hub: import model predictor = model.load("path/to/this/repo", device="cpu") predictor.predict("Subject: You won!\\n\\nClaim your $1000 prize now") # -> {"spam": 0.99, "ham": 0.01} """ from __future__ import annotations import html import re from pathlib import Path import torch from huggingface_hub import PyTorchModelHubMixin from torch import nn from transformers import AutoConfig, AutoModel, AutoTokenizer REPO_ID = "shalev396/email-spam-classifier" FRAMEWORK = "pytorch" WEIGHTS_FILE = "model.safetensors" BASE_MODEL = "distilbert/distilbert-base-uncased" LABELS = ["ham", "spam"] # logit = spam score; sigmoid(logit) = P(spam) def cuda_available() -> bool: return torch.cuda.is_available() # --------------------------------------------------------------------------- text cleaning _SUBJECT_RE = re.compile(r"^\s*subject\s*:\s*", flags=re.I) _HTML_TAG_RE = re.compile(r"<[^>]+>") _URL_RE = re.compile(r"(?:https?://|www\.)\S+") # Keep the punctuation that carries spam signal ($, %, !) plus basic sentence punctuation. _SPECIAL_CHARS_RE = re.compile(r"[^a-z0-9\s.,!?$%'-]") _WHITESPACE_RE = re.compile(r"\s+") def clean_text(text: str) -> str: """Raw email text -> lowercase text without HTML, URLs or unusual characters. Applied to every training example (training/src/data_setup.py) and to every input at inference, so both see exactly the same kind of text. Idempotent. """ if not isinstance(text, str): return "" text = html.unescape(text) text = _HTML_TAG_RE.sub(" ", text) text = _URL_RE.sub(" ", text) text = text.lower() text = _SPECIAL_CHARS_RE.sub(" ", text) text = _WHITESPACE_RE.sub(" ", text) return text.strip() def compose_email(subject: str, body: str) -> str: """(subject, body) -> "Subject: \\n\\n", the text format /predict accepts.""" return f"Subject: {(subject or '').strip()}\n\n{(body or '').strip()}" def prepare_text(text: str) -> str: """Model input text. Drops a leading "Subject:" header, then `clean_text`. The Enron training emails are " " without the header word, so an input like "Subject: \\n\\n" becomes " ", the same shape the model was trained on. """ if not isinstance(text, str): return "" return clean_text(_SUBJECT_RE.sub("", text, count=1)) # --------------------------------------------------------------------------- architecture class SpamClassifier( nn.Module, PyTorchModelHubMixin, tags=["ml-lab", "text-classification", "spam-detection", "distilbert"], repo_url="https://github.com/shalev396/ml-lab", pipeline_tag="text-classification", license="mit", ): """DistilBERT encoder + `Dropout(dropout) -> Linear(hidden, 1)` on the [CLS] token. One logit per email; sigmoid(logit) = P(spam). Args: base_model: Hub id of the pretrained encoder (only downloaded when `pretrained=True`). backbone_config: the encoder's `transformers` config as a dict. Stored in config.json by the mixin so the encoder can be rebuilt offline with `AutoModel.from_config`. dropout: dropout in front of the linear head. max_len: tokens per email (longer emails are truncated). threshold: P(spam) at or above which an email counts as spam. labels: class names, index 1 is the positive (spam) class. pretrained: load the encoder's pretrained weights from the Hub. Only needed to *train*; `load()` always passes `pretrained=False`, so inference never downloads the base model. """ def __init__(self, base_model: str = BASE_MODEL, backbone_config: dict | None = None, dropout: float = 0.3, max_len: int = 256, threshold: float = 0.5, labels: list[str] | None = None, pretrained: bool = False): super().__init__() self.base_model = base_model self.max_len = int(max_len) self.threshold = float(threshold) self.labels = list(labels or LABELS) if pretrained: self.backbone = AutoModel.from_pretrained(base_model) else: config = (AutoConfig.for_model(**backbone_config) if backbone_config else AutoConfig.from_pretrained(base_model)) self.backbone = AutoModel.from_config(config) self.dropout = nn.Dropout(dropout) self.classifier = nn.Linear(self.backbone.config.hidden_size, 1) def encoder_layers(self) -> nn.ModuleList: """The stack of transformer blocks (DistilBERT: `transformer.layer`, BERT-likes: `encoder.layer`).""" if hasattr(self.backbone, "transformer"): return self.backbone.transformer.layer if hasattr(self.backbone, "encoder"): return self.backbone.encoder.layer raise AttributeError(f"no encoder layers found on {type(self.backbone).__name__}") def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor: """(N, T) token ids + mask -> (N,) spam logits.""" hidden = self.backbone(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state return self.classifier(self.dropout(hidden[:, 0])).squeeze(-1) def encode(tokenizer, texts: list[str], max_len: int, device=None) -> dict[str, torch.Tensor]: """Already-prepared texts -> padded (to the longest in the batch) + truncated tensors.""" batch = tokenizer(list(texts), padding=True, truncation=True, max_length=max_len, return_tensors="pt") batch = {"input_ids": batch["input_ids"], "attention_mask": batch["attention_mask"]} return {k: v.to(device) for k, v in batch.items()} if device is not None else batch # --------------------------------------------------------------------------- inference class Predictor: """Loads the fine-tuned weights + tokenizer once and scores emails.""" def __init__(self, model_dir: str | Path, device: str = "cpu"): model_dir = Path(model_dir) if not (model_dir / WEIGHTS_FILE).is_file(): raise FileNotFoundError( f"{WEIGHTS_FILE} not found in {model_dir}: train/export the model first " f"or download it with huggingface_hub.snapshot_download('{REPO_ID}')." ) if not (model_dir / "tokenizer.json").is_file(): raise FileNotFoundError(f"tokenizer.json not found in {model_dir}: export the tokenizer with the weights.") self.device = torch.device(device) self.model = SpamClassifier.from_pretrained( model_dir, pretrained=False, map_location="cpu", strict=True ) self.model.to(self.device).eval() self.tokenizer = AutoTokenizer.from_pretrained(model_dir) self.max_len = self.model.max_len self.threshold = self.model.threshold @torch.inference_mode() def predict_proba(self, texts: list[str], batch_size: int = 32) -> list[float]: """Raw email texts -> P(spam) for each. Batches are grouped by length for speed.""" prepared = [prepare_text(t) for t in texts] order = sorted(range(len(prepared)), key=lambda i: len(prepared[i])) probs = [0.0] * len(prepared) for start in range(0, len(order), batch_size): idx = order[start:start + batch_size] batch = encode(self.tokenizer, [prepared[i] for i in idx], self.max_len, self.device) p = torch.sigmoid(self.model(**batch).float()).cpu().tolist() for i, value in zip(idx, p): probs[i] = float(value) return probs def predict(self, text: str) -> dict[str, float]: """Email text (ideally "Subject: \\n\\n") -> {"spam": p, "ham": 1 - p}.""" p = self.predict_proba([text])[0] return {"spam": p, "ham": 1.0 - p} def label(self, text: str) -> str: """"spam" if P(spam) >= threshold else "ham".""" return "spam" if self.predict(text)["spam"] >= self.threshold else "ham" def load(model_dir: str | Path, device: str = "cpu") -> Predictor: return Predictor(model_dir, device)