Attention Mechanisms and Transformers

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.

95% accurate, and "urgent" is never predicted031701218012938pred urgentpred billingpred routinetrue urgenttrue billingtrue routineThe whole urgent row has drained into the other two columns; the empty diagonal cell is the finding.
Read the diagonal per row, not the total — a class that is 2% of the data can vanish without moving accuracy.

Setting up

Bash
python -m venv .venv && source .venv/bin/activatepip install "torch>=2.0" transformers datasets scikit-learn \            matplotlib seaborn pandas numpy accelerate
Python
import torch, transformers, numpy as np, randomprint("torch", torch.__version__, "| transformers", transformers.__version__)if torch.cuda.is_available():    print("GPU:", torch.cuda.get_device_name(0),          f"{torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")    DEVICE = "cuda"elif torch.backends.mps.is_available():    DEVICE = "mps"          # Apple siliconelse:    DEVICE = "cpu"print("device:", DEVICE)def set_seed(s=42):    random.seed(s); np.random.seed(s)    torch.manual_seed(s); torch.cuda.manual_seed_all(s)set_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
Weights440 MB
Gradients440 MB
Adam first moment440 MB
Adam second moment440 MB
Subtotal1.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

Python
from datasets import load_datasetimport pandas as pdds = load_dataset("stanfordnlp/imdb")         # 25k train / 25k test, balancedtrain_df = pd.DataFrame(ds["train"]).sample(frac=1, random_state=42)test_df  = pd.DataFrame(ds["test"])# carve a validation split OUT OF TRAIN — never out of testval_df   = train_df.iloc[:2500]train_df = train_df.iloc[2500:]print(len(train_df), len(val_df), len(test_df))   # 22500 2500 25000

For your own data, the only requirement is a text column and an integer label column:

Python
df = pd.read_csv("tickets.csv")               # columns: text, labellabel_names = sorted(df["label"].unique())label2id = {name: i for i, name in enumerate(label_names)}df["label_id"] = df["label"].map(label2id)

Look at the data before you model it

Three checks, and none of them is optional.

Python
from transformers import AutoTokenizerMODEL_NAME = "bert-base-uncased"tok = AutoTokenizer.from_pretrained(MODEL_NAME)# 1. Class balance -> sets the baseline you must beatcounts = train_df["label"].value_counts()print(counts)print("majority-class accuracy:", counts.max() / counts.sum())# 2. Token lengths -> sets max_lenlengths = [len(tok.encode(t, truncation=False)) for t in train_df["text"][:2000]]print("percentiles 50/90/95/99:", np.percentile(lengths, [50, 90, 95, 99]))# 3. Read twenty examples with their labels. Actually read them.for _, r in train_df.sample(20, random_state=0).iterrows():    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_lenReviews fully retainedAttention entries per headRelative attention cost
128~12%16,3841×
256~57%65,5364×
512~84%262,14416×

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

Python
from torch.utils.data import Dataset, DataLoaderfrom transformers import DataCollatorWithPaddingclass TextClassificationDataset(Dataset):    def __init__(self, texts, labels, tokenizer, max_len=256):        self.texts = list(texts)        self.labels = list(labels)        self.tok = tokenizer        self.max_len = max_len    def __len__(self):        return len(self.texts)    def __getitem__(self, i):        enc = self.tok(self.texts[i], truncation=True, max_length=self.max_len)        enc["labels"] = self.labels[i]        return enc                      # no padding here — the collator does itcollator = DataCollatorWithPadding(tok, return_tensors="pt")def make_loader(df, shuffle, bs=16, max_len=256):    dset = TextClassificationDataset(df["text"], df["label"], tok, max_len)    return DataLoader(dset, batch_size=bs, shuffle=shuffle,                      collate_fn=collator, num_workers=2, pin_memory=True)train_loader = make_loader(train_df, True)val_loader   = make_loader(val_df,   False)test_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.

