The Hugging Face Trainer runs the training loop so you do not have to: batching, mixed precision, gradient accumulation, distributed training, checkpointing, logging and evaluation. Most people meet it through a tutorial where everything works. Then they need a weighted loss, an extra input column, a token-level metric on a 7B model or a contrastive objective, and the Trainer starts to behave strangely: columns vanish, the loss changes when gradient accumulation changes, and evaluation runs out of memory at the very end.

Those problems all live in the Trainer's customization surface, which is what this article covers. The overall architecture, TrainingArguments, callbacks and resuming are covered in Hugging Face Trainer architecture. Here we follow one batch from dataset row to loss and one evaluation pass from logits to metrics, and look at what you can change at each step. Signatures were checked against the current Trainer documentation and source; check them against your installed version too, because this API keeps evolving.

Advertisement

The path of one batch

In a training step the Trainer pulls rows from the dataset, removes columns the model cannot accept, passes the list of rows to the data collator to build tensors, counts the label tokens across the whole accumulated batch, calls compute_loss from inside training_step, scales the loss for gradient accumulation and calls backward. Every few micro-batches the optimizer and scheduler step. In evaluation, prediction_step returns the loss, logits and labels for each batch, these are accumulated, and compute_metrics sees them all at the end.

Each stage has a lightweight hook and a heavier override. The rule of thumb is to use the lightest one that solves your problem: a constructor argument before a callback, a callback before a subclass, and overriding a narrow method such as compute_loss before a broad one such as training_step.

Where your code plugs into one Trainer step (training above, evaluation below)dataset rowafter column pruningdata_collatorpad, mask, labelsget_batch_samplescount num_items_in_batchcompute_lossor compute_loss_functraining_stepscale for accumulation, backwardoptimizer + schedulerevery N micro-batcheseval batchsame collatorprediction_steploss, logits, labelspreprocess_logitsshrink before cachingaccumulatedevice, then CPUcompute_metrics(EvalPrediction)numpy arrays, whole eval setYellow boxes are the customization points. Each has a cheaper and a more invasive way in.
One training step and one evaluation pass through the Trainer. The yellow boxes are the customization points discussed in this article.

What actually reaches the model: column pruning, labels and collators

By default remove_unused_columns=True. The Trainer inspects the signature of the model's forward method and drops dataset columns that do not match a parameter name. It logs the dropped columns at info level, which most people never see. If your custom loss needs a sample_weight column and the model's forward does not take one, it is gone before your collator runs. Set remove_unused_columns=False and make your collator responsible for deciding what goes into the batch, or wrap the model so that forward accepts the extra field.

label_names tells the Trainer which input keys are labels; by default it looks for keys containing label. This matters in evaluation, where the Trainer separates labels from inputs, and for question-answering models whose labels are start_positions and end_positions. If the Trainer cannot find labels, compute_metrics receives none and the evaluation loss may be missing.

The collator turns a list of examples into a batch. For causal language modelling, padding tokens must be excluded from the loss by setting their labels to -100, the ignore index used by PyTorch's cross-entropy and by the Trainer's own token counting. A collator that pads labels with the pad token ID instead trains the model to predict padding.

The fastest way to see what the model will actually receive is to ask the Trainer itself. Build it, call next(iter(trainer.get_train_dataloader())), and print every key with its shape and dtype, plus the share of label positions equal to -100. That one batch shows pruned columns, wrong padding and missing labels before you spend an hour of GPU time discovering them through a flat loss curve.

from dataclasses import dataclass
import torch

@dataclass
class WeightedCollator:
    tokenizer: object
    def __call__(self, rows):
        batch = self.tokenizer.pad(
            [{"input_ids": r["input_ids"], "attention_mask": r["attention_mask"]} for r in rows],
            return_tensors="pt")
        batch["labels"] = torch.tensor([r["label"] for r in rows])
        batch["sample_weight"] = torch.tensor([r["weight"] for r in rows], dtype=torch.float)
        return batch

# TrainingArguments(..., remove_unused_columns=False) so "weight" survives to the collator
Advertisement

Three ways to write a custom loss

The least invasive way is to let the model compute its own loss: if the batch contains labels and the model supports them, outputs.loss is used. The next step up is the compute_loss_func constructor argument, a function that receives the raw model outputs, the labels and num_items_in_batch. When it is set, the Trainer removes labels from the inputs before calling the model, so the model does not compute a loss you would throw away. For causal LM labels this means your function must do the one-position shift itself.

The most flexible way is to subclass and override compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None). You get the whole input dictionary, including extra fields such as the sample weights above.

import torch.nn.functional as F
from transformers import Trainer

class WeightedTrainer(Trainer):
    def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
        labels = inputs.pop("labels")
        weights = inputs.pop("sample_weight")
        outputs = model(**inputs)
        per_example = F.cross_entropy(outputs.logits, labels, reduction="none")
        weighted = (per_example * weights).sum()
        if num_items_in_batch is not None:      # Trainer will NOT divide by accumulation steps
            loss = weighted / num_items_in_batch
        else:                                   # Trainer divides a per-batch mean itself
            loss = weighted / weights.sum()
        return (loss, outputs) if return_outputs else loss

Return (loss, outputs) when return_outputs is true, because evaluation calls the same method and needs the logits. The two branches are not optional. When the Trainer passes num_items_in_batch, which it does whenever the batch has labels and the model's forward accepts loss keyword arguments, it assumes your loss is already normalised over the whole accumulated batch and does not divide by the accumulation steps. A per-batch mean returned in that case makes gradients roughly accumulation-steps times too large. The next section explains why.

Gradient accumulation and num_items_in_batch

Gradient accumulation runs several micro-batches and steps the optimizer once, to simulate a larger batch. The naive implementation computes a mean loss per micro-batch and divides each by the number of micro-batches. That equals the true mean only if every micro-batch has the same number of items. For language modelling the items are non-ignored label tokens, and their count varies from batch to batch.

A worked example makes the error concrete. Micro-batch A has 100 label tokens with a summed loss of 200; micro-batch B has 900 tokens with a summed loss of 900. The true mean over 1,000 tokens is 1,100 / 1,000 = 1.1. The mean of per-batch means is (2.0 + 1.0) / 2 = 1.5, which weights each of A's tokens nine times more than each of B's. The model trains on a different objective depending on how you split the batch, and loss curves stop being comparable across accumulation settings.

The Trainer fixes this by looking ahead. Before running the micro-batches of one optimizer step it fetches them all, counts labels not equal to -100 across them (skipping the first position when the loss shifts labels), and passes that count as num_items_in_batch. A model whose forward accepts loss keyword arguments, or a compute_loss_func, can then return the summed loss divided by the global count, and the Trainer will not divide again by the accumulation steps. When the count is not passed, whatever loss you return is divided by the accumulation steps. The decision comes from whether the count was passed, not from inspecting your loss. Under distributed training, the average_tokens_across_devices argument gathers the count across ranks so that every token counts equally.

def token_ce(outputs, labels, num_items_in_batch=None):
    logits = outputs.logits[..., :-1, :].contiguous()      # predict token t+1 from t
    labels = labels[..., 1:].contiguous()
    loss = F.cross_entropy(logits.view(-1, logits.size(-1)), labels.view(-1),
                           ignore_index=-100, reduction="sum")
    if num_items_in_batch is None:                          # e.g. evaluation
        return loss / (labels != -100).sum()
    return loss / num_items_in_batch

trainer = Trainer(model=model, args=args, train_dataset=train, eval_dataset=dev,
                  data_collator=collator, compute_loss_func=token_ce)

If you override compute_loss instead, follow the same rule: use reduction='sum' and divide by num_items_in_batch when it is provided. Recent source also exposes a loss_is_scaled_for_ga class attribute so a subclass can say which kind of loss it returns; check whether your installed version has it before relying on it. The quickest test is to train for a few steps with accumulation 1 and batch size 8, then accumulation 4 and batch size 2, and compare the logged losses; how to read them is covered in transformer training loss curves.

Evaluation: prediction_step, EvalPrediction and compute_metrics

In evaluation the Trainer calls prediction_step(model, inputs, prediction_loss_only, ignore_keys=None) for every batch. It returns the loss, the logits and the labels. Those are concatenated across batches, padded where sequence lengths differ, gathered across processes and handed to compute_metrics as an EvalPrediction whose predictions and label_ids are NumPy arrays covering the whole evaluation set.

