How Cross Entropy Loss Works for Classification: Inside Karpathy’s nn-zero-to-hero Implementation
Cross entropy loss measures the discrepancy between predicted logits and true class labels by converting logits to probabilities via softmax, computing the negative log-likelihood of the correct class, and averaging across the batch.
In the educational repository karpathy/nn-zero-to-hero, Andrej Karpathy provides a line-by-line breakdown of deep learning fundamentals. The makemore_part4_backprop.ipynb notebook contains a bit-exact manual implementation of cross entropy loss that mirrors PyTorch’s internal logic, making it an ideal reference for understanding how classification loss actually functions under the hood.
What Is Cross Entropy Loss?
Cross entropy loss quantifies how far a model’s predicted probability distribution diverges from the true distribution defined by the ground-truth labels. For a single training example with C classes, the loss is calculated as:
[ \mathcal{L} = -\log \big(p_{y}\big) ]
Here, p_y represents the model’s predicted probability for the correct class y. When processing a batch of N examples, the loss becomes the mean (or sum) of individual losses:
[ \mathcal{L}{\text{batch}} = -\frac{1}{N}\sum{i=1}^{N}\log\big(p_{i,,y_i}\big) ]
The value approaches zero as the model assigns higher probability to the correct class, increasing as confidence in wrong answers grows.
How PyTorch Computes Cross Entropy Loss
The torch.nn.functional.cross_entropy function (and its module counterpart nn.CrossEntropyLoss) performs three distinct operations internally:
- Logits → Probabilities: Applies the softmax function to convert raw model outputs into a valid probability distribution summing to 1.
- Log-Probability: Computes the natural logarithm of these probabilities for numerical stability.
- Negative Log-Likelihood: Extracts the log-probability corresponding to the target class and aggregates it across the batch.
This compound operation is more numerically stable than manually applying softmax followed by NLLLoss, as it avoids the risk of taking the logarithm of zero.
Manual Implementation in nn-zero-to-hero
The lectures/makemore/makemore_part4_backprop.ipynb file (lines 237‑245) contains a manual implementation that produces bit-exact gradients matching PyTorch’s autograd. This explicit version keeps every intermediate tensor to facilitate educational backpropagation exercises:
# logits from the final linear layer
logits = h @ W2 + b2
# ---- manual cross‑entropy (same as F.cross_entropy(logits, Yb)) ----
logit_maxes = logits.max(1, keepdim=True).values # 1️⃣ max for stability
norm_logits = logits - logit_maxes # 2️⃣ shift logits
counts = norm_logits.exp() # 3️⃣ exp → un‑normalized probs
counts_sum = counts.sum(1, keepdims=True) # 4️⃣ sum over classes
counts_sum_inv = counts_sum**-1 # 5️⃣ 1 / sum
probs = counts * counts_sum_inv # 6️⃣ soft‑max probabilities
logprobs = probs.log() # 7️⃣ log‑probabilities
loss = -logprobs[range(n), Yb].mean() # 8️⃣ NLL loss (mean over batch)
Understanding the Code Steps
Each line in the manual implementation corresponds to a specific mathematical operation:
logits.max(1, keepdim=True).values: Identifies the maximum logit in each row to prevent numerical overflow during exponentiation.norm_logits = logits - logit_maxes: Subtracts the max from each logit, shifting the distribution without changing the relative probabilities (softmax is shift-invariant).counts = norm_logits.exp(): Computes the unnormalized exponentials, effectively measuring the "confidence" for each class.probs = counts * counts_sum_inv: Normalizes the counts into probabilities by dividing by the sum, implementing the softmax function.logprobs[range(n), Yb]: Indexes into the log-probability tensor using integer targetsYbto select only the probabilities of the correct classes..mean(): Averages the negative log-probabilities across the batch dimension.
Numerical Stability Tricks
The subtraction of logit_maxes before exponentiation is critical for numerical stability. Without this offset, large positive logits would produce enormous exponential values, risking floating-point overflow. By centering the logits around zero (subtracting the maximum), the largest value entering exp() becomes zero, ensuring the computation remains numerically safe while preserving the relative probabilities.
Three Ways to Calculate Cross Entropy in PyTorch
While the manual implementation educational value is high, production code typically uses these optimized PyTorch APIs:
1. Using torch.nn.functional.cross_entropy
import torch
import torch.nn.functional as F
logits = model(inputs) # shape: (N, C)
loss = F.cross_entropy(logits, targets)
loss.backward()
2. Using the nn.CrossEntropyLoss module
import torch.nn as nn
criterion = nn.CrossEntropyLoss(reduction='mean') # or 'sum'
logits = model(inputs)
loss = criterion(logits, targets)
loss.backward()
3. Manual implementation (educational)
def manual_cross_entropy(logits, targets):
# Stability trick
logits = logits - logits.max(dim=1, keepdim=True).values
# Softmax
probs = logits.exp()
probs = probs / probs.sum(dim=1, keepdim=True)
# Negative log-likelihood
log_probs = probs.log()
loss = -log_probs[torch.arange(logits.shape[0]), targets].mean()
return loss
All three methods yield mathematically equivalent gradients, though the manual version explicitly surfaces every intermediate tensor for inspection.
Summary
- Cross entropy loss combines softmax activation with negative log-likelihood to penalize incorrect predictions based on confidence.
- The
nn-zero-to-herorepository implements this manually inmakemore_part4_backprop.ipynbto demonstrate backpropagation through every operation. - Numerical stability is achieved by subtracting the maximum logit before exponentiation.
- PyTorch’s
F.cross_entropyis the preferred production method, as it handles the softmax and log operations in a fused, numerically stable kernel.
Frequently Asked Questions
Why subtract the maximum logit before applying softmax?
Subtracting the maximum logit from all logits prevents numerical overflow when computing exponentials. Since softmax is shift-invariant ($e^{x+c}/\sum e^{x+c} = e^x/\sum e^x$), this operation doesn't change the final probabilities but keeps floating-point values in a safe range.
What's the difference between CrossEntropyLoss and NLLLoss?
CrossEntropyLoss accepts raw logits and internally applies log-softmax before computing the negative log-likelihood. NLLLoss expects log-probabilities (already passed through log_softmax) as input. Using CrossEntropyLoss is generally preferred as it handles the numerical stability tricks automatically.
Does cross entropy loss work with one-hot encoded targets?
PyTorch’s F.cross_entropy and nn.CrossEntropyLoss expect integer class indices as targets (shape [N]), not one-hot vectors. If you have one-hot labels, use torch.argmax to convert them to indices before passing them to the loss function.
How does the reduction parameter affect training?
The reduction parameter (default 'mean') determines how the loss aggregates across the batch. Setting it to 'mean' averages the loss, making the gradient magnitude invariant to batch size. Setting it to 'sum' adds losses together, meaning gradient scales with batch size. 'none' returns per-example losses without reduction, useful for custom weighting schemes.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →