Loss functions
Measure how wrong a network is with mean squared error and cross-entropy, computed stably from logits.
- Choose a loss for regression and for classification
- Compute mean squared error, binary cross-entropy and categorical cross-entropy
- Explain why cross-entropy punishes confident mistakes, and compute it stably from logits
Training needs a single number that says how wrong the network is: the loss. Training is nothing more than nudging the weights to make that number smaller - so the loss defines what “good” means.
| Task | Output | Usual loss |
|---|---|---|
| Regression (a price, a temperature) | one number | mean squared error |
| Yes/no classification | one probability (sigmoid) | binary cross-entropy |
| One of K classes | K probabilities (softmax) | categorical cross-entropy |
Mean squared error
Squaring makes every error positive and punishes big misses much more than small ones - which also makes MSE sensitive to outliers. Mean absolute error (MAE), the average of , is gentler with them; the Huber loss blends the two.
Cross-entropy
For classification, the network outputs probabilities, and cross-entropy is the negative log of the probability it gave the correct class:
Gave the right class 0.9? Loss 0.105. Gave it 0.5? Loss 0.693. Confidently wrong, with 0.01? Loss 4.6. The log makes confident mistakes very expensive - exactly the pressure you want. For a yes/no output with probability p and label y ∈ {0, 1}, the same idea is binary cross-entropy: .
1import numpy as np
2
3predicted = np.array([2.5, 0.0, 2.0, 8.0])
4actual = np.array([3.0, -0.5, 2.0, 7.0])
5print(f"MSE {np.mean((predicted - actual) ** 2):.3f} MAE {np.mean(np.abs(predicted - actual)):.3f}")
6
7for p in [0.9, 0.5, 0.1, 0.01]:
8 print(f"p(correct) = {p:<4} cross-entropy = {-np.log(p):.3f}")MSE 0.375 MAE 0.500 p(correct) = 0.9 cross-entropy = 0.105 p(correct) = 0.5 cross-entropy = 0.693 p(correct) = 0.1 cross-entropy = 2.303 p(correct) = 0.01 cross-entropy = 4.605
From logits, stably
Taking np.log(softmax(z)) can hit log(0) = -inf when a probability underflows. Instead, combine the two with the log-sum-exp trick:
where is the largest logit. This is why frameworks offer losses that take raw logits (PyTorch’s CrossEntropyLoss, Keras’ from_logits=True) - pass logits, not softmax output, and you get stability for free.
1import numpy as np
2
3def cross_entropy(logits, labels):
4 shifted = logits - logits.max(axis=1, keepdims=True)
5 log_sum_exp = np.log(np.exp(shifted).sum(axis=1))
6 correct = shifted[np.arange(len(labels)), labels]
7 return np.mean(log_sum_exp - correct)
8
9logits = np.array([[2.0, 1.0, 0.1],
10 [0.5, 2.5, 0.3],
11 [1000.0, 0.0, -1000.0]])
12labels = np.array([0, 1, 2])
13print(f"{cross_entropy(logits, labels):.3f}")666.879
The third example is wildly, confidently wrong (it bet everything on class 0, the answer was class 2), so the average loss is huge - but finite. The naive version would print inf. Note the fancy indexing shifted[np.arange(N), labels], which picks each row’s correct-class logit.
Try it
Pick the loss
Choose the loss you’d start with for each task.
“Predict tomorrow’s temperature”
“Spam or not spam?”
“Which of 10 digits is in the image?”
“Estimate a doodle’s ink coverage (0-1)”
“Is this photo tagged “dragon”? (and separately “castle”?)”
“Which of 50,000 tokens comes next?”
Key takeaways
The loss is the single number training minimizes; it defines what “good” means.
MSE for regression (sensitive to outliers); MAE or Huber when outliers matter.
Cross-entropy, −log p(correct), for classification - it heavily punishes confident mistakes.
Compute cross-entropy from logits with log-sum-exp; frameworks’ logit losses do this for you.
Lesson quiz
7 questions · pass with 5 correct · up to 50 XP
Passing this quiz completes the lesson and keeps your streak going. Questions you miss come back in review sessions later.
Practice: write Python
Write Python in the editor and run it against sample inputs. Python runs locally in your browser using a WebAssembly runtime.
Cross-entropy from logits
Each input line is label logit0 logit1 ... for one example. Compute each example’s cross-entropy stably from the logits (log-sum-exp), print it to 4 decimals, then print mean= and the average to 4 decimals. Some logits are enormous.
- Huge logits
- Ordinary logits
Python runs in a sandboxed browser worker with a 60 second time limit. Its runtime loads from the Pyodide CDN; your code stays in this browser.
Outliers: MSE vs MAE
Each input line is predicted actual. Print the MSE and the MAE to 2 decimals, then the same two losses without the single worst example (the one with the largest absolute error), in the format all: MSE 4.25 MAE 1.50 / without worst: MSE 0.33 MAE 0.33. Notice which loss the outlier dominates.
- One outlier
- No outlier
Python runs in a sandboxed browser worker with a 60 second time limit. Its runtime loads from the Pyodide CDN; your code stays in this browser.
Questions about this lesson
Stuck? Ask. Figured something out? Share it. Explaining is one of the best ways to learn.
Loading posts…