Skip to main content

Trainer overview

Trainer is the batteries-included training loop for PyTorch Transformers models.

Adapted for this playbook from the 🤗 Transformers documentation by Hugging Face. Official page: https://huggingface.co/docs/transformers/trainer. Images © Hugging Face (documentation-images) unless noted. This is not a substitute for the upstream docs — verify against the current version.

Covers Hugging Face pages: trainer, trainer_customize, trainer_callbacks, data_collators, optimizers, hpo_train, trainer_recipes

Minimal shape​

from transformers import AutoModelForSequenceClassification, AutoTokenizer, Trainer, TrainingArguments
from datasets import load_dataset

model_id = "bert-base-uncased"
tok = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForSequenceClassification.from_pretrained(model_id, num_labels=2)

# prepare tokenized datasets ...
args = TrainingArguments(
output_dir="./out",
learning_rate=2e-5,
per_device_train_batch_size=16,
num_train_epochs=3,
eval_strategy="epoch",
fp16=True,
)
trainer = Trainer(model=model, args=args, train_dataset=train_ds, eval_dataset=eval_ds, processing_class=tok)
trainer.train()

Customisation surfaces​

HookDoc
Subclass methodstrainer_customize
Callbackstrainer_callbacks
Data collatorsdata_collators
Optimisers / schedulersoptimizers
Hyperparameter searchhpo_train

Training granularities illustration — Source: Hugging Face

Source: Hugging Face Transformers documentation.

Discussion

Comments​

Share feedback or questions about this page. No account required.

Loading comments…