Python
import torch.nn as nnfrom transformers import AutoModelclass TransformerClassifier(nn.Module):    def __init__(self, model_name=MODEL_NAME, num_labels=2, dropout=0.1):        super().__init__()        self.encoder = AutoModel.from_pretrained(model_name)        hidden = self.encoder.config.hidden_size        self.dropout = nn.Dropout(dropout)        self.classifier = nn.Linear(hidden, num_labels)    def forward(self, input_ids, attention_mask,                token_type_ids=None, output_attentions=False):        out = self.encoder(input_ids=input_ids,                           attention_mask=attention_mask,                           token_type_ids=token_type_ids,                           output_attentions=output_attentions)        cls = out.last_hidden_state[:, 0]        # the [CLS] position        logits = self.classifier(self.dropout(cls))        return (logits, out.attentions) if output_attentions else logits

The head is 768×2+2=1,538768 \times 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

Python
from dataclasses import dataclass@dataclassclass Config:    model_name: str = MODEL_NAME    max_len: int    = 256    batch_size: int = 16    lr: float       = 2e-5    epochs: int     = 3    warmup_ratio: float = 0.1    weight_decay: float = 0.01    max_grad_norm: float = 1.0    accum_steps: int = 1    seed: int = 42cfg = Config()
Python
import torch.nn.functional as Ffrom transformers import get_linear_schedule_with_warmupfrom sklearn.metrics import f1_score, accuracy_scorefrom torch.amp import autocast, GradScalerdef build_optimizer(model, cfg):    decay, no_decay = [], []    for name, p in model.named_parameters():        if not p.requires_grad:            continue        (no_decay if any(k in name for k in ("bias", "LayerNorm.weight"))         else decay).append(p)    return torch.optim.AdamW(        [{"params": decay,    "weight_decay": cfg.weight_decay},         {"params": no_decay, "weight_decay": 0.0}],        lr=cfg.lr, eps=1e-8)@torch.no_grad()def evaluate(model, loader, device):    model.eval()    preds, golds = [], []    for batch in loader:        batch = {k: v.to(device) for k, v in batch.items()}        labels = batch.pop("labels")        logits = model(**batch)        preds.extend(logits.argmax(-1).cpu().tolist())        golds.extend(labels.cpu().tolist())    return (accuracy_score(golds, preds),            f1_score(golds, preds, average="macro"),            np.array(golds), np.array(preds))def train(model, train_loader, val_loader, cfg, device=DEVICE):    model.to(device)    opt = build_optimizer(model, cfg)    total_steps = (len(train_loader) // cfg.accum_steps) * cfg.epochs    sched = get_linear_schedule_with_warmup(        opt, int(total_steps * cfg.warmup_ratio), total_steps)    scaler = GradScaler(device) if device == "cuda" else None    best_f1, best_state = 0.0, None    for epoch in range(1, cfg.epochs + 1):        model.train()        running = 0.0        opt.zero_grad(set_to_none=True)        for step, batch in enumerate(train_loader, 1):            batch = {k: v.to(device) for k, v in batch.items()}            labels = batch.pop("labels")            if scaler:                with autocast("cuda", dtype=torch.float16):                    loss = F.cross_entropy(model(**batch), labels)                scaler.scale(loss / cfg.accum_steps).backward()            else:                loss = F.cross_entropy(model(**batch), labels)                (loss / cfg.accum_steps).backward()            if step % cfg.accum_steps == 0:                if scaler:                    scaler.unscale_(opt)                       # BEFORE clipping                torch.nn.utils.clip_grad_norm_(model.parameters(),                                               cfg.max_grad_norm)                if scaler:                    scaler.step(opt); scaler.update()                else:                    opt.step()                sched.step()                opt.zero_grad(set_to_none=True)            running += loss.item()        acc, f1, _, _ = evaluate(model, val_loader, device)        print(f"epoch {epoch}  train_loss {running/len(train_loader):.4f}  "              f"val_acc {acc:.4f}  val_macro_f1 {f1:.4f}")        if f1 > best_f1:            best_f1 = f1            best_state = {k: v.detach().cpu().clone()                          for k, v in model.state_dict().items()}    model.load_state_dict(best_state)    return model, best_f1

Four lines in that loop are the ones that silently ruin runs if you get them wrong.

LineWhat breaks without it
scaler.unscale_(opt) before clippingGradients still scaled by ~65,536 look like outliers, so every batch is crushed to 1/655361/65536 and the model never moves
loss / cfg.accum_stepsGradients sum rather than average, giving an effective learning rate multiplied by accum_steps
sched.step() inside the accumulation guardThe schedule advances faster than the optimiser; warmup ends early
Restoring best_stateYou ship the last epoch, which on a small dataset is usually the most overfitted one
Python
set_seed(cfg.seed)model = TransformerClassifier(cfg.model_name, num_labels=2)model, best_f1 = train(model, train_loader, val_loader, cfg)test_acc, test_f1, y_true, y_pred = evaluate(model, test_loader, DEVICE)print(f"TEST  accuracy {test_acc:.4f}  macro-F1 {test_f1:.4f}")torch.save({"state_dict": model.state_dict(), "config": cfg.__dict__},           "classifier.pt")

A healthy run on IMDB with these settings looks roughly like:

Text
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.9326

Note 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

Python
from sklearn.metrics import confusion_matrix, classification_reportimport seaborn as sns, matplotlib.pyplot as pltcm = confusion_matrix(y_true, y_pred)sns.heatmap(cm, annot=True, fmt="d", cmap="Blues",            xticklabels=["neg", "pos"], yticklabels=["neg", "pos"])plt.xlabel("predicted"); plt.ylabel("true"); plt.tight_layout()print(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 negpredicted posrow total
true neg44060500
true pos45455500

Work every metric out by hand once, because the shape of these formulas is what tells you which one to care about.

accuracy=440+4551000=0.895\text{accuracy} = \frac{440 + 455}{1000} = 0.895

For the positive class: 455 true positives, 60 false positives, 45 false negatives.

Ppos=455455+60=455515=0.8835,Rpos=455455+45=455500=0.910P_{pos} = \frac{455}{455 + 60} = \frac{455}{515} = 0.8835, \qquad R_{pos} = \frac{455}{455+45} = \frac{455}{500} = 0.910
F1pos=2⋅0.8835⋅0.9100.8835+0.910=1.60801.7935=0.8966F1_{pos} = \frac{2 \cdot 0.8835 \cdot 0.910}{0.8835 + 0.910} = \frac{1.6080}{1.7935} = 0.8966

For the negative class: 440 true, 45 false positives, 60 false negatives.

Pneg=440485=0.9072,Rneg=440500=0.880,F1neg=2⋅0.9072⋅0.8801.7872=0.8934P_{neg} = \frac{440}{485} = 0.9072, \qquad R_{neg} = \frac{440}{500} = 0.880, \qquad F1_{neg} = \frac{2\cdot0.9072\cdot0.880}{1.7872} = 0.8934

macro-F1=0.8966+0.89342=0.8950\text{macro-F1} = \frac{0.8966 + 0.8934}{2} = 0.8950

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:

accuracy=9501000=0.950\text{accuracy} = \frac{950}{1000} = 0.950

but Rpos=0/50=0R_{pos} = 0/50 = 0, so F1pos=0F1_{pos} = 0, while Pneg=950/1000=0.95P_{neg} = 950/1000 = 0.95 and Rneg=1.0R_{neg} = 1.0 give F1neg=2(0.95)(1)/1.95=0.9744F1_{neg} = 2(0.95)(1)/1.95 = 0.9744. Therefore

macro-F1=0.9744+02=0.4872\text{macro-F1} = \frac{0.9744 + 0}{2} = 0.4872

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.

Python
@torch.no_grad()def collect_errors(model, loader, texts, device=DEVICE, k=25):    model.eval()    rows, offset = [], 0    for batch in loader:        batch = {k_: v.to(device) for k_, v in batch.items()}        labels = batch.pop("labels")        probs = torch.softmax(model(**batch), dim=-1)        conf, pred = probs.max(-1)        for j in range(len(labels)):            if pred[j] != labels[j]:                rows.append({                    "text": texts[offset + j][:300],                    "true": labels[j].item(),                    "pred": pred[j].item(),                    "confidence": round(conf[j].item(), 4),                })        offset += len(labels)    # most confident mistakes first — these are the most informative    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 findWhat it meansWhat to do
Sarcasm and irony ("brilliant, another two hours of my life gone")Surface sentiment words contradict the meaningA 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 bothReconsider whether one label per document is the right formulation
Errors cluster on the longest documentsTruncation is cutting the decisive textRaise max_len, chunk and pool, or truncate from the other end
The model is confidently wrong on a label that a human also finds wrongAnnotation noiseFix the labels. This is often the highest-value hour in the project
Errors cluster on one topic or vocabularyDomain gap from pre-trainingDomain-adaptive pre-training, or a domain-specific checkpoint

Attention inspection

Python
@torch.no_grad()def cls_attention(model, text, tokenizer, layer=-1, device=DEVICE):    model.eval()    # The default SDPA kernel never materialises attention weights, so switch    # to the explicit ("eager") path for inspection; switch back to "sdpa" to serve.    model.encoder.set_attn_implementation("eager")    enc = tokenizer(text, truncation=True, max_length=128,                    return_tensors="pt").to(device)    _, attentions = model(**enc, output_attentions=True)    # attentions[layer]: (batch, heads, seq, seq)    w = attentions[layer][0]                 # (heads, seq, seq)    cls_row = w[:, 0, :].mean(0)             # averaged over heads, [CLS] row    tokens = tokenizer.convert_ids_to_tokens(enc["input_ids"][0])    return list(zip(tokens, cls_row.cpu().tolist()))for token, weight in cls_attention(model, "The plot was dull but the acting saved it.", tok):    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

Python
class SentimentPredictor:    def __init__(self, checkpoint="classifier.pt", device=DEVICE):        ckpt = torch.load(checkpoint, map_location=device)        self.cfg = Config(**ckpt["config"])        self.tok = AutoTokenizer.from_pretrained(self.cfg.model_name)        self.model = TransformerClassifier(self.cfg.model_name, num_labels=2)        self.model.load_state_dict(ckpt["state_dict"])        self.model.to(device).eval()        self.device = device        self.labels = ["negative", "positive"]    @torch.no_grad()    def predict(self, texts, batch_size=32):        if isinstance(texts, str):            texts = [texts]        results = []        for i in range(0, len(texts), batch_size):            enc = self.tok(texts[i:i + batch_size], truncation=True,                           max_length=self.cfg.max_len, padding=True,                           return_tensors="pt").to(self.device)            probs = torch.softmax(self.model(**enc), dim=-1)            conf, pred = probs.max(-1)            results.extend(                {"label": self.labels[p], "confidence": round(c, 4)}                for p, c in zip(pred.cpu().tolist(), conf.cpu().tolist()))        return resultspredictor = SentimentPredictor()print(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.

Python
from flask import Flask, request, jsonifyapp = Flask(__name__)predictor = SentimentPredictor()@app.post("/predict")def predict():    payload = request.get_json(silent=True) or {}    texts = payload.get("texts")    if not texts:        return jsonify({"error": "field 'texts' is required"}), 400    if isinstance(texts, str):        texts = [texts]    if len(texts) > 128:        return jsonify({"error": "at most 128 texts per request"}), 400    return jsonify({"predictions": predictor.predict(texts)})@app.get("/health")def health():    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

StageWhat must be true
DataClass balance printed; majority-class baseline recorded; 20 examples read by a human; validation split taken from train, never from test
LengthToken-length percentiles measured; max_len chosen against them, not guessed
Wiring100 examples overfit to near-zero loss before the real run
TrainingWarmup, gradient clipping, best-checkpoint restore all present
EvaluationMacro-F1 and per-class recall reported, not accuracy alone; test set used exactly once
ReliabilityThree seeds run; mean and spread reported
Analysis25 most-confident errors read; failure modes named
ServingBatched inference; input size capped; model loaded once at startup

Reference numbers on IMDB, so you know whether yours are reasonable:

ModelTest accuracyParametersRelative training time
Logistic regression on TF-IDF~0.88—seconds
DistilBERT-base, 2 epochs~0.9266 M0.5×
BERT-base, 3 epochs~0.93110 M1×
RoBERTa-base, 3 epochs~0.95125 M1×

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.