That design is convenient for classification and dangerous for language models. Keeping full logits for 2,000 evaluation sequences of 512 tokens with a 32,000-token vocabulary in float32 needs 2,000 × 512 × 32,000 × 4 bytes, about 131 GB. By default the accumulated tensors stay on the accelerator until the loop ends, so evaluation dies on the last batch after running for an hour.

Two settings fix this. preprocess_logits_for_metrics(logits, labels) runs on each batch before caching and can reduce the logits to what the metric needs, usually the argmax token IDs, which is 32,000 times smaller. eval_accumulation_steps moves the accumulated outputs to the CPU every N steps instead of holding everything on the device. Use the first always for token-level tasks and the second when even the reduced outputs are large.

import numpy as np

def keep_argmax(logits, labels):
    if isinstance(logits, tuple):           # some models return extra tensors
        logits = logits[0]
    return logits.argmax(dim=-1)

def token_accuracy(p):
    preds = p.predictions[:, :-1]           # same shift as the loss
    labels = p.label_ids[:, 1:]
    mask = labels != -100
    return {"token_acc": float((preds[mask] == labels[mask]).mean())}

trainer = Trainer(..., compute_metrics=token_accuracy,
                  preprocess_logits_for_metrics=keep_argmax)

Metric names are prefixed with eval_ in the logs, so metric_for_best_model can refer to token_acc. For generation metrics such as BLEU or ROUGE, the base Trainer's argmax over teacher-forced logits is not generation; use Seq2SeqTrainer with generation enabled, or run generation in a callback on a sample.

Other override points, and when to stop

A few other methods are commonly overridden. get_train_dataloader lets you supply a custom sampler, for example length-grouped or class-balanced sampling, though check first whether a TrainingArguments flag already does it. create_optimizer controls parameter groups, such as a lower learning rate for the backbone; the optimizer_cls_and_kwargs constructor argument covers the simpler case of choosing an optimizer class. training_step is the broadest override and the riskiest, because it also handles mixed precision, accumulation scaling and distributed backends that you would have to reproduce.

If you find yourself overriding training_step and the dataloader and the evaluation loop, the Trainer is no longer saving you work. A hand-written loop on Hugging Face Accelerate gives the same device and distributed handling with a loop you fully control. Data preparation belongs in the datasets library, where it is cached, not in the collator.

Failure modes

SymptomLikely causeFix
KeyError for a custom field in compute_lossremove_unused_columns dropped itSet it to False and build the batch in the collator
Loss changes with accumulation stepsPer-batch mean loss on token dataSum and divide by num_items_in_batch
Gradients too large with accumulationCustom loss ignores a passed num_items_in_batchDivide by num_items_in_batch when it is given
Model learns to emit paddingLabels padded with the pad IDPad labels with -100
OOM at the end of evaluationFull logits accumulated on devicepreprocess_logits_for_metrics, eval_accumulation_steps
compute_metrics gets no labelslabel_names does not match the batch keysSet label_names explicitly
Evaluation crashes in compute_lossOverride ignores return_outputsReturn (loss, outputs) when asked
Metric off by one tokenShift applied in loss but not in metricApply the same shift in both

What to do next

  1. Print one collated batch and check every key, shape and the -100 positions before training.
  2. Decide whether you need extra columns; if so set remove_unused_columns to False and own the collator.
  3. Pick the lightest loss hook: model loss, compute_loss_func, then a compute_loss override.
  4. For token-level losses, sum and divide by num_items_in_batch, then verify with two accumulation settings.
  5. Add preprocess_logits_for_metrics to every language-model evaluation and set eval_accumulation_steps if outputs are large.
  6. Make compute_metrics apply the same shift and masking as the loss.
  7. Check the signatures above against your installed transformers version.
Key takeaway: Customize the Trainer at the lightest point that works. Know that remove_unused_columns prunes inputs to the model's forward signature, that labels must be padded with -100, and that label_names controls what evaluation sees. Write custom losses with compute_loss_func or a compute_loss override that returns outputs when asked. For token-level losses, sum and divide by num_items_in_batch so gradient accumulation does not change the objective, and keep evaluation memory bounded with preprocess_logits_for_metrics and eval_accumulation_steps.