Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Mini-Project: Transformer-Based Text Classifier
A model reports 95% accuracy on a support-ticket classifier and gets deployed. Two weeks later the team discovers it has never once predicted the "urgent" class. Nothing was broken. The training set was 95% routine tickets, and predicting "routine" every single time scores 95%.
Building a text classifier is not hard. Building one whose reported number means something is, and that is what this project is really about. You will assemble a complete pipeline — data, tokenisation, a fine-tuned transformer, evaluation, error analysis, attention inspection, and a serving endpoint — and at each stage you will check the thing that stops the number from lying to you.
The target is a working sentiment classifier with macro-F1 above 0.90, an error analysis that names its actual failure modes, and an inference path you could put behind an API.
Setting up
1python -m venv .venv && source .venv/bin/activate23pip install "torch>=2.0" transformers datasets scikit-learn \4 matplotlib seaborn pandas numpy accelerate1import torch, transformers, numpy as np, random23print("torch", torch.__version__, "| transformers", transformers.__version__)4if torch.cuda.is_available():5 print("GPU:", torch.cuda.get_device_name(0),6 f"{torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")7 DEVICE = "cuda"8elif torch.backends.mps.is_available():9 DEVICE = "mps" # Apple silicon10else:11 DEVICE = "cpu"12print("device:", DEVICE)1314def set_seed(s=42):15 random.seed(s); np.random.seed(s)16 torch.manual_seed(s); torch.cuda.manual_seed_all(s)1718set_seed(42)Set the seed before anything else, and remember what it does and does not buy you. It makes one run reproducible. It does not make your result reliable — BERT fine-tuning on small datasets is genuinely seed-sensitive, and two seeds can differ by several points. A single number from a single seed is a data point, not a finding.
Memory planning, since this is where most runs die. Fine-tuning BERT-base with AdamW needs, before any activations:
| Item (110 M parameters, fp32) | Memory |
|---|---|
| Weights | 440 MB |
| Gradients | 440 MB |
| Adam first moment | 440 MB |
| Adam second moment | 440 MB |
| Subtotal | 1.76 GB |
| Activations, batch 32 at length 128 | ~2–3 GB |
So roughly 5 GB. If you hit an out-of-memory error, halve the batch size and use gradient accumulation to keep the effective batch the same.
Data
1from datasets import load_dataset2import pandas as pd34ds = load_dataset("stanfordnlp/imdb") # 25k train / 25k test, balanced5train_df = pd.DataFrame(ds["train"]).sample(frac=1, random_state=42)6test_df = pd.DataFrame(ds["test"])78# carve a validation split OUT OF TRAIN — never out of test9val_df = train_df.iloc[:2500]10train_df = train_df.iloc[2500:]1112print(len(train_df), len(val_df), len(test_df)) # 22500 2500 25000For your own data, the only requirement is a text column and an integer label column:
1df = pd.read_csv("tickets.csv") # columns: text, label2label_names = sorted(df["label"].unique())3label2id = {name: i for i, name in enumerate(label_names)}4df["label_id"] = df["label"].map(label2id)Look at the data before you model it
Three checks, and none of them is optional.
1from transformers import AutoTokenizer23MODEL_NAME = "bert-base-uncased"4tok = AutoTokenizer.from_pretrained(MODEL_NAME)56# 1. Class balance -> sets the baseline you must beat7counts = train_df["label"].value_counts()8print(counts)9print("majority-class accuracy:", counts.max() / counts.sum())1011# 2. Token lengths -> sets max_len12lengths = [len(tok.encode(t, truncation=False)) for t in train_df["text"][:2000]]13print("percentiles 50/90/95/99:", np.percentile(lengths, [50, 90, 95, 99]))1415# 3. Read twenty examples with their labels. Actually read them.16for _, r in train_df.sample(20, random_state=0).iterrows():17 print(f"[{r['label']}] {r['text'][:160]}...")On IMDB the length percentiles come out around [230, 620, 795, 1153]. That is a real decision point. BERT's position table has exactly 512 rows, so anything beyond 512 tokens is discarded silently. At max_len=256 you keep about 57% of reviews intact; at 512 about 84%, at four times the attention cost.
Sequence length drives cost quadratically inside attention, so:
max_len | Reviews fully retained | Attention entries per head | Relative attention cost |
|---|---|---|---|
| 128 | ~12% | 16,384 | 1× |
| 256 | ~57% | 65,536 | 4× |
| 512 | ~84% | 262,144 | 16× |
For IMDB, 256 is the sensible compromise — sentiment is usually stated early and repeated. For a task where the decisive information sits at the end, truncation from the right is exactly wrong and you would truncate from the left instead.
The dataset class and the model
1from torch.utils.data import Dataset, DataLoader2from transformers import DataCollatorWithPadding34class TextClassificationDataset(Dataset):5 def __init__(self, texts, labels, tokenizer, max_len=256):6 self.texts = list(texts)7 self.labels = list(labels)8 self.tok = tokenizer9 self.max_len = max_len1011 def __len__(self):12 return len(self.texts)1314 def __getitem__(self, i):15 enc = self.tok(self.texts[i], truncation=True, max_length=self.max_len)16 enc["labels"] = self.labels[i]17 return enc # no padding here — the collator does it1819collator = DataCollatorWithPadding(tok, return_tensors="pt")2021def make_loader(df, shuffle, bs=16, max_len=256):22 dset = TextClassificationDataset(df["text"], df["label"], tok, max_len)23 return DataLoader(dset, batch_size=bs, shuffle=shuffle,24 collate_fn=collator, num_workers=2, pin_memory=True)2526train_loader = make_loader(train_df, True)27val_loader = make_loader(val_df, False)28test_loader = make_loader(test_df, False)Padding in the collator rather than in __getitem__ is worth 30–50% of training time on variable-length data. Each batch is padded to its own longest member instead of to a global 256, and most batches are far shorter than that.
1import torch.nn as nn2from transformers import AutoModel34class TransformerClassifier(nn.Module):5 def __init__(self, model_name=MODEL_NAME, num_labels=2, dropout=0.1):6 super().__init__()7 self.encoder = AutoModel.from_pretrained(model_name)8 hidden = self.encoder.config.hidden_size9 self.dropout = nn.Dropout(dropout)10 self.classifier = nn.Linear(hidden, num_labels)1112 def forward(self, input_ids, attention_mask,13 token_type_ids=None, output_attentions=False):14 out = self.encoder(input_ids=input_ids,15 attention_mask=attention_mask,16 token_type_ids=token_type_ids,17 output_attentions=output_attentions)18 cls = out.last_hidden_state[:, 0] # the [CLS] position19 logits = self.classifier(self.dropout(cls))20 return (logits, out.attentions) if output_attentions else logitsThe head is 768×2+2=1,538 parameters on top of 109,482,240. That is 0.0014% of the model — a reminder that you are learning a readout, not a capability. Everything the classifier knows about language came from pre-training.
The [CLS] position works as a summary because BERT's attention is bidirectional: that token has no word of its own, so across twelve layers it accumulates whatever the model finds worth aggregating from all other positions. In a causal model this would not work — position 0 can only see itself — and you would pool the last real token instead.
Training
1from dataclasses import dataclass23@dataclass4class Config:5 model_name: str = MODEL_NAME6 max_len: int = 2567 batch_size: int = 168 lr: float = 2e-59 epochs: int = 310 warmup_ratio: float = 0.111 weight_decay: float = 0.0112 max_grad_norm: float = 1.013 accum_steps: int = 114 seed: int = 421516cfg = Config()1import torch.nn.functional as F2from transformers import get_linear_schedule_with_warmup3from sklearn.metrics import f1_score, accuracy_score4from torch.amp import autocast, GradScaler56def build_optimizer(model, cfg):7 decay, no_decay = [], []8 for name, p in model.named_parameters():9 if not p.requires_grad:10 continue11 (no_decay if any(k in name for k in ("bias", "LayerNorm.weight"))12 else decay).append(p)13 return torch.optim.AdamW(14 [{"params": decay, "weight_decay": cfg.weight_decay},15 {"params": no_decay, "weight_decay": 0.0}],16 lr=cfg.lr, eps=1e-8)1718@torch.no_grad()19def evaluate(model, loader, device):20 model.eval()21 preds, golds = [], []22 for batch in loader:23 batch = {k: v.to(device) for k, v in batch.items()}24 labels = batch.pop("labels")25 logits = model(**batch)26 preds.extend(logits.argmax(-1).cpu().tolist())27 golds.extend(labels.cpu().tolist())28 return (accuracy_score(golds, preds),29 f1_score(golds, preds, average="macro"),30 np.array(golds), np.array(preds))3132def train(model, train_loader, val_loader, cfg, device=DEVICE):33 model.to(device)34 opt = build_optimizer(model, cfg)35 total_steps = (len(train_loader) // cfg.accum_steps) * cfg.epochs36 sched = get_linear_schedule_with_warmup(37 opt, int(total_steps * cfg.warmup_ratio), total_steps)38 scaler = GradScaler(device) if device == "cuda" else None3940 best_f1, best_state = 0.0, None4142 for epoch in range(1, cfg.epochs + 1):43 model.train()44 running = 0.045 opt.zero_grad(set_to_none=True)4647 for step, batch in enumerate(train_loader, 1):48 batch = {k: v.to(device) for k, v in batch.items()}49 labels = batch.pop("labels")5051 if scaler:52 with autocast("cuda", dtype=torch.float16):53 loss = F.cross_entropy(model(**batch), labels)54 scaler.scale(loss / cfg.accum_steps).backward()55 else:56 loss = F.cross_entropy(model(**batch), labels)57 (loss / cfg.accum_steps).backward()5859 if step % cfg.accum_steps == 0:60 if scaler:61 scaler.unscale_(opt) # BEFORE clipping62 torch.nn.utils.clip_grad_norm_(model.parameters(),63 cfg.max_grad_norm)64 if scaler:65 scaler.step(opt); scaler.update()66 else:67 opt.step()68 sched.step()69 opt.zero_grad(set_to_none=True)7071 running += loss.item()7273 acc, f1, _, _ = evaluate(model, val_loader, device)74 print(f"epoch {epoch} train_loss {running/len(train_loader):.4f} "75 f"val_acc {acc:.4f} val_macro_f1 {f1:.4f}")7677 if f1 > best_f1:78 best_f1 = f179 best_state = {k: v.detach().cpu().clone()80 for k, v in model.state_dict().items()}8182 model.load_state_dict(best_state)83 return model, best_f1Four lines in that loop are the ones that silently ruin runs if you get them wrong.
| Line | What breaks without it |
|---|---|
scaler.unscale_(opt) before clipping | Gradients still scaled by ~65,536 look like outliers, so every batch is crushed to 1/65536 and the model never moves |
loss / cfg.accum_steps | Gradients sum rather than average, giving an effective learning rate multiplied by accum_steps |
sched.step() inside the accumulation guard | The schedule advances faster than the optimiser; warmup ends early |
Restoring best_state | You ship the last epoch, which on a small dataset is usually the most overfitted one |
1set_seed(cfg.seed)2model = TransformerClassifier(cfg.model_name, num_labels=2)3model, best_f1 = train(model, train_loader, val_loader, cfg)45test_acc, test_f1, y_true, y_pred = evaluate(model, test_loader, DEVICE)6print(f"TEST accuracy {test_acc:.4f} macro-F1 {test_f1:.4f}")78torch.save({"state_dict": model.state_dict(), "config": cfg.__dict__},9 "classifier.pt")A healthy run on IMDB with these settings looks roughly like:
epoch 1 train_loss 0.2814 val_acc 0.9248 val_macro_f1 0.9247epoch 2 train_loss 0.1436 val_acc 0.9336 val_macro_f1 0.9336epoch 3 train_loss 0.0721 val_acc 0.9312 val_macro_f1 0.9311TEST accuracy 0.9327 macro-F1 0.9326Note that epoch 3 is slightly worse than epoch 2 on validation while the training loss halved again. That is the beginning of overfitting, and it is why the best checkpoint is kept rather than the last.
Reading the results properly
The confusion matrix, and the arithmetic on it
1from sklearn.metrics import confusion_matrix, classification_report2import seaborn as sns, matplotlib.pyplot as plt34cm = confusion_matrix(y_true, y_pred)5sns.heatmap(cm, annot=True, fmt="d", cmap="Blues",6 xticklabels=["neg", "pos"], yticklabels=["neg", "pos"])7plt.xlabel("predicted"); plt.ylabel("true"); plt.tight_layout()89print(classification_report(y_true, y_pred, target_names=["neg", "pos"], digits=4))Take a concrete matrix on a 1,000-example balanced test set:
| predicted neg | predicted pos | row total | |
|---|---|---|---|
| true neg | 440 | 60 | 500 |
| true pos | 45 | 455 | 500 |
Work every metric out by hand once, because the shape of these formulas is what tells you which one to care about.
For the positive class: 455 true positives, 60 false positives, 45 false negatives.
For the negative class: 440 true, 45 false positives, 60 false negatives.
On balanced data, accuracy and macro-F1 agree — 0.895 both ways. Now change the data. Suppose the test set is 950 negative and 50 positive, and the model predicts "negative" for everything:
but Rpos=0/50=0, so F1pos=0, while Pneg=950/1000=0.95 and Rneg=1.0 give F1neg=2(0.95)(1)/1.95=0.9744. Therefore
Accuracy says 95%. Macro-F1 says 49%. The second number is the honest one, and it is the one that would have caught the ticket classifier in the opening.
Report macro-F1 and per-class recall on any data that is not close to balanced. Accuracy on imbalanced data measures the class distribution, not the model.
Error analysis
Aggregate metrics tell you how much you are wrong. Reading the errors tells you why, and that is what determines what to do next.
1@torch.no_grad()2def collect_errors(model, loader, texts, device=DEVICE, k=25):3 model.eval()4 rows, offset = [], 05 for batch in loader:6 batch = {k_: v.to(device) for k_, v in batch.items()}7 labels = batch.pop("labels")8 probs = torch.softmax(model(**batch), dim=-1)9 conf, pred = probs.max(-1)10 for j in range(len(labels)):11 if pred[j] != labels[j]:12 rows.append({13 "text": texts[offset + j][:300],14 "true": labels[j].item(),15 "pred": pred[j].item(),16 "confidence": round(conf[j].item(), 4),17 })18 offset += len(labels)19 # most confident mistakes first — these are the most informative20 return pd.DataFrame(rows).sort_values("confidence", ascending=False).head(k)Sorting by confidence is the important part. A wrong prediction at 0.51 confidence is a borderline case. A wrong prediction at 0.99 confidence means the model has learned something false, and reading twenty of those usually reveals a pattern.
| Pattern you find | What it means | What to do |
|---|---|---|
| Sarcasm and irony ("brilliant, another two hours of my life gone") | Surface sentiment words contradict the meaning | A genuinely hard class. More data of this type, or accept the ceiling |
| Mixed sentiment ("great acting, terrible script") | The label is arguably wrong; the document has both | Reconsider whether one label per document is the right formulation |
| Errors cluster on the longest documents | Truncation is cutting the decisive text | Raise max_len, chunk and pool, or truncate from the other end |
| The model is confidently wrong on a label that a human also finds wrong | Annotation noise | Fix the labels. This is often the highest-value hour in the project |
| Errors cluster on one topic or vocabulary | Domain gap from pre-training | Domain-adaptive pre-training, or a domain-specific checkpoint |
Attention inspection
1@torch.no_grad()2def cls_attention(model, text, tokenizer, layer=-1, device=DEVICE):3 model.eval()4 # The default SDPA kernel never materialises attention weights, so switch5 # to the explicit ("eager") path for inspection; switch back to "sdpa" to serve.6 model.encoder.set_attn_implementation("eager")7 enc = tokenizer(text, truncation=True, max_length=128,8 return_tensors="pt").to(device)9 _, attentions = model(**enc, output_attentions=True)1011 # attentions[layer]: (batch, heads, seq, seq)12 w = attentions[layer][0] # (heads, seq, seq)13 cls_row = w[:, 0, :].mean(0) # averaged over heads, [CLS] row14 tokens = tokenizer.convert_ids_to_tokens(enc["input_ids"][0])15 return list(zip(tokens, cls_row.cpu().tolist()))1617for token, weight in cls_attention(model, "The plot was dull but the acting saved it.", tok):18 print(f"{token:>12} {weight:.4f} {'#' * int(weight * 100)}")The [CLS] attention row sums to 1 over the sequence, so you can read it as a distribution over tokens. On the example above a trained model typically concentrates on dull, but and saved — with but often carrying substantial weight, because it is the token that signals which clause wins.
Two cautions. Averaging over heads blurs twelve different patterns into one bland map; inspect individual heads when a map looks uninformative. And high attention weight shows that information flowed along an edge, not that the prediction depended on it — the value vector could be near zero, and later layers can discard whatever arrived. Attention maps generate hypotheses; they do not prove them.
Serving it
1class SentimentPredictor:2 def __init__(self, checkpoint="classifier.pt", device=DEVICE):3 ckpt = torch.load(checkpoint, map_location=device)4 self.cfg = Config(**ckpt["config"])5 self.tok = AutoTokenizer.from_pretrained(self.cfg.model_name)6 self.model = TransformerClassifier(self.cfg.model_name, num_labels=2)7 self.model.load_state_dict(ckpt["state_dict"])8 self.model.to(device).eval()9 self.device = device10 self.labels = ["negative", "positive"]1112 @torch.no_grad()13 def predict(self, texts, batch_size=32):14 if isinstance(texts, str):15 texts = [texts]16 results = []17 for i in range(0, len(texts), batch_size):18 enc = self.tok(texts[i:i + batch_size], truncation=True,19 max_length=self.cfg.max_len, padding=True,20 return_tensors="pt").to(self.device)21 probs = torch.softmax(self.model(**enc), dim=-1)22 conf, pred = probs.max(-1)23 results.extend(24 {"label": self.labels[p], "confidence": round(c, 4)}25 for p, c in zip(pred.cpu().tolist(), conf.cpu().tolist()))26 return results2728predictor = SentimentPredictor()29print(predictor.predict(["Utterly wonderful.", "I want those two hours back."]))Batching inside predict is not a nicety. Sending 1,000 texts one at a time through a GPU wastes almost all of it — the per-call overhead dominates, and throughput can be an order of magnitude worse than batched.
1from flask import Flask, request, jsonify23app = Flask(__name__)4predictor = SentimentPredictor()56@app.post("/predict")7def predict():8 payload = request.get_json(silent=True) or {}9 texts = payload.get("texts")10 if not texts:11 return jsonify({"error": "field 'texts' is required"}), 40012 if isinstance(texts, str):13 texts = [texts]14 if len(texts) > 128:15 return jsonify({"error": "at most 128 texts per request"}), 40016 return jsonify({"predictions": predictor.predict(texts)})1718@app.get("/health")19def health():20 return jsonify({"status": "ok", "device": DEVICE})The length cap matters. Without it one request with 100,000 texts occupies the process for minutes and every other caller times out.
Definition of done
| Stage | What must be true |
|---|---|
| Data | Class balance printed; majority-class baseline recorded; 20 examples read by a human; validation split taken from train, never from test |
| Length | Token-length percentiles measured; max_len chosen against them, not guessed |
| Wiring | 100 examples overfit to near-zero loss before the real run |
| Training | Warmup, gradient clipping, best-checkpoint restore all present |
| Evaluation | Macro-F1 and per-class recall reported, not accuracy alone; test set used exactly once |
| Reliability | Three seeds run; mean and spread reported |
| Analysis | 25 most-confident errors read; failure modes named |
| Serving | Batched inference; input size capped; model loaded once at startup |
Reference numbers on IMDB, so you know whether yours are reasonable:
| Model | Test accuracy | Parameters | Relative training time |
|---|---|---|---|
| Logistic regression on TF-IDF | ~0.88 | — | seconds |
| DistilBERT-base, 2 epochs | ~0.92 | 66 M | 0.5× |
| BERT-base, 3 epochs | ~0.93 | 110 M | 1× |
| RoBERTa-base, 3 epochs | ~0.95 | 125 M | 1× |
The first row is the one to take seriously. A TF-IDF baseline takes five minutes to build and reaches 0.88. If your transformer lands at 0.89, you have spent a GPU and a day for one point, and something in your pipeline is wrong — almost always truncation, a learning rate that is too high, or too many epochs.
What this means when you build something
The pattern here transfers to essentially every classification task you will meet, and the parts that transfer are not the model code.
Establish the floor before you build the ceiling. Majority-class accuracy and a TF-IDF baseline take ten minutes together and they define what your result means. Without them, 0.93 is a number with no interpretation.
Measure the input distribution before choosing hyperparameters. max_len, batch size and truncation direction all follow from token-length percentiles. Guessing 512 "to be safe" costs sixteen times the attention compute of 128 and often protects a handful of documents.
Spend an hour reading errors. Twenty-five confidently wrong predictions will tell you more about what to do next than any hyperparameter sweep. It routinely reveals label noise, and fixing labels is usually the cheapest available accuracy gain.
Use the test set once. Every time you look at test performance and change something, you leak a little information into your design choices. Tune on validation. Touch test at the end. If you tuned on test, your reported number is optimistic by an amount you cannot estimate.
When you extend this to a harder problem, the changes are small and predictable: more than two classes means changing num_labels and using macro-F1 (which you already are); severe imbalance means class weights in the loss or a resampled sampler; multi-label means BCEWithLogitsLoss and a per-label threshold instead of argmax; documents longer than 512 tokens mean chunking with overlap and pooling the chunk predictions, or a long-context model. The training loop, the evaluation discipline, and the error analysis stay exactly as they are.