Losses with a fused backward
Overview
Section titled “Overview”| Module | L0.3 · build · Python · Pass 2 · 3 to 4 h, plus your graded tests (rung R2) |
| You build | python/tinyllm/autograd/losses.py: cross_entropy (ignore_index, label smoothing, three reductions), mse, bce_with_logits (pos_weight); and your own tests in python/tests/l0-3-losses/ |
| Contract | course/contracts/py/tinyllm/autograd/losses.pyi |
| Tests | course/tests/L0.3/ (what they check: section 4), golden values from torch 2.14 in course/fixtures/L0.3/losses_torch.npz; your tests are graded by mutation, threshold 0.60 |
| Needs | L0.1 from_op · M09.2 stable log_softmax · M08.3 cross_entropy_vjp (the differential test) · M11.1 cross-entropy (the differential test) · reading: craft.03 how your tests are graded (or --ref-deps) |
| Used by | L0.5 every training step of the course computes one of these · later: L11.1, L12.1, L2.2, L3.6, L4.1, L5.5, L6.1, L6.2, L6.3, L6.5 |
| Milestone | MS-L0 (the bigram and the digits MLP train on cross_entropy) |
| Optional depth | Goodfellow, Bengio, Courville, Deep Learning (free online), sections 6.2.2 and 7.5; Szegedy et al., “Rethinking the Inception Architecture” (2016), section 7 (label smoothing) |
Key Takeaways
Section titled “Key Takeaways”- Softmax cross-entropy is one graph node: its forward uses
log_softmax, and its gradient is the closed form (test_matches_m08_cross_entropy_vjp,test_gradcheck_losses). - Padding rows (
ignore_index) add no loss, get no gradient, and do not count in the mean; a batch of only padding is loss 0, notnan(test_ignore_index_excluded,test_all_ignored_is_zero). - Label smoothing spreads over every class, the target included, so a row’s loss is the cross-entropy against the smoothed target (
test_hand_example_label_smoothing,test_rows_match_m11_cross_entropy). - Stable forms keep logits of finite:
log_softmaxfor cross-entropy, softplus for binary cross-entropy (test_large_logits_stay_finite). - Your own tests, written from the names in section 4, must kill at least 60% of the planted faults: the first grade of your tests in the course.
How to work this chapter
Section titled “How to work this chapter”ol start L0.3 # stubs losses.py into your repo; prints your test path and rungol tests L0.3 # the course tests: the exemplars your own tests imitateol check L0.3 # course tests first, then the mutation grade of your testsol mutate L0.3 # the full grade, cached by your test files' hashol check L0.3 --ref-deps # only if L0.1 or a math dependency is not passing yetol diff L0.3 # after passing: your code against the referenceRead craft.03 first if you have not: it explains what a mutant is, how the score is computed, and why your tests may import only the contract.
1. Why now
Section titled “1. Why now”L0.2 gives you log_softmax and gather, so cross-entropy could be written as three ops: -gather(log_softmax(z), t).mean(). It would be correct and slow, and it would be fragile. Backward through that chain keeps a [N, V] array alive at each step, and for the bigram (V = 256) or a real vocabulary (V = 49,152 for SmolLM2) that is most of the memory a training step uses. Padding makes it worse: batches of token windows have rows that must not count, and a mean that divides by the padded length silently shrinks the loss. Every training loop from here on (L0.5, L2.2, L6.1, the capstone) calls this file once per step, so the loss must be one node with a closed-form gradient, correct about padding, and finite for any logits.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| logits of row ( rows, classes) | float[N, V] (or [..., V]) | |
target class of row , or ignore_index | int[N] | |
| predicted distribution | float[V] | |
1 when row is kept ( ignore_index) | ||
| number of kept rows | int | |
| label smoothing, in | float | |
| smoothed target, the one-hot vector | float[V] | |
row loss, (M11.1) | float | |
| logit and target of binary cross-entropy, | float | |
pos_weight, the weight on positive examples | float, default 1 | |
| , | functions |
The fused gradient. For one row with a one-hot target, . Differentiate: . With a smoothed target, (because ), so . That is the whole backward: one subtraction on the probabilities, which the forward already computed as . M08.3 derived it; the course test compares your fused VJP with M08.3’s cross_entropy_vjp on the same inputs.
The forward is stable. (M09.2). Computing instead underflows: at softmax returns exactly in float64 and .
Reductions and padding. The kept mask zeroes ignored rows in both directions: their loss is 0 and their gradient row is 0. "none" returns the per-row losses (); "sum" is ; "mean" is that sum divided by , the number of kept rows. Dividing by would make a batch of 3 real rows and 5 padding rows report 3/8 of its true loss and train at 3/8 of the learning rate. When the mean is ; the course defines it as loss 0 with a zero gradient (torch returns nan, and one nan step poisons every weight). The upstream gradient scales the rows: by for the mean, by for the sum, row by row for "none".
Label smoothing. A target of exactly one class asks the model to push to 1, which needs . Smoothing moves of the target mass onto a uniform distribution over all classes, so and every other class gets . The contract writes the same loss as . Spreading over classes (excluding the target) is a different, also published, variant; torch and this course use .
Mean squared error. over all elements; the gradient is .
Binary cross-entropy on logits. For one logit and target : . Using and , this is
and never exponentiates a positive number. The gradient is , divided by the element count for the mean.
3. Worked example by hand
Section titled “3. Worked example by hand”Cross-entropy with a padding row. Logits for two rows over classes: row 0 is with target 2; row 1 has target ignore_index (-100), whatever its logits.
- , sum 6, so .
- .
- Row 1 is ignored: , , so the mean loss is .
- Gradient of row 0: , divided by . Row 1: zeros.
Label smoothing. Same row 0, : each class gets and the target keeps more, so .
- .
- Gradient .
Both gradients sum to zero, as every softmax gradient must.
These are test_hand_example_cross_entropy and test_hand_example_label_smoothing.
4. The interface
Section titled “4. The interface”def cross_entropy(logits: Tensor, targets: ArrayLike, ignore_index: int = -100, label_smoothing: float = 0.0, reduction: Literal["mean", "sum", "none"] = "mean") -> Tensordef mse(pred: Tensor, target: ArrayLike) -> Tensordef bce_with_logits(logits: Tensor, targets: ArrayLike, pos_weight: Optional[ArrayLike] = None) -> TensorEach returns one from_op node whose VJP is the closed form above. Validate targets before anything else: integers, shape logits.shape[:-1], each in or equal to ignore_index.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example_cross_entropy | unit | section 3: loss , gradient , zeros on the ignored row | you and the test agree on the definition |
test_hand_example_label_smoothing | unit | section 3: , loss and | the smoothing every LM config can turn on |
test_matches_torch | golden | loss and gradient equal torch.nn.functional on every option | training loops ported from torch |
test_gradcheck_losses | gradcheck | the fused VJPs against frozen central differences | a closed form off by a factor fails |
test_matches_m08_cross_entropy_vjp | differential | your fused gradient equals M08.3’s matrix | the derivation and the code agree |
test_rows_match_m11_cross_entropy | differential | each smoothed row loss is M11.1’s | the definition, computed the slow way |
test_ignore_index_excluded | property | appending ignored rows changes neither the loss nor the kept gradients | padded batches (L0.6 windows, L12.1 packing) |
test_all_ignored_is_zero | boundary | only padding gives loss 0 and a zero gradient | one nan step ruins a run |
test_large_logits_stay_finite | boundary | logits of and bce logits of 1000 stay finite | trained logits are large |
test_reductions_agree | property | none, sum, mean agree with each other | evaluation sums none (L6.7) |
test_float32_stays_float32 | boundary | float32 logits give float32 loss and gradient | in-place float32 updates |
test_rejects_bad_inputs | boundary | out-of-range or float targets, unknown reduction, bad , shape mismatches | data and config bugs fail here |
Your graded tests (rung R2)
Section titled “Your graded tests (rung R2)”Write them in python/tests/l0-3-losses/ (any test_*.py file there). The names and what each must show are given; the bodies are yours. Import only names the contract declares (tinyllm.autograd.losses, tinyllm.autograd.tensor), plus numpy, pytest, and the standard library: your tests run against the reference with one fault planted at a time, so they must not depend on anything else of yours. Compute expected values by hand (this is rung R1’s habit: the number in the assertion comes from you, not from running the code).
test_hand_example_matches_section_3: logits , target 2: loss , gradient .test_ignored_rows_get_no_loss_and_no_gradient: a row whose target isignore_indexadds no loss and gets a zero gradient.test_mean_divides_by_kept_rows: two kept rows and one ignored row: the mean is the kept sum over 2.test_all_ignored_batch_is_zero: only padding gives loss 0 and a zero gradient, nevernan.test_label_smoothing_hand_value: on : the loss and of section 3.test_large_logits_are_finite: logits of and bce logits of 1000 give finite losses and gradients.test_out_of_range_target_raises: targets outside , other thanignore_index, are aValueError.test_mse_and_bce_gradients_by_finite_differences: both gradients match central differences you write; bce withpos_weight2 at , is .
ol check L0.3 passes when the course tests pass and your tests kill at least 60% of the planted faults (ol mutate L0.3 prints the score and the survivors).
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. counting ignored rows in the mean, or giving them gradient | padded batches report a smaller loss; padding logits train toward a fake target | test_ignore_index_excluded, test_hand_example_cross_entropy (mutants s01, s02) |
| 2. trusting targets | -1 silently trains the last class (numpy indexes from the end) | test_rejects_bad_inputs (mutant s07) |
| 3. dividing by | an all-padding batch gives nan, then every weight is nan | test_all_ignored_is_zero (mutant s03) |
| 4. and | inf loss at logits of | test_large_logits_stay_finite (mutants s04, s09) |
| 5. smoothing over classes, or summing instead of averaging | disagrees with torch; the loss grows with the vocabulary size | test_hand_example_label_smoothing (mutants s05, s06) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Forward | L12.1 | Registered module relationship. |
| Forward | L11.1 | Registered call site uses this module. |
| Forward | L2.2 | Registered call site uses this module. |
| Forward | L3.6 | Registered call site uses this module. |
| Forward | L4.1 | Registered call site uses this module. |
| Forward | L5.5 | Registered call site uses this module. |
| Forward | L6.1 | Registered call site uses this module. |
| Forward | L6.2 | Registered call site uses this module. |
| Forward | L6.3 | Registered call site uses this module. |
| Forward | L6.5 | Registered call site uses this module. |
| Direction | Module | How it uses this |
|---|---|---|
| Back | L0.1 | each loss is one from_op node |
| Back | M09.2 | log_softmax for the stable forward |
| Back | M08.3 | cross_entropy_vjp, the derivation the differential test compares with |
| Back | M11.1 | , the definition of a smoothed row’s loss |
| Forward | L0.5 | the bigram and the MLPs of the training-loop tests minimize cross_entropy, mse, and bce_with_logits |
If you skip this module, ol check L0.5 stops with L0.5 needs L0.3: build it, or rerun with --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
fused cross_entropy | Liger Kernel fused_linear_cross_entropy | fuses the output projection too: the [N, V] logits are never materialized, chunked over rows | liger_kernel/ops/fused_linear_cross_entropy.py |
ignore_index | PyTorch nll_loss | the same semantics, plus per-class weights | aten/src/ATen/native/LossNLL.cpp |
| label smoothing | torch.nn.CrossEntropyLoss(label_smoothing=...) | combined with class weights and probability targets | aten/src/ATen/native/Loss.cpp |
bce_with_logits | torch.nn.functional.binary_cross_entropy_with_logits | the same softplus form, vectorized on GPU | aten/src/ATen/native/Loss.cpp |