The Trainer API
Trainer is the training loop already written — you hand it a model, a dataset and a settings object, and it runs the epochs, evaluation, logging and checkpointing for you.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Trainer is a ready-made training loop. You supply the model, the data and the settings; it does the repeating, checking and saving.
Think of a pressure cooker with a preset for dal. You add the lentils and water, pick the setting, and walk away. The timing, the heat and the whistle are already decided by someone who cooked dal a thousand times. You could do it in an open pot, and sometimes you should — but not on a Tuesday.
Training a model means showing it examples over and over so it improves. The loop that does the showing is nearly identical in every project. Trainer is that loop, written once.
Why it exists
Writing your own loop is a good exercise and a bad habit. The loop itself is short. What surrounds it is not: evaluating at the right moments, logging numbers, saving checkpoints, resuming after a crash, handling GPUs.
Each of those is a place to introduce a silent bug. A model left in the wrong mode during evaluation. Gradients not cleared. A metric computed on the wrong split. Every one of them produces plausible numbers and a wrong conclusion.
Trainer removes that whole class of mistake by having the boring parts already correct.
How it works
model ─┐
data ─┼─→ Trainer → for each epoch:
settings┘ show batches, adjust the model
check the score on held-out data
write the numbers to the screen
save a copy if askedThree things go in. A model with the right head. A dataset already turned into numbers. And a settings object — how many passes, how big the batches, how fast to learn.
One thing comes out: a trained model, plus a log of how the score moved.
A real example you have seen
Exam preparation with a mock-test schedule. Study a chapter, take a mock test, look at the score, continue. The studying is yours; the schedule and the marking are somebody else's system. Trainer is that schedule.
Remember this
- Trainer is the standard training loop, already written and debugged.
- You give it a model, a training set, an evaluation set and a settings object.
- It reports a score on held-out data as it goes, so you can see learning happen.
What to learn next
- Accelerate — the layer under Trainer, and the way out when you need your own loop.
- Experiment tracking — where those logged numbers should end up on a real project.
- Fine-tuning — the beginner-level story of why any of this is worth doing.
Developer — Code and libraries.
Setup
pip install transformers torch datasetsTested with transformers 5.6, torch 2.5 and datasets 5.0. The model is google/bert_uncased_L-2_H-128_A-2 — BERT-tiny, 4.4M parameters, about 18 MB of files. The whole run finishes in about a second on a laptop CPU.
A complete fine-tune, start to finish
import numpy as np
from datasets import Dataset
from transformers import (AutoTokenizer, AutoModelForSequenceClassification,
DataCollatorWithPadding, Trainer, TrainingArguments, set_seed)
set_seed(0)
NAME = "google/bert_uncased_L-2_H-128_A-2"
tok = AutoTokenizer.from_pretrained(NAME)
dishes = ["biryani", "dosa", "chai", "vada pav", "filter coffee", "paneer tikka", "idli", "samosa"]
praise = ["was delicious", "was fresh and hot", "was perfect", "tasted wonderful",
"was excellent", "made my day"]
insult = ["was cold", "was stale", "was awful", "tasted terrible",
"was disgusting", "ruined my evening"]
rows = [{"text": f"the {d} {p}", "label": 1} for d in dishes for p in praise]
rows += [{"text": f"the {d} {p}", "label": 0} for d in dishes for p in insult]
split = Dataset.from_list(rows).shuffle(seed=0).train_test_split(test_size=0.25, seed=0)
enc = split.map(lambda b: tok(b["text"], truncation=True, max_length=16),
batched=True, remove_columns=["text"])
print("train:", enc["train"].num_rows, " test:", enc["test"].num_rows)
model = AutoModelForSequenceClassification.from_pretrained(NAME, num_labels=2)
def compute_metrics(pred):
return {"accuracy": float((np.argmax(pred.predictions, -1) == pred.label_ids).mean())}
args = TrainingArguments(
output_dir="tiny-run",
num_train_epochs=5,
per_device_train_batch_size=8,
learning_rate=5e-4, # high, because the head starts random and the model is tiny
eval_strategy="epoch",
logging_strategy="epoch",
save_strategy="no",
report_to=[], # no experiment tracker for a toy run
seed=0,
disable_tqdm=True,
)
trainer = Trainer(model=model, args=args,
train_dataset=enc["train"], eval_dataset=enc["test"],
data_collator=DataCollatorWithPadding(tok),
compute_metrics=compute_metrics)
trainer.train()
print(trainer.evaluate())train: 72 test: 24
[transformers] BertForSequenceClassification LOAD REPORT from: google/bert_uncased_L-2_H-128_A-2
Key | Status |
-------------------------------------------+------------+-
cls.predictions.transform.dense.weight | UNEXPECTED |
cls.seq_relationship.weight | UNEXPECTED |
classifier.bias | MISSING |
classifier.weight | MISSING |
{'loss': '0.6929', 'grad_norm': '1.362', 'learning_rate': '0.0004111', 'epoch': '1'}
{'eval_loss': '0.7406', 'eval_accuracy': '0.3333', 'epoch': '1'}
{'loss': '0.6682', 'grad_norm': '3.408', 'learning_rate': '0.0003111', 'epoch': '2'}
{'eval_loss': '0.6582', 'eval_accuracy': '0.7917', 'epoch': '2'}
{'loss': '0.6008', 'grad_norm': '2.392', 'learning_rate': '0.0002111', 'epoch': '3'}
{'eval_loss': '0.5171', 'eval_accuracy': '0.9167', 'epoch': '3'}
{'loss': '0.4421', 'grad_norm': '3.424', 'learning_rate': '0.0001111', 'epoch': '4'}
{'eval_loss': '0.3376', 'eval_accuracy': '1', 'epoch': '4'}
{'loss': '0.3465', 'grad_norm': '3.188', 'learning_rate': '1.111e-05', 'epoch': '5'}
{'eval_loss': '0.289', 'eval_accuracy': '1', 'epoch': '5'}
{'train_runtime': '0.6563', 'train_samples_per_second': '548.5', 'epoch': '5'}
{'eval_loss': 0.28904154896736145, 'eval_accuracy': 1.0, 'epoch': 5.0}The real load report lists a few more UNEXPECTED rows and each line carries timing fields; both are trimmed above. Exact losses shift a little with your CPU and library build.
About that 1.0 accuracy. The sentences were built from two word lists, so "delicious" versus "disgusting" separates them completely. A perfect score here means the task is trivial, not that the model is good. Read it as proof the wiring works — nothing more.
The walkthrough
The load report is the same one from lesson one. classifier.weight | MISSING says the head arrived random, which is exactly why we are training. UNEXPECTED rows are the pretraining heads we do not want, discarded on the way in. See which AutoModel class to use.
Watch loss fall and accuracy rise together. Training loss went 0.69 → 0.35, evaluation accuracy 0.33 → 1.0. Epoch 1's accuracy below 0.5 is the random head guessing worse than a coin. When these two numbers disagree — loss falling while accuracy sits still — overfitting and underfitting is the next place to look.
compute_metrics receives raw scores, not predictions. pred.predictions holds the logits, one row per example, one column per label. argmax turns them into a chosen class. Without this function, Trainer reports loss and nothing else — which tells you little about whether the model is useful.
The column must be labels by the time the model sees it. We named it label, and Trainer renamed it automatically because that is a known alias. Call it target and the model receives no labels, computes no loss, and training fails with a confusing error about a missing loss.
eval_strategy="epoch" is the one setting to never skip. Without it you learn the final score and nothing about the journey. With it, overfitting is visible: evaluation loss turning upward while training loss keeps falling. Add load_best_model_at_end=True with save_strategy="epoch" and Trainer keeps the best checkpoint rather than the last.
report_to=[] keeps the run local. By default Trainer logs to any tracker it finds installed. That is genuinely useful on real projects — see experiment tracking — and pure noise on a toy one.
Common mistakes
A learning rate copied from the wrong place. 5e-4 suits a 4-million-parameter model with a random head. Full fine-tuning of a base-size encoder wants 2e-5 to 5e-5. Off by a factor of ten upward and the loss goes flat or explodes; downward and nothing moves. Start from the published value for your model size — gradient descent explains what the step size is doing.
Leaving a string column in the dataset. Trainer drops columns the model's forward method does not accept, which is usually helpful and occasionally hides a mistake. Being explicit with remove_columns at map time is better than trusting the guess.
Evaluating on the data you trained on. The score will be excellent and meaningless. Split first, always, and read train/test split if the reason is not yet in your bones.
Filling the disk with checkpoints. Default saving keeps several full copies of the model plus optimizer state, several times the model's size each. Set save_total_limit=2 on anything larger than a toy.
Reaching for Trainer when the loop is the point. Trainer assumes a standard supervised setup. Multiple models, adversarial objectives, or custom gradient flow fight the abstraction. Subclass and override compute_loss for small deviations; write a plain loop with Accelerate for large ones.
Try it yourself
Set num_train_epochs=30 and learning_rate=5e-3 and watch what a too-large step does to the loss column. Then add a real evaluation split by holding out one whole dish — train on seven dishes, test on "samosa" only — and see whether the accuracy survives. It is a much fairer test than the random split.
What to learn next
- Accelerate — the layer under Trainer, and the way out when you need your own loop.
- Experiment tracking — where those logged numbers should end up on a real project.
- Fine-tuning — the beginner-level story of why any of this is worth doing.
Researcher — Mathematics and papers.
What the abstraction actually contains
Trainer.train wraps a standard loop with a large amount of correctness bookkeeping: model.train() / model.eval() toggling, optimizer.zero_grad() placement, gradient accumulation with correct loss scaling, gradient clipping, learning-rate scheduling per optimizer step rather than per batch, mixed-precision autocast plus loss scaling, distributed wrapping, and deterministic dataloader seeding for resumption. Defaults: AdamW (Loshchilov and Hutter, 2019) with linear decay and no warmup, weight_decay=0.0, and no decay applied to bias and LayerNorm parameters — the convention set by the original BERT code.
Device placement and distribution are delegated to Accelerate, which is why TrainingArguments exposes DeepSpeed and FSDP configuration without Trainer implementing either.
Extension points, in increasing order of intrusiveness
compute_metrics(EvalPrediction)for metrics. Note thatpredictionsare gathered across all processes and held in memory; on generative tasks this can dominate evaluation memory, whicheval_accumulation_stepsandpreprocess_logits_for_metricsexist to bound.TrainerCallbackfor behaviour that observes but does not change the computation: early stopping, custom logging, learning-rate probes.- Subclassing
compute_lossfor a different objective on the same loop — label smoothing, class weighting, auxiliary losses, knowledge distillation. optimizers=(optimizer, scheduler)to bypass the defaults entirely.
Beyond that the abstraction is fighting you, and the honest move is a plain loop.
Reproducibility, and its limits
seed fixes Python, NumPy and torch generators plus the dataloader sampler; data_seed separates data order from initialisation, which matters because Dodge et al. (2020) show fine-tuning variance on small datasets is dominated by exactly those two factors. Reporting a single seed's result on a small benchmark is not a measurement — report the spread over several. full_determinism=True additionally forces deterministic kernels at a real throughput cost, and even then results are only reproducible on identical hardware and library versions, for the floating-point reasons in random seeds and reproducibility.
Where Trainer sits among the alternatives
TRL's SFTTrainer, DPOTrainer and GRPOTrainer subclass Trainer, adding data formatting and objective changes for post-training while inheriting the same loop. PyTorch Lightning solves the same problem with a different decomposition. The comparison worth making is not features but blast radius: a bug in your own loop is yours to find, while a bug in a shared loop is found by thousands of users first.
What to learn next
- Accelerate — the layer under Trainer, and the way out when you need your own loop.
- Experiment tracking — where those logged numbers should end up on a real project.
- Fine-tuning — the beginner-level story of why any of this is worth doing.