# How Cross Entropy Loss Works for Classification: Inside Karpathy’s nn-zero-to-hero Implementation

> Understand cross entropy loss for classification. Learn how it measures discrepancies by converting logits to probabilities and calculating negative log-likelihood using Karpathy's nn-zero-to-hero implementation.

- Repository: [Andrej/nn-zero-to-hero](https://github.com/karpathy/nn-zero-to-hero)
- Tags: deep-dive
- Published: 2026-05-23

---

**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:

1. **Logits → Probabilities**: Applies the **softmax** function to convert raw model outputs into a valid probability distribution summing to 1.
2. **Log-Probability**: Computes the natural logarithm of these probabilities for numerical stability.
3. **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:

```python

# 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 targets `Yb` to 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`**

```python
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**

```python
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)**

```python
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-hero` repository implements this manually in `makemore_part4_backprop.ipynb` to demonstrate backpropagation through every operation.
- **Numerical stability** is achieved by subtracting the maximum logit before exponentiation.
- PyTorch’s `F.cross_entropy` is 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.