Building Models with nn.Module

Writing a custom loss function

A loss function in PyTorch is any expression built from tensor operations — write one when the built-ins reward the wrong thing, and verify it against a built-in before trusting it.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A custom loss function is your own scoring rule for the model's mistakes, written so that training can push against it.

Think about how a marking scheme shapes students. If an exam gives nine marks for theory and one for practicals, students memorise theory and skip the lab. Change the marking scheme and you change what everyone works on — without saying a word to any student.

A loss function is the marking scheme for a model. Training does one thing only: adjust the model to score better on it. So the loss is the instruction. Write a new one, and the model chases a new goal.

Why it exists

The standard losses — covered in loss functions — treat every mistake as equally important. Often that is wrong for your problem.

A fraud model that misses real fraud causes far more damage than one that flags a good payment for review. A demand forecast that runs out of stock may cost more than one that overstocks. The standard marking scheme does not know this. Your business does.

So you fold the real cost into the score. Mistakes that hurt more, cost more marks. The model, chasing marks as always, starts avoiding the expensive mistakes first.

How it works

prediction, truth
      |
      v
[ your scoring rule ]      e.g. "wrong on fraud counts 3x"
      |
      v
one number: the loss
      |
      v
training pushes every knob downhill against YOUR rule

One requirement makes it work: the rule must be written from smooth building blocks, so training can feel which direction improves it. PyTorch's tensor operations are those blocks.

A real example you have seen

Spam filters err on the side of letting borderline mail through, because deleting a real email angers you more than showing one spam. That asymmetry was engineered — someone encoded "these two mistakes are not equal" into the score being minimised.

Remember this

  • The loss is the marking scheme; the model becomes whatever the scheme rewards.
  • Custom losses encode your costs when standard ones treat all errors alike.
  • Built from PyTorch operations, a custom loss trains with no extra machinery.

What to learn next

  • Transfer learning — pointing a pretrained model at your problem, custom loss included.
  • Loss functions — the standard schemes worth checking yours against.
  • RLHF — what happens when the objective cannot be written down at all.

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on CPU. All inputs are fixed, so this output is exact.

Two custom losses, and a correctness check

custom_losses.py
import torch
from torch import nn
import torch.nn.functional as F

torch.manual_seed(0)

def weighted_mse(pred, target, weights):
    """MSE where some samples matter more than others."""
    return (weights * (pred - target) ** 2).mean()

class Huberish(nn.Module):
    """Squared error for small mistakes, linear for big ones."""
    def __init__(self, threshold=1.0):
        super().__init__()
        self.threshold = threshold

    def forward(self, pred, target):
        error = (pred - target).abs()
        squared = 0.5 * error ** 2
        linear = self.threshold * (error - 0.5 * self.threshold)
        return torch.where(error < self.threshold, squared, linear).mean()

pred = torch.tensor([2.5, 0.0, 8.0], requires_grad=True)
target = torch.tensor([3.0, 0.5, 2.0])

loss = weighted_mse(pred, target, torch.tensor([1.0, 1.0, 3.0]))
loss.backward()
print(f"weighted mse: {loss.item():.4f}")
print("grad on pred:", [round(g, 3) for g in pred.grad.tolist()])

mine = Huberish()(pred.detach(), target)
theirs = F.huber_loss(pred.detach(), target, delta=1.0)
print(f"my huber {mine.item():.4f}  vs torch's {theirs.item():.4f}")
Output
weighted mse: 36.1667
grad on pred: [-0.333, -0.333, 12.0]
my huber 1.9167  vs torch's 1.9167

Read the gradient line — it is the whole lesson

The third sample got weight 3 and has the biggest error, and its gradient is 12.0 against 0.333 for the others. Training pressure now concentrates exactly where you declared it should. You changed nothing about the model, the optimizer, or the loop. The marking scheme alone redirected the learning.

The last line is the professional habit: check your loss against a trusted one on the cases where they should agree. Huberish with threshold 1.0 must equal F.huber_loss with delta=1.0, and it does, to four decimals. Only after that do you trust it where they differ.

Function or class?

A plain function is enough for a loss with no settings. Reach for nn.Module when it carries configuration (the threshold above) or state — same registration machinery as any module. There is no performance difference; nn.MSELoss is a thin wrapper over F.mse_loss.

The three rules that keep autograd alive

  1. Stay inside tensor operations. torch.where, .abs(), .mean() — every step above is differentiable, so loss.backward() flows through automatically. No extra code buys you gradients; they come from the building blocks.
  2. Never round-trip through Python or NumPy. loss.item(), .detach(), .numpy(), or an if on a tensor's value inside the computation cuts the graph. The loss still computes — it silently stops teaching. (torch.where is the differentiable substitute for if.)
  3. Return a scalar. .mean() over the batch, conventionally. .sum() also works but couples the gradient's size to the batch size, which quietly changes your effective learning rate.

Common mistakes

A loss that ignores the prediction still runs. Botch the broadcasting — say, shapes (3,) versus (3,1) — and you average a 3x3 mistake matrix without an error. Assert pred.shape == target.shape at the top; this one line has saved careers.

Non-differentiable ambitions. Accuracy, F1, "count of correct" are steps, not slopes — gradients are zero almost everywhere. You need a smooth stand-in that tracks the metric: that is exactly what cross-entropy is for accuracy.

Numerical cliffs. torch.log(x) explodes as x nears zero. Clamp (x.clamp(min=1e-8)) or, for probabilities, keep everything in log-space — the reason F.cross_entropy takes raw scores, not softmax outputs.

Testing only on the model. Test the loss alone first, on tiny hand-made tensors where you can compute the right answer on paper — like the 36.1667 above: (0.25 + 0.25 + 3·36)/3.

Try it yourself

Write overstock_loss: squared error, but multiply by 4 whenever pred > target. Feed it predictions straddling the target and confirm the gradient is steeper on the overstocked side.

What to learn next

  • Transfer learning — pointing a pretrained model at your problem, custom loss included.
  • Loss functions — the standard schemes worth checking yours against.
  • RLHF — what happens when the objective cannot be written down at all.

Researcher — Mathematics and papers.

The loss defines the estimator

Training minimises empirical risk $\frac{1}{N}\sum_i \ell(f_\theta(x_i), y_i)$, and the choice of $\ell$ fixes what the optimum estimates. Squared error is minimised by the conditional mean; absolute error by the conditional median; the pinball loss

$$ \ell_\tau(y, \hat{y}) = \max\big(\tau (y - \hat{y}),\; (\tau - 1)(y - \hat{y})\big) $$

by the conditional $\tau$-quantile — where $\tau \in (0,1)$ is the target quantile and $y, \hat y$ truth and prediction. Asymmetric business costs are therefore often better served by choosing the statistic (a high quantile for stockouts) than by ad-hoc weighting. Classification losses are surrogates: 0-1 loss is intractable, and hinge, logistic, and cross-entropy are its convex upper bounds; Bartlett et al. (2006), Convexity, Classification, and Risk Bounds, gives the calibration conditions under which minimising the surrogate minimises the true risk.

Gradient shaping, read as design

The gradient of the loss w.r.t. the prediction is the per-sample teaching signal, and celebrated losses are best understood by that gradient:

  • Focal loss (Lin et al., 2017): $\ell = -(1-p_t)^\gamma \log p_t$, with $p_t$ the probability on the true class and $\gamma$ the focusing exponent — the factor $(1-p_t)^\gamma$ decays the gradient of already-easy examples, rebalancing dense detection where easy negatives outnumber hard positives thousands to one.
  • Label smoothing (Szegedy et al., 2016): cross-entropy against $(1-\epsilon)$ one-hot plus $\epsilon/K$ uniform; caps the gradient's demand for infinite logits, improving calibration.
  • Huber (1964, from robust statistics): gradient saturates at $\pm\delta$, so outliers vote with bounded strength — the torch.where construction in the developer block is its literal transcription.

When a desired objective is genuinely non-differentiable (ranking metrics, BLEU, human preference), the escape routes are smooth relaxations (soft ranking), policy-gradient estimators (REINFORCE — score-function gradients need only the reward, not its derivative), or learning a differentiable proxy of the metric — the path that leads to RLHF, covered in RLHF.

Implementation notes at the edge

Reductions interact with distributed training: with data parallelism, mean over the local batch then mean over replicas equals the global mean only for equal shard sizes — uneven final batches bias it, one reason drop_last appears in reference training scripts. Mixed precision adds a constraint: losses composing exp/log should use fused, stabilised primitives (logsumexp, F.cross_entropy on logits) since fp16 overflows at 65504. For fully custom backward behaviour — non-standard gradients, straight-through estimators — subclass torch.autograd.Function with explicit forward/backward, and validate with torch.autograd.gradcheck against finite differences in double precision. That tool is the final word on "is my gradient right".

Reading

  • Lin et al. (2017), Focal Loss for Dense Object Detection.
  • Bartlett, Jordan, McAuliffe (2006), Convexity, Classification, and Risk Bounds.
  • Szegedy et al. (2016), Rethinking the Inception Architecture — label smoothing's origin.

What to learn next

  • Transfer learning — pointing a pretrained model at your problem, custom loss included.
  • Loss functions — the standard schemes worth checking yours against.
  • RLHF — what happens when the objective cannot be written down at all.